Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 34 additions & 9 deletions tests-unit/assets_test/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@
import pytest
import requests

from .helpers import assert_hash_fields_consistent


def pytest_addoption(parser: pytest.Parser) -> None:
"""
Expand All @@ -28,6 +30,18 @@ def pytest_addoption(parser: pytest.Parser) -> None:
default=os.environ.get("ASSETS_TEST_DB_URL"),
help="SQLAlchemy DB URL (e.g. sqlite:///path/to/db.sqlite3)",
)
parser.addoption(
"--enable-asset-hashing",
action="store_true",
help="Start the assets subprocess with hash-mode behavior enabled.",
)


def pytest_configure(config: pytest.Config) -> None:
config.addinivalue_line(
"markers",
"hashing_on: exercises the subprocess harness with --enable-asset-hashing",
)


def _free_port() -> int:
Expand Down Expand Up @@ -103,8 +117,7 @@ def comfy_url_and_proc(comfy_tmp_base_dir: Path, request: pytest.FixtureRequest)
if not (comfy_root / "main.py").is_file():
raise FileNotFoundError(f"main.py not found under {comfy_root}")

proc = subprocess.Popen(
args=[
command = [
sys.executable,
"main.py",
f"--base-directory={str(comfy_tmp_base_dir)}",
Expand All @@ -115,7 +128,19 @@ def comfy_url_and_proc(comfy_tmp_base_dir: Path, request: pytest.FixtureRequest)
"--port",
str(port),
"--cpu",
],
]
if (
request.config.getoption("--enable-asset-hashing")
or "hashing_on" in request.config.getoption("markexpr")
or any(
item.get_closest_marker("hashing_on")
for item in request.session.items
)
):
command.append("--enable-asset-hashing")

proc = subprocess.Popen(
args=command,
stdout=out_log,
stderr=err_log,
cwd=str(comfy_root),
Expand Down Expand Up @@ -190,8 +215,9 @@ def _post_multipart_asset(
@pytest.fixture
def make_asset_bytes() -> Callable[[str, int], bytes]:
# Salt content per test so it never collides with assets left over from
# earlier tests. Delete is now always a soft delete (content is preserved),
# so the suite can no longer rely on hard-deleting content for isolation.
# earlier tests. Delete hard-deletes the record but preserves content
# (content rows and files are untouched), so the suite cannot rely on delete
# removing content for isolation.
# Deterministic within a test: the same (name, size) yields the same bytes.
salt = uuid.uuid4().bytes

Expand Down Expand Up @@ -236,9 +262,9 @@ def seeded_asset(request: pytest.FixtureRequest, http: requests.Session, api_bas
if tags is None:
tags = ["models", "model_type:checkpoints", "unit-tests", "alpha"]
meta = {"purpose": "test", "epoch": 1, "flags": ["x", "y"], "nullable": None}
# Unique content per test so the seed always creates a fresh asset (201).
# Delete is now always a soft delete, so content from a prior test survives
# and would otherwise dedup this upload into an existing asset (200).
# Unique content per test so the seed also owns a fresh content row. Delete
# preserves content (only the record is hard-deleted), so content from a
# prior test survives and would otherwise be reused behind this record.
content = uuid.uuid4().bytes + b"A" * (4096 - 16)
files = {"file": (name, content, "application/octet-stream")}
form_data = {
Expand All @@ -249,7 +275,6 @@ def seeded_asset(request: pytest.FixtureRequest, http: requests.Session, api_bas
r = http.post(api_base + "/api/assets", files=files, data=form_data, timeout=120)
body = r.json()
assert r.status_code == 201, body
from helpers import assert_hash_fields_consistent
assert_hash_fields_consistent(body)
return body

Expand Down
122 changes: 121 additions & 1 deletion tests-unit/assets_test/helpers.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,124 @@
"""Helper functions for assets integration tests."""
from __future__ import annotations

import json
import time
import uuid
from collections.abc import Iterator, Mapping
from dataclasses import dataclass
from datetime import datetime, timedelta
from typing import NotRequired, TypeAlias, TypedDict

import pytest
import requests
from aiohttp import web
from aiohttp.test_utils import make_mocked_request
from sqlalchemy import Engine, create_engine
from sqlalchemy.orm import Session

from app.assets.api import routes
from app.assets.database.models import Asset
from app.assets.database.queries.records import create_content, create_record
from app.database.models import Base


class AssetItem(TypedDict):
id: str
name: str
preview_id: NotRequired[str]


class AssetListBody(TypedDict):
assets: list[AssetItem]
total: int
has_more: bool
next_cursor: NotRequired[str]


class ErrorItem(TypedDict):
code: str


class ErrorBody(TypedDict):
error: ErrorItem


@dataclass(frozen=True, slots=True)
class RecordSeed:
name: str
tags: tuple[str, ...] = ()
size_bytes: int = 0


RouteDatabase: TypeAlias = tuple[Engine, Session]


@pytest.fixture(autouse=True)
def autoclean_unit_test_assets() -> Iterator[None]:
yield


@pytest.fixture
def route_database(monkeypatch: pytest.MonkeyPatch) -> Iterator[RouteDatabase]:
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(engine)
monkeypatch.setattr(routes, "create_session", lambda: Session(engine))
monkeypatch.setattr(routes, "_ASSETS_ENABLED", True)
with Session(engine) as session:
yield engine, session
engine.dispose()


def seed_record(session: Session, seed: RecordSeed) -> Asset:
content = create_content(
session,
path=f"/output/{uuid.uuid4()}-{seed.name}",
size_bytes=seed.size_bytes,
)
return create_record(
session,
content_id=content.id,
name=seed.name,
mime_type="image/png",
tags=seed.tags,
)


@pytest.fixture
def sortable_record_ids(route_database: RouteDatabase) -> tuple[str, str]:
_, session = route_database
older = seed_record(session, RecordSeed("z.png", ("sort-case",), 100))
newer = seed_record(session, RecordSeed("a.png", ("sort-case",), 200))
base_time = datetime(2026, 1, 1)
older.created_at = base_time
newer.created_at = base_time + timedelta(days=1)
newer.updated_at = base_time
older.updated_at = base_time + timedelta(days=1)
older.last_access_time = base_time
newer.last_access_time = base_time + timedelta(days=1)
session.commit()
return newer.id, older.id


async def request_assets(query: str = "") -> web.StreamResponse:
suffix = f"?{query}" if query else ""
return await routes.list_assets_route(
make_mocked_request("GET", f"/api/assets{suffix}")
)


def asset_list_body(response: web.StreamResponse) -> AssetListBody:
assert isinstance(response, web.Response)
body = response.body
assert isinstance(body, bytes | bytearray)
return json.loads(body)


def error_body(response: web.StreamResponse) -> ErrorBody:
assert isinstance(response, web.Response)
body = response.body
assert isinstance(body, bytes | bytearray)
return json.loads(body)


def trigger_sync_seed_assets(session: requests.Session, base_url: str) -> None:
Expand All @@ -28,7 +145,10 @@ def get_asset_filename(asset_hash: str, extension: str) -> str:
return asset_hash.removeprefix("blake3:") + extension


def assert_hash_fields_consistent(body: dict, expected_hash: str | None = None) -> None:
def assert_hash_fields_consistent(
body: Mapping[str, str | None],
expected_hash: str | None = None,
) -> None:
"""Assert hash and asset_hash invariants on an Asset response.

Both must be present or both absent (so a regression that drops only one
Expand Down
12 changes: 12 additions & 0 deletions tests-unit/assets_test/queries/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,18 @@
from app.assets.database.models import Base


@pytest.fixture(scope="session", autouse=True)
def assert_asset_metadata_tables():
assert set(Base.metadata.tables) == {
"assets",
"asset_contents",
"asset_meta",
"asset_tags",
"tags",
"asset_system_state",
}


@pytest.fixture
def session():
"""In-memory SQLite session for fast unit tests."""
Expand Down
127 changes: 24 additions & 103 deletions tests-unit/assets_test/queries/test_asset_reference_keyset.py
Original file line number Diff line number Diff line change
@@ -1,112 +1,33 @@
"""Keyset-pagination tiebreaker tests for list_references_page.

When multiple rows share the same primary sort value (e.g. four assets
created in the same microsecond), the secondary `ORDER BY id` is what keeps
keyset pagination from losing or repeating rows. This file exercises that
branch directly against an in-memory SQLite session — engineering identical
timestamps via HTTP is unreliable enough that we work at the query layer.
"""
import uuid
from datetime import datetime

import pytest
from sqlalchemy.orm import Session

from app.assets.database.models import Asset, AssetReference
from app.assets.database.queries.asset_reference import list_references_page
from app.assets.database.queries import create_content, create_record, list_records_page
from app.assets.database.queries.records import RecordCursorBoundary, RecordPageSpec


def _make_ref(session: Session, created_at: datetime, name: str, owner: str = "") -> AssetReference:
asset = Asset(hash=f"blake3:{uuid.uuid4().hex}", size_bytes=1024)
session.add(asset)
def test_record_keyset_cursor_pages_in_creation_order(session: Session) -> None:
records = [
create_record(session, create_content(session, f"/output/{name}").id, name)
for name in ("one.png", "two.png", "three.png")
]
for index, record in enumerate(records, start=1):
record.id = f"00000000-0000-0000-0000-{index:012d}"
session.flush()
ref = AssetReference(
id=str(uuid.uuid4()),
asset_id=asset.id,
owner_id=owner,
name=name,
file_path=f"/tmp/{name}",
created_at=created_at,
updated_at=created_at,
last_access_time=created_at,
is_missing=False,
)
session.add(ref)
return ref


@pytest.mark.parametrize("order", ["desc", "asc"])
def test_tiebreaker_walks_duplicate_sort_values(session: Session, order: str):
"""Four rows with the SAME created_at must paginate cleanly under cursor
mode — no row dropped, no row repeated, despite the primary sort column
being non-discriminating.
"""
shared_ts = datetime(2024, 5, 20, 12, 0, 0) # naive UTC, like the DB stores
refs = [_make_ref(session, shared_ts, f"tie_{i}.png") for i in range(4)]
session.commit()

expected_ids = sorted([r.id for r in refs], reverse=(order == "desc"))

# Walk the cursor by hand: page size 2, take 3 pages (2 + 2 + 0).
seen: list[str] = []
after_value = None
after_id = None
for _ in range(4): # generous loop bound; ought to be 2 iterations
page, _tag_map, _total = list_references_page(
session,
limit=2,
sort="created_at",
order=order,
after_cursor_value=after_value,
after_cursor_id=after_id,
)
if not page:
break
seen.extend(p.id for p in page)
# Use the last row's (created_at, id) as the next cursor input.
last = page[-1]
after_value, after_id = last.created_at, last.id
if len(page) < 2:
break

assert seen == expected_ids, (
f"keyset tiebreaker failed for order={order}: expected {expected_ids}, got {seen}"
first_page, _, _ = list_records_page(
session,
RecordPageSpec(limit=2, order="asc"),
)


def test_tiebreaker_no_duplicates_under_mixed_collisions(session: Session):
"""Some rows share a timestamp, some don't. The cursor must still walk
every row exactly once regardless of where ties sit relative to a
page boundary."""
t1 = datetime(2024, 5, 20, 12, 0, 0)
t2 = datetime(2024, 5, 20, 12, 0, 1)
layout = [t1, t1, t1, t2, t2] # three rows at t1, two at t2
refs = [_make_ref(session, ts, f"mix_{i}.png") for i, ts in enumerate(layout)]
session.commit()

all_ids = {r.id for r in refs}
seen_set: set[str] = set()
seen_list: list[str] = []
after_value = None
after_id = None
for _ in range(6):
page, _, _ = list_references_page(
session,
boundary_record = first_page[-1]
second_page, _, _ = list_records_page(
session,
RecordPageSpec(
limit=2,
sort="created_at",
order="desc",
after_cursor_value=after_value,
after_cursor_id=after_id,
)
if not page:
break
for p in page:
assert p.id not in seen_set, f"duplicate row {p.id} appeared in cursor walk"
seen_set.add(p.id)
seen_list.append(p.id)
last = page[-1]
after_value, after_id = last.created_at, last.id
if len(page) < 2:
break
order="asc",
after=RecordCursorBoundary(
value=boundary_record.created_at,
id=boundary_record.id,
),
),
)

assert seen_set == all_ids, f"missing rows: expected {all_ids}, got {seen_set}"
assert [record.id for record in first_page + second_page] == [record.id for record in records]
Loading
Loading