diff --git a/CHANGELOG.md b/CHANGELOG.md index 455d2b5..4814246 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,18 @@ the fuller account of each version, including verification notes. ### Added +- `models.list()` and `models.schema()`, so you can discover Comfy Router models from Python + as the TypeScript SDK already can. `list(cursor=, limit=, timeout=)` returns an iterable that + walks the catalog (`GET /v2/models`), following `next_cursor` while `has_more` is true, and + yields `CatalogModel` entries (`id`, `provider`, `model`, `billing`). `list(...).page()` + returns one `ModelPage` (`data`, `has_more`, `next_cursor`, `limit`, `request_id`). + `schema(model, etag=, timeout=)` reads `GET /v2/models/{provider}/{model}/openapi.json` into a + `SchemaResult`. With `etag=`, it sends `If-None-Match`, and a `304` returns `unchanged=True` + with `document=None` rather than raising. Both methods use the Router host and the client's + credential, raise the same typed Router exceptions as `models.run`, retry under the client's + policy (a keyless read also retries a `5xx` or read timeout whenever that policy retries at + all), and default to a 30-second timeout. `AsyncComfy` has the same methods + (`async for ... in client.models.list()`, `await client.models.schema(...)`). - `RouterRunResult.credits_used` — what Comfy Router reported a run cost, lifted from the `X-Comfy-Credits-Used` response header onto what `models.run_detailed()` returns. It is a price rather than a settled ledger entry, absent means "not reported" and never "free", and diff --git a/README.md b/README.md index 9a83981..33c6e39 100644 --- a/README.md +++ b/README.md @@ -403,7 +403,9 @@ two variables. `base_url` and `timeout` are a read-only view of that configuration; model operations are added to this namespace as they land. There are two ways to run a model on it — `run`, which waits, and `submit`, which queues — and they send -the same request. +the same request. Two more read-only calls tell you what to run before you run +it: `list`, the model catalog, and `schema`, one model's input and output +schemas. ### `models.run` — one call, one result @@ -484,9 +486,10 @@ Path("hello.mp3").write_bytes(result.content) ``` `BinaryResult` is importable from `comfy_sdk` for exactly this `isinstance` -check. Which shape a given model returns is in its own contract — `GET -/v2/models/{provider}/{model}/openapi.json`, whose `200` is `application/json` -for a JSON model and `*/*` with `format: binary` for a bytes one. +check. Which shape a given model returns is in its own contract — +[`models.schema`](#modelsschema--what-a-model-takes-and-what-it-returns) +fetches it — whose `200` is `application/json` for a JSON model and `*/*` with +`format: binary` for a bytes one. Note `content_type` keeps the header's **parameters**, because for some media types the parameters are part of what the bytes are — ElevenLabs' `pcm_*` output @@ -506,6 +509,76 @@ contract marks that header required on every answer it sends, so an HTML interstitial from a proxy in front of it has none), and `not result.content` means nothing was delivered. Check them before writing `content` to disk. +### `models.list` — what you can run + +```python +from comfy_sdk import Comfy + +with Comfy(api_key="comfyui-...") as client: + for model in client.models.list(): + print(model.id, model.billing) +``` + +That walks Router's model catalog — `GET https://api.comfy.org/v2/models` — +page by page, following `next_cursor` while `has_more` is true, and yields one +`CatalogModel` per entry: `id` (the `{provider}/{model}` id `models.run` takes), +`provider`, `model`, and `billing` (per-model billing facts such as +`charges_on_policy_rejection`, never prices). Nothing is fetched until you +iterate, and each loop is a fresh walk. + +For one page and its paging facts instead, call `.page()`: + +```python +page = client.models.list(limit=50).page() +page.data # tuple of CatalogModel +page.has_more # walk on this, not on a short page +page.next_cursor # pass back as list(cursor=...) for the next page +page.limit # the page size the server actually served +page.request_id # X-Comfy-Request-Id, for a support request +``` + +`limit` is sent as given. The server defaults to 20 and clamps anything above +100 down to 100 rather than rejecting it, so read `page.limit` for the size you +got. The cursor is opaque and only good for the walk that produced it. On +`AsyncComfy` it is `async for model in client.models.list():` and +`await client.models.list().page()`. + +### `models.schema` — what a model takes, and what it returns + +```python +with Comfy(api_key="comfyui-...") as client: + result = client.models.schema("bfl/flux-2-pro") + document = result.document # the model's own OpenAPI document + etag = result.etag # keep it for the next call +``` + +That is `GET https://api.comfy.org/v2/models/bfl/flux-2-pro/openapi.json`: the +model's input schema (the body `models.run` sends) and its output schema, as a +standalone OpenAPI document. The id is validated exactly as `models.run` +validates it, before any request. + +Pass a tag you stored earlier to make the read conditional. When the document +has not changed, the server answers `304` with no body, and you get +`unchanged=True` back rather than an exception: + +```python +result = client.models.schema("bfl/flux-2-pro", etag=etag) +if result.unchanged: + ... # your cached document is still current; result.document is None +else: + document, etag = result.document, result.etag +``` + +The SDK keeps no cache, so storing the tag and the document is up to you. An +unknown model raises `ModelNotFound`, and `await client.models.schema(...)` is the +async form. + +Both methods use the same host and credential as `models.run`, raise the same +typed Router exceptions (see +[Catching Comfy Router errors](#catching-comfy-router-errors)), and default to +a 30-second timeout. Pass `timeout=` seconds, an `httpx.Timeout`, or `None` to +wait indefinitely. + ### Image to image — upload an asset first An image-to-image model takes an image *as input*, and Router forwards the @@ -558,8 +631,10 @@ result = client.models.run( ) ``` -Which form a model takes is in its input schema — `GET -/v2/models/{provider}/{model}/openapi.json`, or the model's page in the +Which form a model takes is in its input schema — +`client.models.schema("bfl/flux-2-pro").document` (see +[`models.schema`](#modelsschema--what-a-model-takes-and-what-it-returns)), or +the model's page in the [Router model catalog](https://docs.comfy.org/development/comfy-router/models). Because the server may legitimately hold the connection for minutes, `run` uses diff --git a/src/comfy_low/transport.py b/src/comfy_low/transport.py index 818d3dc..f2f1f23 100644 --- a/src/comfy_low/transport.py +++ b/src/comfy_low/transport.py @@ -38,6 +38,11 @@ for exactly the reason the run path was, and they are the one part of this change a spec sync is expected to correct. +``get_model_catalog`` / ``get_model_schema`` bind the two discovery reads of +the same contract (``listRouterModels``, ``getRouterModelInputSchema``), and +follow the run path's rule: their routes live in :data:`_MODEL_CATALOG_PATH` / +:data:`_MODEL_SCHEMA_PATH_TEMPLATE` alone, pinned against the vendored file. + This layer contains no orchestration, retries, hashing, or reconnection — those live in ``comfy_sdk``. """ @@ -125,6 +130,21 @@ _MODEL_REQUEST_STATUS_PATH_TEMPLATE = _MODEL_REQUEST_PATH_TEMPLATE + "/status" _MODEL_REQUEST_CANCEL_PATH_TEMPLATE = _MODEL_REQUEST_PATH_TEMPLATE + "/cancel" +#: Routes for model *discovery* — the catalog (``operationId: +#: listRouterModels``) and one model's input/output schema document +#: (``operationId: getRouterModelInputSchema``), verbatim from +#: ``spec/router-openapi.yaml``. Pinned against the vendored file by +#: ``tests/test_router_spec_contract.py`` exactly as +#: :data:`_MODEL_RUN_PATH_TEMPLATE` is, so a sync that moves either route fails. +_MODEL_CATALOG_PATH = "/v2/models" +_MODEL_SCHEMA_PATH_TEMPLATE = _MODEL_RUN_PATH_TEMPLATE + "/openapi.json" + +#: Default timeout for a discovery read. Both routes answer from a catalog the +#: server already holds, so nothing here waits on a generation — the bound is +#: the client's ordinary 30s, stated here so it does not silently follow a +#: client configured with a much longer (or shorter) timeout for other calls. +DISCOVERY_TIMEOUT = httpx.Timeout(30.0) + #: Longest request id accepted into a path. The contract mints UUIDs (36 #: characters); the bound exists so a server-controlled value that is NOT one #: cannot reach the public handle, a log line or an exception message unbounded. @@ -435,6 +455,80 @@ def model_request_path(model: str, request_id: str, template: str) -> str: ) +def model_catalog_path(cursor: str | None = None, limit: int | None = None) -> str: + """Sans-IO path (with query) for one page of the Router model catalog. + + Each parameter is sent only when given, so a bare call is the first page at + the server's default size. ``limit`` is passed through unchanged, including + a value above the declared maximum of 100: the route clamps rather than + rejects it and echoes the size it actually served, so refusing it here would + only disagree with the server. + """ + params: list[tuple[str, str]] = [] + if cursor is not None: + params.append(("cursor", cursor)) + if limit is not None: + params.append(("limit", str(limit))) + query = urlencode(params) + return f"{_MODEL_CATALOG_PATH}?{query}" if query else _MODEL_CATALOG_PATH + + +def model_schema_path(model: str) -> str: + """Sans-IO path for one model's input/output schema document. + + Addressed by the same ``{provider}/{model}`` id, validated and + percent-encoded exactly as :func:`model_run_request` does it, so an id that + runs is an id whose schema can be read. + """ + provider, name = parse_model_id(model) + return _MODEL_SCHEMA_PATH_TEMPLATE.format( + provider=quote(provider, safe=""), model=quote(name, safe="") + ) + + +def model_schema_headers(etag: str | None) -> dict[str, str] | None: + """``If-None-Match`` for a conditional schema read, or ``None`` for a plain one. + + Raises ``ValueError`` before any request for an empty or non-ASCII tag: an + empty header makes the read effectively unconditional, so a ``304`` to it + could not honestly mean "your copy is current", and a non-ASCII value would + fail inside httpx's header encoding as an untyped ``UnicodeEncodeError``. + """ + if etag is None: + return None + if not isinstance(etag, str): + raise TypeError(f"etag must be a str, got {type(etag).__name__}") + if not etag or not etag.isascii(): + raise ValueError(f"etag must be a non-empty ASCII string; got {etag!r}") + return {"If-None-Match": etag} + + +def model_schema_answer( + p: _Prepared, resp: httpx.Response, etag: str | None +) -> dict[str, Any] | None: + """The schema document, or ``None`` for a ``304`` to a conditional read. + + ``None`` is reserved for the ``304`` so the SDK can read it as "unchanged": + a ``200`` whose body decodes to anything but a JSON object (``null``, a list, + a scalar) is raised as ``invalid_response`` rather than passed through. + """ + # Only an answer to a conditional read: a 304 to a request that sent no tag + # cannot mean "your copy is current", so it falls through and is raised + # like any other unexpected status. + if resp.status_code == 304 and etag is not None: + return None + body: Any = p.parse_or_raise(resp, (200,)) + if not isinstance(body, dict): + raise ApiError( + f"The {resp.status_code} schema response is not a JSON object", + code="invalid_response", + http_status=resp.status_code, + request_id=_request_id(resp), + body_excerpt=_body_excerpt(resp), + ) + return body + + def _build_user_agent(client_info: str | None) -> str: """SDK identity sent on every request. This is request metadata (not telemetry — no phone-home), so adoption is measurable server-side from @@ -1326,6 +1420,47 @@ def put_model_request_cancel( resp = self.raw_request("PUT", url, timeout=timeout) return self._p.parse_or_raise(resp, (200, 202, 204)), resp.headers + # -- models: discovery ------------------------------------------------ + def get_model_catalog( + self, + *, + cursor: str | None = None, + limit: int | None = None, + timeout: Any = DISCOVERY_TIMEOUT, + ) -> tuple[dict[str, Any], httpx.Headers]: + """GET ``{router_base_url}/v2/models`` — one page of the model catalog. + + ``listRouterModels`` of ``spec/router-openapi.yaml``, hand-bound — see + :data:`_MODEL_CATALOG_PATH`. Returns ``(body, headers)`` so the caller + can read ``X-Comfy-Request-Id`` off a success. + """ + url = self._p.router_base_url + model_catalog_path(cursor, limit) + resp = self.raw_request("GET", url, timeout=timeout) + return self._p.parse_or_raise(resp, (200,)), resp.headers + + def get_model_schema( + self, + model: str, + *, + etag: str | None = None, + timeout: Any = DISCOVERY_TIMEOUT, + ) -> tuple[dict[str, Any] | None, httpx.Headers]: + """GET ``{router_base_url}/v2/models/{provider}/{model}/openapi.json``. + + ``getRouterModelInputSchema`` of ``spec/router-openapi.yaml``. With + ``etag`` the request carries ``If-None-Match``, and a ``304`` — the + document is unchanged — is a success, returned as a ``None`` body + rather than raised: it is the answer the caller asked for. + + Raises ``TypeError``/``ValueError`` before any request when ``model`` + is not a ``{provider}/{model}`` id (:func:`parse_model_id`) or ``etag`` + is not a non-empty ASCII string (:func:`model_schema_headers`). + """ + url = self._p.router_base_url + model_schema_path(model) + headers = model_schema_headers(etag) + resp = self.raw_request("GET", url, headers=headers, timeout=timeout) + return model_schema_answer(self._p, resp, etag), resp.headers + class AsyncComfyLow: """Asynchronous protocol bindings — mirrors :class:`ComfyLow`.""" @@ -1696,6 +1831,32 @@ async def put_model_request_cancel( resp = await self.raw_request("PUT", url, timeout=timeout) return self._p.parse_or_raise(resp, (200, 202, 204)), resp.headers + # -- models: discovery ------------------------------------------------ + async def get_model_catalog( + self, + *, + cursor: str | None = None, + limit: int | None = None, + timeout: Any = DISCOVERY_TIMEOUT, + ) -> tuple[dict[str, Any], httpx.Headers]: + """Async :meth:`ComfyLow.get_model_catalog`.""" + url = self._p.router_base_url + model_catalog_path(cursor, limit) + resp = await self.raw_request("GET", url, timeout=timeout) + return self._p.parse_or_raise(resp, (200,)), resp.headers + + async def get_model_schema( + self, + model: str, + *, + etag: str | None = None, + timeout: Any = DISCOVERY_TIMEOUT, + ) -> tuple[dict[str, Any] | None, httpx.Headers]: + """Async :meth:`ComfyLow.get_model_schema` — a ``304`` is a ``None`` body.""" + url = self._p.router_base_url + model_schema_path(model) + headers = model_schema_headers(etag) + resp = await self.raw_request("GET", url, headers=headers, timeout=timeout) + return model_schema_answer(self._p, resp, etag), resp.headers + def _looks_like_path(s: str) -> bool: return s.startswith("http") or s.startswith("/") diff --git a/src/comfy_sdk/__init__.py b/src/comfy_sdk/__init__.py index 39bb3e5..b38380f 100644 --- a/src/comfy_sdk/__init__.py +++ b/src/comfy_sdk/__init__.py @@ -83,6 +83,7 @@ except PackageNotFoundError: # running from a source tree, not installed __version__ = "0+unknown" +from .model_catalog import AsyncModelList, CatalogModel, ModelList, ModelPage, SchemaResult from .models import RouterRunResult __all__ = [ @@ -96,6 +97,12 @@ "AsyncComfy", # model runs "RouterRunResult", + # model discovery + "CatalogModel", + "ModelList", + "AsyncModelList", + "ModelPage", + "SchemaResult", # assets / workflows / jobs / outputs "Asset", "AsyncAsset", diff --git a/src/comfy_sdk/model_catalog.py b/src/comfy_sdk/model_catalog.py new file mode 100644 index 0000000..51f30ec --- /dev/null +++ b/src/comfy_sdk/model_catalog.py @@ -0,0 +1,388 @@ +"""Model discovery on Comfy Router — ``models.list()`` and ``models.schema()``. + +The two read-only routes that tell a caller *what* it can run before it runs +anything: the catalog (``GET /v2/models``, ``listRouterModels``) and one model's +input/output schema as a standalone OpenAPI document +(``GET /v2/models/{provider}/{model}/openapi.json``, +``getRouterModelInputSchema``). Both go to the same host, with the same +credential, as :meth:`comfy_sdk.models.Models.run`, and a failure raises the +same typed :class:`~comfy_sdk.router_exceptions.RouterError` subclass, read off +``X-Comfy-Error-Type``. + +The shapes mirror the TypeScript SDK's ``comfy.models.list()`` / +``comfy.models.schema()`` so the two surfaces answer the same question the same +way: + +* ``list()`` returns a :class:`ModelList` — iterate it to walk every page, + following ``next_cursor`` while ``has_more`` is true, or call + :meth:`ModelList.page` for exactly one page and its paging facts; +* ``schema()`` returns a :class:`SchemaResult`. With ``etag=`` the request is + conditional, and a ``304`` comes back as ``unchanged=True`` rather than as an + exception, because "nothing changed" is the answer the caller asked for. + Storing the tag between calls is the caller's business; the SDK keeps no + cache. +""" + +from __future__ import annotations + +import asyncio +import time +from collections.abc import AsyncIterator, Awaitable, Callable, Iterator, Mapping +from dataclasses import dataclass, replace +from typing import Any, TypeVar, cast + +import httpx + +from comfy_low.errors import ApiError, clean_request_id +from comfy_low.transport import DISCOVERY_TIMEOUT, AsyncComfyLow, ComfyLow, parse_model_id + +from .exceptions import ComfyError, translating +from .retry import Retrier, RetryPolicy +from .router_exceptions import REQUEST_ID_HEADER, RouterError + +#: See :data:`comfy_sdk.models._CANDIDATE_FAILURES` — restated rather than +#: imported because :mod:`comfy_sdk.models` imports this module. +_CANDIDATE_FAILURES = (ApiError, RouterError, httpx.TransportError) + +_now = time.monotonic + +_T = TypeVar("_T") + +Timeout = float | httpx.Timeout | None + + +@dataclass(frozen=True, slots=True) +class CatalogModel: + """One entry in the Router model catalog — the identity of a runnable model.""" + + id: str + """The canonical ``{provider}/{model}`` id — pass it straight to ``models.run``.""" + + provider: str + """The ``provider`` segment of :attr:`id`.""" + + model: str + """The ``model`` segment of :attr:`id`.""" + + billing: Mapping[str, Any] + """Per-model billing facts a caller needs before invoking — never prices. + + Carried as the decoded JSON object (today it holds + ``charges_on_policy_rejection``) rather than as a class, so a field the + server adds reaches the caller without an SDK release. + """ + + def __hash__(self) -> int: + # The generated hash would cover `billing`, a dict, and raise; the id + # triple is enough to be consistent with the generated `__eq__`, and + # lets `set(client.models.list())` work. + return hash((self.id, self.provider, self.model)) + + +@dataclass(frozen=True, slots=True) +class ModelPage: + """One page of the Router model catalog, as :meth:`ModelList.page` returns it.""" + + data: tuple[CatalogModel, ...] + """The models on this page, at most :attr:`limit` of them.""" + + has_more: bool + """Whether another page exists. Walk on this, never on a short :attr:`data`.""" + + next_cursor: str | None + """Pass as ``cursor=`` to fetch the next page; ``None`` on the last one.""" + + limit: int + """The page size the server actually served. + + A requested ``limit`` above the maximum (100) is clamped down rather than + rejected, so this can be smaller than the value asked for. + """ + + request_id: str | None + """``X-Comfy-Request-Id`` — the id to quote in a support request.""" + + +@dataclass(frozen=True, slots=True) +class SchemaResult: + """What :meth:`comfy_sdk.models.Models.schema` returns. + + ``unchanged`` is the branch to take: ``False`` carries the full + :attr:`document`; ``True`` means the ``etag`` passed in still matches, the + server sent no body, and :attr:`document` is ``None``. + """ + + unchanged: bool + """``True`` on a ``304`` — the document is the one the caller's ``etag`` names.""" + + document: dict[str, Any] | None + """The input/output schemas as a standalone OpenAPI document; ``None`` on a ``304``.""" + + etag: str | None + """``ETag`` of the current document — send it back as ``etag=`` next time.""" + + request_id: str | None + """``X-Comfy-Request-Id`` — the id to quote in a support request.""" + + +def _invalid(message: str, request_id: str | None) -> ComfyError: + return ComfyError(message, code="invalid_response", request_id=request_id) + + +def _entry(raw: Any, request_id: str | None) -> CatalogModel: + """One catalog entry, or ``invalid_response`` if it cannot address a run. + + ``id``/``provider``/``model`` are required and checked, because they are + what a caller hands back to ``models.run``: ``id`` must parse exactly as + ``models.run`` would parse it, and ``provider``/``model`` must be its two + segments, so a caller that allowlists on ``entry.provider`` and then runs + ``entry.id`` runs the provider it checked. ``billing`` is read leniently, + because nothing downstream is addressed by it. + """ + if not isinstance(raw, dict): + raise _invalid("catalog entry is not a JSON object", request_id) + model_id, provider, model = raw.get("id"), raw.get("provider"), raw.get("model") + if not all(isinstance(value, str) and value for value in (model_id, provider, model)): + raise _invalid("catalog entry is missing id, provider or model", request_id) + try: + segments = parse_model_id(cast(str, model_id)) + except ValueError: + raise _invalid("catalog entry id is not a runnable model id", request_id) from None + if segments != (provider, model): + raise _invalid("catalog entry id does not match its provider and model", request_id) + billing = raw.get("billing") + return CatalogModel( + id=cast(str, model_id), + provider=cast(str, provider), + model=cast(str, model), + billing=dict(billing) if isinstance(billing, dict) else {}, + ) + + +def _page(body: Any, headers: Mapping[str, str]) -> ModelPage: + request_id = clean_request_id(headers.get(REQUEST_ID_HEADER)) + # The transport returns the decoded body unchecked, so a 200 carrying a + # top-level list, string or `null` has to be refused here, inside the same + # `invalid_response` as every other malformed page. + if not isinstance(body, Mapping): + raise _invalid("catalog page is not a JSON object", request_id) + data = body.get("data") + has_more = body.get("has_more") + limit = body.get("limit") + if not isinstance(data, list) or not isinstance(has_more, bool): + raise _invalid("catalog page is missing data or has_more", request_id) + next_cursor = body.get("next_cursor") + return ModelPage( + data=tuple(_entry(item, request_id) for item in data), + has_more=has_more, + # Only a page that has more names a next one: a stale cursor on the + # last page would send a caller paging by hand round again. + next_cursor=( + next_cursor if has_more and isinstance(next_cursor, str) and next_cursor else None + ), + # `limit` is required by the contract, but a page that omits it is + # still a usable page; its own length is the honest stand-in. + limit=limit if isinstance(limit, int) and not isinstance(limit, bool) else len(data), + request_id=request_id, + ) + + +def _next_cursor(page: ModelPage, seen: set[str]) -> str | None: + """The cursor to walk to next, or ``None`` once the catalog is exhausted. + + A page that says ``has_more`` yet names no cursor, or names one this walk + already followed, would otherwise loop forever re-reading the same page; + both are raised as ``invalid_response`` instead. + """ + if not page.has_more: + return None + cursor = page.next_cursor + if cursor is None: + raise _invalid("catalog page says has_more but names no next_cursor", page.request_id) + if cursor in seen: + raise _invalid("catalog walk was handed a cursor it already followed", page.request_id) + seen.add(cursor) + return cursor + + +def _schema_result( + body: dict[str, Any] | None, headers: Mapping[str, str], sent: str | None +) -> SchemaResult: + # The transport returns `None` for a 304 to a conditional read and nothing + # else — a non-object 200 is already `invalid_response` there. + unchanged = body is None + etag = headers.get("ETag") + if unchanged and etag is None: + # A 304 just confirmed the tag that was sent; an intermediary that + # strips `ETag` off it must not make the caller drop that tag and read + # unconditionally from then on. + etag = sent + return SchemaResult( + unchanged=unchanged, + document=body, + etag=etag, + request_id=clean_request_id(headers.get(REQUEST_ID_HEADER)), + ) + + +def _read_policy(policy: RetryPolicy) -> RetryPolicy: + """``policy`` for a discovery read, which is a plain ``GET``. + + :attr:`~comfy_sdk.retry.RetryPolicy.retry_possibly_in_flight` is off by + default because a run's resend may bill a second generation or be refused + for its reused key. Neither applies to a read that carries no key and bills + nothing, so a ``503`` or a read timeout here is retried whenever the policy + retries at all; its budgets and backoff are still the caller's. + """ + return replace(policy, retry_possibly_in_flight=True) + + +def _call(policy: RetryPolicy, send: Callable[[], _T]) -> _T: + """Run one discovery read under the client's retry policy, translated.""" + retrier = Retrier(_read_policy(policy), now=_now) + with translating(): + while True: + try: + return send() + except _CANDIDATE_FAILURES as exc: + delay = retrier.delay_before_retry(exc) + if delay is None: + raise + time.sleep(delay) + + +async def _acall(policy: RetryPolicy, send: Callable[[], Awaitable[_T]]) -> _T: + """Awaitable :func:`_call`.""" + retrier = Retrier(_read_policy(policy), now=_now) + with translating(): + while True: + try: + return await send() + except _CANDIDATE_FAILURES as exc: + delay = retrier.delay_before_retry(exc) + if delay is None: + raise + await asyncio.sleep(delay) + + +class ModelList: + """The Router model catalog, returned by :meth:`comfy_sdk.models.Models.list`. + + Iterating it walks every page from ``cursor`` (the first page when + omitted), yielding :class:`CatalogModel` entries and following + ``next_cursor`` while ``has_more`` is true. Nothing is fetched until + iteration starts, and each ``for`` loop starts a fresh walk. Call + :meth:`page` instead for exactly one page and its paging facts. + """ + + def __init__( + self, + low: ComfyLow, + retry: RetryPolicy, + *, + cursor: str | None, + limit: int | None, + timeout: Timeout, + ) -> None: + self._low = low + self._retry = retry + self._cursor = cursor + self._limit = limit + self._timeout = timeout + + def _fetch(self, cursor: str | None) -> ModelPage: + body, headers = _call( + self._retry, + lambda: self._low.get_model_catalog( + cursor=cursor, limit=self._limit, timeout=self._timeout + ), + ) + return _page(body, headers) + + def page(self) -> ModelPage: + """Fetch one page — the one ``cursor`` names — and return it whole.""" + return self._fetch(self._cursor) + + def __iter__(self) -> Iterator[CatalogModel]: + cursor = self._cursor + seen: set[str] = set() if cursor is None else {cursor} + while True: + page = self._fetch(cursor) + yield from page.data + cursor = _next_cursor(page, seen) + if cursor is None: + return + + +class AsyncModelList: + """The Router model catalog on ``AsyncComfy`` — mirrors :class:`ModelList`. + + ``async for`` walks every page; ``await catalog.page()`` fetches one. + """ + + def __init__( + self, + low: AsyncComfyLow, + retry: RetryPolicy, + *, + cursor: str | None, + limit: int | None, + timeout: Timeout, + ) -> None: + self._low = low + self._retry = retry + self._cursor = cursor + self._limit = limit + self._timeout = timeout + + async def _fetch(self, cursor: str | None) -> ModelPage: + body, headers = await _acall( + self._retry, + lambda: self._low.get_model_catalog( + cursor=cursor, limit=self._limit, timeout=self._timeout + ), + ) + return _page(body, headers) + + async def page(self) -> ModelPage: + """Awaitable :meth:`ModelList.page`.""" + return await self._fetch(self._cursor) + + async def __aiter__(self) -> AsyncIterator[CatalogModel]: + cursor = self._cursor + seen: set[str] = set() if cursor is None else {cursor} + while True: + page = await self._fetch(cursor) + for entry in page.data: + yield entry + cursor = _next_cursor(page, seen) + if cursor is None: + return + + +def get_schema( + low: ComfyLow, retry: RetryPolicy, model: str, *, etag: str | None, timeout: Timeout +) -> SchemaResult: + """One ``openapi.json`` read, sync — see :meth:`comfy_sdk.models.Models.schema`.""" + body, headers = _call(retry, lambda: low.get_model_schema(model, etag=etag, timeout=timeout)) + return _schema_result(body, headers, etag) + + +async def aget_schema( + low: AsyncComfyLow, retry: RetryPolicy, model: str, *, etag: str | None, timeout: Timeout +) -> SchemaResult: + """Awaitable :func:`get_schema`.""" + body, headers = await _acall( + retry, lambda: low.get_model_schema(model, etag=etag, timeout=timeout) + ) + return _schema_result(body, headers, etag) + + +__all__ = [ + "DISCOVERY_TIMEOUT", + "AsyncModelList", + "CatalogModel", + "ModelList", + "ModelPage", + "SchemaResult", +] diff --git a/src/comfy_sdk/models.py b/src/comfy_sdk/models.py index c1afe42..149a786 100644 --- a/src/comfy_sdk/models.py +++ b/src/comfy_sdk/models.py @@ -68,6 +68,14 @@ from ._core import new_idempotency_key, validate_idempotency_key from .exceptions import IdempotencyKeyReuse, _stamp, to_sdk_error, translating +from .model_catalog import ( + DISCOVERY_TIMEOUT, + AsyncModelList, + ModelList, + SchemaResult, + aget_schema, + get_schema, +) from .model_requests import ( _CANCEL_FAILURES, _CANCEL_TIMEOUT, @@ -883,6 +891,66 @@ def handle(self, model: str, request_id: str) -> RequestHandle: parse_request_id(request_id) return RequestHandle(cast(ComfyLow, self._low), model, request_id, self._retry) + # -- discovery: what can run, and what it takes ------------------------ + def list( + self, + *, + cursor: str | None = None, + limit: int | None = None, + timeout: float | httpx.Timeout | None = DISCOVERY_TIMEOUT, + ) -> ModelList: + """The Router model catalog — every model this client can :meth:`run`. + + ``GET {router_base_url}/v2/models``, with the same host and credential + as :meth:`run`. Iterate the returned + :class:`~comfy_sdk.model_catalog.ModelList` to walk the whole catalog — + it follows ``next_cursor`` page by page while ``has_more`` is true — or + call its ``page()`` for exactly one page plus ``has_more``, + ``next_cursor``, the ``limit`` actually served, and ``request_id``. + Nothing is fetched until you do one or the other. + + ``cursor`` starts from a page an earlier walk handed you (it is opaque + and only valid for that walk). ``limit`` is the page size, sent as + given: the server defaults to 20 and clamps anything above 100 down to + 100 rather than rejecting it, and reports the size it served on the + page. Each is sent only when set. + + ``timeout`` bounds each page request, defaulting to 30 seconds + (:data:`~comfy_low.transport.DISCOVERY_TIMEOUT`); pass seconds, an + ``httpx.Timeout``, or ``None`` to wait indefinitely. A failure raises + the same typed :class:`~comfy_sdk.router_exceptions.RouterError` + subclass :meth:`run` would, read off ``X-Comfy-Error-Type``. + """ + return ModelList( + cast(ComfyLow, self._low), self._retry, cursor=cursor, limit=limit, timeout=timeout + ) + + def schema( + self, + model: str, + *, + etag: str | None = None, + timeout: float | httpx.Timeout | None = DISCOVERY_TIMEOUT, + ) -> SchemaResult: + """``model``'s input and output schemas, as a standalone OpenAPI document. + + ``GET {router_base_url}/v2/models/{provider}/{model}/openapi.json``. + ``model`` is the same canonical ``{provider}/{model}`` id :meth:`run` + takes, validated the same way before any request. + + Returns a :class:`~comfy_sdk.model_catalog.SchemaResult` carrying the + parsed ``document``, its ``etag`` and the ``request_id``. Pass a stored + ``etag`` to make the read conditional (``If-None-Match``): when the + document has not changed the server answers ``304`` with no body, and + this returns ``unchanged=True`` with ``document=None`` — not an + exception. The SDK keeps no cache; storing the tag is the caller's. + + ``timeout`` defaults to 30 seconds, as on :meth:`list`. An unknown + model raises :class:`~comfy_sdk.router_exceptions.ModelNotFound`, and + every other failure the typed Router exception :meth:`run` would. + """ + return get_schema(cast(ComfyLow, self._low), self._retry, model, etag=etag, timeout=timeout) + class AsyncModels(_ModelsBase): """``client.models`` on :class:`~comfy_sdk.client.AsyncComfy` — mirrors :class:`Models`.""" @@ -1108,5 +1176,37 @@ async def handle(self, model: str, request_id: str) -> AsyncRequestHandle: parse_request_id(request_id) return AsyncRequestHandle(cast(AsyncComfyLow, self._low), model, request_id, self._retry) + # -- discovery: what can run, and what it takes ------------------------ + def list( + self, + *, + cursor: str | None = None, + limit: int | None = None, + timeout: float | httpx.Timeout | None = DISCOVERY_TIMEOUT, + ) -> AsyncModelList: + """The Router model catalog on ``AsyncComfy`` — see :meth:`Models.list`. + + Not itself awaited: it returns an + :class:`~comfy_sdk.model_catalog.AsyncModelList` at once, which + ``async for`` walks page by page and whose ``page()`` is awaited for a + single page. That keeps ``async for m in client.models.list():`` the + one-liner it is in the sync client. + """ + return AsyncModelList( + cast(AsyncComfyLow, self._low), self._retry, cursor=cursor, limit=limit, timeout=timeout + ) + + async def schema( + self, + model: str, + *, + etag: str | None = None, + timeout: float | httpx.Timeout | None = DISCOVERY_TIMEOUT, + ) -> SchemaResult: + """Awaitable :meth:`Models.schema` — same arguments, same result.""" + return await aget_schema( + cast(AsyncComfyLow, self._low), self._retry, model, etag=etag, timeout=timeout + ) + __all__ = ["Models", "AsyncModels", "BinaryResult", "RouterRunResult"] diff --git a/tests/conftest.py b/tests/conftest.py index c790148..208b37d 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -24,7 +24,7 @@ from dataclasses import dataclass, field from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Any -from urllib.parse import unquote +from urllib.parse import parse_qs, unquote, urlsplit import pytest @@ -202,6 +202,49 @@ class ServerState: # The bucket sent on `X-Comfy-Error-Type` alongside it. model_run_validation_error_type: str = "invalid_input" + # --- model discovery: GET /v2/models and .../openapi.json --- + # Catalog pages keyed by the `cursor` that fetches them (`None` is the first + # page). Each value is the whole response body, so a test states the + # paging facts (`has_more`, `next_cursor`, `limit`) it is about. + catalog_pages: dict[str | None, dict[str, Any]] = field( + default_factory=lambda: { + None: { + "data": [ + { + "id": "bfl/flux-2-pro", + "provider": "bfl", + "model": "flux-2-pro", + "billing": {"charges_on_policy_rejection": "unknown"}, + } + ], + "has_more": False, + "next_cursor": None, + "limit": 20, + } + } + ) + # (status, code) answered instead of a catalog page. + catalog_error: tuple[int, str] | None = None + # Catalog requests answered 429 rate_limited (Retry-After: 1) before one is + # served — the transient path the client's retry policy rides out. + catalog_fail_times: int = 0 + # Every catalog request's decoded query, in arrival order. + catalog_queries: list[dict[str, list[str]]] = field(default_factory=list) + # The document `.../openapi.json` serves, and the ETag it is served under. + schema_document: dict[str, Any] = field( + default_factory=lambda: {"openapi": "3.1.0", "info": {"title": "bfl/flux-2-pro"}} + ) + schema_etag: str = '"schema-v1"' + # Models the schema route knows; any other id answers 404 model_not_found. + schema_models: set[str] = field(default_factory=lambda: {"bfl/flux-2-pro"}) + # (status, code) answered instead of the document, for every model. + schema_error: tuple[int, str] | None = None + # Raw request path and If-None-Match of every schema request. + schema_paths: list[str] = field(default_factory=list) + schema_if_none_match: list[str | None] = field(default_factory=list) + # The `X-Comfy-Request-Id` both discovery routes stamp on their answers. + discovery_request_id: str = "req_discovery_01" + # --- the queued model surface (submit / status / result / cancel) --- # POST .../requests answers this status with a body naming a request id. queue_submit_status: int = 200 @@ -557,6 +600,15 @@ def do_GET(self) -> None: if m: self._serve_queue_result(m.group(3)) return + # Comfy Router's discovery routes — the catalog (query-string + # paged) and one model's schema document. + if urlsplit(self.path).path == "/v2/models": + self._serve_catalog() + return + m = re.match(r"/v2/models/([^/]+)/([^/]+)/openapi\.json$", self.path) + if m: + self._serve_schema(unquote(m.group(1)), unquote(m.group(2))) + return m = re.match(r"/api/v2/jobs/([^/]+)/events$", self.path) if m: self._serve_events(m.group(1)) @@ -845,6 +897,50 @@ def _put_queue_cancel(self, request_id: str) -> None: }, ) + def _serve_catalog(self) -> None: + query = parse_qs(urlsplit(self.path).query) + state.catalog_queries.append(query) + if state.catalog_fail_times > 0: + state.catalog_fail_times -= 1 + self._router_err(429, "rate_limited", retry_after="1") + return + if state.catalog_error: + status, code = state.catalog_error + self._router_err(status, code) + return + cursor = query.get("cursor", [None])[0] + if cursor not in state.catalog_pages: + self._router_err(400, "invalid_input", "unknown cursor") + return + self._json( + 200, + state.catalog_pages[cursor], + headers={"X-Comfy-Request-Id": state.discovery_request_id}, + ) + + def _serve_schema(self, provider: str, model: str) -> None: + state.schema_paths.append(self.path) + state.schema_if_none_match.append(self.headers.get("If-None-Match")) + if state.schema_error: + status, code = state.schema_error + self._router_err(status, code) + return + if f"{provider}/{model}" not in state.schema_models: + self._router_err(404, "model_not_found", "no such model") + return + headers = { + "X-Comfy-Request-Id": state.discovery_request_id, + "ETag": state.schema_etag, + "Cache-Control": "public, max-age=300", + } + if self.headers.get("If-None-Match") == state.schema_etag: + self.send_response(304) + for k, v in headers.items(): + self.send_header(k, v) + self.end_headers() + return + self._json(200, state.schema_document, headers=headers) + def _router_err( self, status: int, code: str, message: str = "err", retry_after: str | None = None ) -> None: diff --git a/tests/test_models_discovery.py b/tests/test_models_discovery.py new file mode 100644 index 0000000..c83e113 --- /dev/null +++ b/tests/test_models_discovery.py @@ -0,0 +1,498 @@ +"""Model discovery: ``models.list()`` and ``models.schema()``. + +Driven against the stub in ``conftest.py`` — ``server.state.catalog_pages`` and +the ``schema_*`` fields set the scenario, and the fixture points the SDK at the +stub through ``COMFY_ROUTER_BASE_URL``, exactly as the ``models.run`` tests do. + +What each group here is for: + +* the **catalog walk** — every page, followed by ``next_cursor`` while + ``has_more`` is true, and the single-page form beside it; +* ``limit`` passed through **unchanged** above the maximum (the server clamps); +* the **conditional schema read** — ``If-None-Match`` sent, and a ``304`` + returned as ``unchanged`` rather than raised; +* the **same host, credential and typed errors** as ``models.run``; +* and every one of those on the async client too. +""" + +from __future__ import annotations + +import inspect +from typing import Any + +import httpx +import pytest + +from comfy_low.transport import DISCOVERY_TIMEOUT +from comfy_sdk import ( + ROUTER_BASE_URL_ENV_VAR, + AsyncComfy, + AsyncModelList, + CatalogModel, + Comfy, + ModelList, + ModelPage, + SchemaResult, +) +from comfy_sdk.exceptions import ComfyError +from comfy_sdk.models import AsyncModels, Models +from comfy_sdk.retry import NO_RETRY +from comfy_sdk.router_exceptions import ( + Forbidden, + InternalError, + ModelNotFound, + RouterError, + ServiceUnavailable, + Unauthorized, +) + +MODEL = "bfl/flux-2-pro" + + +def _entry(provider: str, model: str) -> dict[str, Any]: + return { + "id": f"{provider}/{model}", + "provider": provider, + "model": model, + "billing": {"charges_on_policy_rejection": "no"}, + } + + +def _three_pages(server) -> None: + server.state.catalog_pages = { + None: { + "data": [_entry("bfl", "flux-2-pro"), _entry("bfl", "flux-2-dev")], + "has_more": True, + "next_cursor": "c-2", + "limit": 2, + }, + "c-2": { + "data": [_entry("wan", "wan2.5-i2i-preview")], + "has_more": True, + "next_cursor": "c-3", + "limit": 2, + }, + # Empty but not last would be legal; the walk ends on `has_more`, not + # on a short page — so the last page here is deliberately short *and* + # carries a stale cursor the walk must not follow. + "c-3": { + "data": [_entry("acme", "fast-sdxl")], + "has_more": False, + "next_cursor": "c-stale", + "limit": 2, + }, + } + + +EXPECTED_IDS = ["bfl/flux-2-pro", "bfl/flux-2-dev", "wan/wan2.5-i2i-preview", "acme/fast-sdxl"] + + +# --- list(): the walk ------------------------------------------------------ + + +def test_list_walks_every_page_following_next_cursor(server) -> None: + _three_pages(server) + with Comfy() as client: + models = list(client.models.list(limit=2)) + assert [m.id for m in models] == EXPECTED_IDS + cursors = [q.get("cursor", [None])[0] for q in server.state.catalog_queries] + assert cursors == [None, "c-2", "c-3"] + assert all(q["limit"] == ["2"] for q in server.state.catalog_queries) + + +async def test_async_list_walks_every_page_following_next_cursor(server) -> None: + _three_pages(server) + async with AsyncComfy() as client: + models = [m async for m in client.models.list(limit=2)] + assert [m.id for m in models] == EXPECTED_IDS + cursors = [q.get("cursor", [None])[0] for q in server.state.catalog_queries] + assert cursors == [None, "c-2", "c-3"] + + +def test_list_entries_expose_id_provider_model_and_billing(server) -> None: + with Comfy() as client: + (entry,) = list(client.models.list()) + assert entry == CatalogModel( + id="bfl/flux-2-pro", + provider="bfl", + model="flux-2-pro", + billing={"charges_on_policy_rejection": "unknown"}, + ) + + +def test_list_is_lazy_and_each_iteration_is_a_fresh_walk(server) -> None: + with Comfy() as client: + catalog = client.models.list() + assert isinstance(catalog, ModelList) + assert server.state.catalog_queries == [] + assert len(list(catalog)) == 1 + assert len(list(catalog)) == 1 + assert len(server.state.catalog_queries) == 2 + + +def test_list_starts_from_the_cursor_given(server) -> None: + _three_pages(server) + with Comfy() as client: + models = list(client.models.list(cursor="c-2")) + assert [m.id for m in models] == EXPECTED_IDS[2:] + + +def test_list_sends_no_query_when_nothing_is_set(server) -> None: + with Comfy() as client: + list(client.models.list()) + assert server.state.catalog_queries == [{}] + + +@pytest.mark.parametrize("limit", [101, 500]) +def test_limit_above_the_maximum_is_passed_through_unchanged(server, limit: int) -> None: + # The route clamps rather than rejects, and reports the size it served — + # so the SDK sends what it was given and reads the served size back. + server.state.catalog_pages[None]["limit"] = 100 + with Comfy() as client: + page = client.models.list(limit=limit).page() + assert server.state.catalog_queries == [{"limit": [str(limit)]}] + assert page.limit == 100 + + +def test_a_has_more_page_without_a_cursor_raises_rather_than_looping(server) -> None: + server.state.catalog_pages[None].update(has_more=True, next_cursor=None) + with Comfy() as client, pytest.raises(ComfyError) as info: + list(client.models.list()) + assert info.value.code == "invalid_response" + assert len(server.state.catalog_queries) == 1 + + +def test_a_cursor_cycle_raises_rather_than_looping(server) -> None: + _three_pages(server) + server.state.catalog_pages["c-3"].update(has_more=True, next_cursor="c-2") + with Comfy() as client, pytest.raises(ComfyError) as info: + list(client.models.list()) + assert info.value.code == "invalid_response" + assert len(server.state.catalog_queries) == 3 + + +# --- list().page(): the single-page form ----------------------------------- + + +def test_page_returns_one_page_with_its_paging_facts(server) -> None: + _three_pages(server) + with Comfy() as client: + page = client.models.list(limit=2).page() + assert isinstance(page, ModelPage) + assert [m.id for m in page.data] == EXPECTED_IDS[:2] + assert page.has_more is True + assert page.next_cursor == "c-2" + assert page.limit == 2 + assert page.request_id == "req_discovery_01" + assert len(server.state.catalog_queries) == 1 + + +async def test_async_page_returns_one_page_with_its_paging_facts(server) -> None: + _three_pages(server) + async with AsyncComfy() as client: + catalog = client.models.list(cursor="c-3") + assert isinstance(catalog, AsyncModelList) + page = await catalog.page() + assert [m.id for m in page.data] == EXPECTED_IDS[3:] + assert page.has_more is False + assert page.request_id == "req_discovery_01" + + +# --- schema() ---------------------------------------------------------------- + + +def test_schema_returns_the_document_etag_and_request_id(server) -> None: + with Comfy() as client: + result = client.models.schema(MODEL) + assert result == SchemaResult( + unchanged=False, + document=server.state.schema_document, + etag='"schema-v1"', + request_id="req_discovery_01", + ) + assert server.state.schema_paths == ["/v2/models/bfl/flux-2-pro/openapi.json"] + # No tag given, so the read is unconditional. + assert server.state.schema_if_none_match == [None] + + +async def test_async_schema_returns_the_document(server) -> None: + async with AsyncComfy() as client: + result = await client.models.schema(MODEL) + assert result.unchanged is False + assert result.document == server.state.schema_document + assert result.etag == '"schema-v1"' + + +def test_schema_sends_if_none_match_and_a_304_is_unchanged_not_an_error(server) -> None: + with Comfy() as client: + result = client.models.schema(MODEL, etag='"schema-v1"') + assert server.state.schema_if_none_match == ['"schema-v1"'] + assert result == SchemaResult( + unchanged=True, document=None, etag='"schema-v1"', request_id="req_discovery_01" + ) + + +async def test_async_schema_304_is_unchanged(server) -> None: + async with AsyncComfy() as client: + result = await client.models.schema(MODEL, etag='"schema-v1"') + assert result.unchanged is True + assert result.document is None + assert result.etag == '"schema-v1"' + assert result.request_id == "req_discovery_01" + + +def test_a_stale_etag_gets_the_new_document(server) -> None: + server.state.schema_etag = '"schema-v2"' + with Comfy() as client: + result = client.models.schema(MODEL, etag='"schema-v1"') + assert result.unchanged is False + assert result.document == server.state.schema_document + assert result.etag == '"schema-v2"' + + +def test_schema_percent_encodes_each_segment(server) -> None: + server.state.schema_models = {"acme/fast sdxl"} + with Comfy() as client: + client.models.schema("acme/fast sdxl") + assert server.state.schema_paths == ["/v2/models/acme/fast%20sdxl/openapi.json"] + + +@pytest.mark.parametrize("bad", ["bfl", "bfl/flux-2-pro/fp8", "bfl/..", "a//b"]) +def test_schema_rejects_a_malformed_id_before_any_request(server, bad: str) -> None: + with Comfy() as client, pytest.raises(ValueError): + client.models.schema(bad) + assert server.state.schema_paths == [] + + +def test_schema_404_for_an_unknown_model_raises_model_not_found(server) -> None: + with Comfy(retry=NO_RETRY) as client, pytest.raises(ModelNotFound) as info: + client.models.schema("bfl/no-such-model") + assert info.value.http_status == 404 + assert info.value.error_type == "model_not_found" + + +async def test_async_schema_404_for_an_unknown_model_raises_model_not_found(server) -> None: + async with AsyncComfy(retry=NO_RETRY) as client: + with pytest.raises(ModelNotFound): + await client.models.schema("bfl/no-such-model") + + +# --- same host, credential, timeout and typed errors as run() ------------- + + +@pytest.mark.parametrize( + ("status", "code", "cls"), + [ + (401, "unauthorized", Unauthorized), + (403, "forbidden", Forbidden), + (500, "internal_error", InternalError), + (503, "service_unavailable", ServiceUnavailable), + ], +) +def test_discovery_failures_raise_the_typed_router_exception( + server, status: int, code: str, cls: type[RouterError] +) -> None: + server.state.catalog_error = (status, code) + server.state.schema_error = (status, code) + with Comfy(retry=NO_RETRY) as client: + with pytest.raises(cls) as listed: + list(client.models.list()) + with pytest.raises(cls) as schema: + client.models.schema(MODEL) + assert listed.value.http_status == status + assert schema.value.http_status == status + + +async def test_async_list_failure_raises_the_typed_router_exception(server) -> None: + server.state.catalog_error = (403, "forbidden") + async with AsyncComfy(retry=NO_RETRY) as client: + with pytest.raises(Forbidden): + await client.models.list().page() + + +def test_discovery_goes_to_the_router_with_the_clients_credential( + monkeypatch, server, second_server +) -> None: + monkeypatch.setenv(ROUTER_BASE_URL_ENV_VAR, second_server.base_url) + second_server.state.require_auth = True + with Comfy(api_key="k-discover") as client: + list(client.models.list()) + assert second_server.state.last_auth_header == "Bearer k-discover" + client.models.schema(MODEL) + assert second_server.state.last_auth_header == "Bearer k-discover" + assert len(second_server.state.catalog_queries) == 1 + assert len(second_server.state.schema_paths) == 1 + # ...and nothing went to the v2 deployment. + assert server.state.catalog_queries == [] + assert server.state.schema_paths == [] + + +def test_discovery_retries_a_transient_failure_under_the_client_policy(server, monkeypatch) -> None: + # A 429 naming Retry-After is retried by the default policy, exactly as on + # every other Router route this namespace owns. + import comfy_sdk.model_catalog as catalog + + monkeypatch.setattr(catalog.time, "sleep", lambda _s: None) + server.state.catalog_fail_times = 1 + with Comfy() as client: + assert len(list(client.models.list())) == 1 + assert len(server.state.catalog_queries) == 2 + + +@pytest.mark.parametrize("cls", [Models, AsyncModels]) +@pytest.mark.parametrize("method", ["list", "schema"]) +def test_discovery_defaults_to_a_30_second_timeout(cls: type, method: str) -> None: + default = inspect.signature(getattr(cls, method)).parameters["timeout"].default + assert default is DISCOVERY_TIMEOUT + assert isinstance(default, httpx.Timeout) + assert (default.connect, default.read, default.write, default.pool) == (30.0,) * 4 + + +def test_a_timeout_reaches_the_request(server, monkeypatch) -> None: + seen: list[Any] = [] + with Comfy() as client: + real = client._low.raw_request + + def spy(*args: Any, **kwargs: Any) -> httpx.Response: + seen.append(kwargs.get("timeout")) + return real(*args, **kwargs) + + monkeypatch.setattr(client._low, "raw_request", spy) + list(client.models.list()) + client.models.schema(MODEL, timeout=5.0) + assert seen == [DISCOVERY_TIMEOUT, 5.0] + + +def test_a_304_to_an_unconditional_read_is_not_reported_as_unchanged(server) -> None: + # Only a request that sent If-None-Match can be told "your copy is + # current"; a caller with no tag has no copy, so a stray 304 must raise. + with Comfy(retry=NO_RETRY) as client: + real = client._low.raw_request + + def as_304(*args: Any, **kwargs: Any) -> httpx.Response: + real(*args, **kwargs) + return httpx.Response(304, headers={"ETag": '"x"'}) + + client._low.raw_request = as_304 # type: ignore[method-assign] + with pytest.raises(ComfyError): + client.models.schema(MODEL) + + +# --- malformed answers are invalid_response, never an untyped error -------- + + +@pytest.mark.parametrize("body", [[], "page", None]) +def test_a_non_object_catalog_page_is_invalid_response(server, body: Any) -> None: + server.state.catalog_pages[None] = body + with Comfy() as client, pytest.raises(ComfyError) as info: + client.models.list().page() + assert info.value.code == "invalid_response" + + +@pytest.mark.parametrize( + "entry", + [ + # No slash: parses as no run id at all. + {"id": "flux-2-pro", "provider": "bfl", "model": "flux-2-pro"}, + # A runnable id that names a different provider than the entry claims. + {"id": "acme/flux-2-pro", "provider": "bfl", "model": "flux-2-pro"}, + {"id": "bfl/flux-2-dev", "provider": "bfl", "model": "flux-2-pro"}, + ], +) +def test_an_entry_whose_id_does_not_address_its_provider_and_model_is_invalid_response( + server, entry: dict[str, Any] +) -> None: + server.state.catalog_pages[None]["data"] = [entry] + with Comfy() as client, pytest.raises(ComfyError) as info: + list(client.models.list()) + assert info.value.code == "invalid_response" + + +def test_the_last_page_never_hands_back_a_cursor(server) -> None: + # `c-3` is the last page and carries a stale cursor; paging by hand must + # not be told to follow it. + _three_pages(server) + with Comfy() as client: + page = client.models.list(cursor="c-3").page() + assert page.has_more is False + assert page.next_cursor is None + + +def test_catalog_entries_are_hashable(server) -> None: + _three_pages(server) + with Comfy() as client: + models = set(client.models.list()) + assert {m.id for m in models} == set(EXPECTED_IDS) + + +@pytest.mark.parametrize("document", [None, [], "doc"]) +def test_a_non_object_schema_200_is_invalid_response_not_unchanged(server, document: Any) -> None: + # `None` is how a 304 reaches the SDK, so a 200 carrying JSON `null` must + # not be read as "your copy is current". + server.state.schema_document = document + with Comfy(retry=NO_RETRY) as client, pytest.raises(ComfyError) as info: + client.models.schema(MODEL) + assert info.value.code == "invalid_response" + + +def test_a_304_without_an_etag_keeps_the_tag_that_was_sent(server) -> None: + with Comfy() as client: + real = client._low.raw_request + + def stripped(*args: Any, **kwargs: Any) -> httpx.Response: + real(*args, **kwargs) + return httpx.Response(304) + + client._low.raw_request = stripped # type: ignore[method-assign] + result = client.models.schema(MODEL, etag='"schema-v1"') + assert result.unchanged is True + assert result.etag == '"schema-v1"' + + +@pytest.mark.parametrize("bad", ["", '"café"']) +def test_schema_rejects_an_empty_or_non_ascii_etag_before_any_request(server, bad: str) -> None: + with Comfy() as client, pytest.raises(ValueError): + client.models.schema(MODEL, etag=bad) + assert server.state.schema_paths == [] + + +async def test_async_schema_rejects_an_empty_etag_before_any_request(server) -> None: + async with AsyncComfy() as client: + with pytest.raises(ValueError): + await client.models.schema(MODEL, etag="") + assert server.state.schema_paths == [] + + +def test_discovery_retries_a_read_timeout_under_the_default_policy(server, monkeypatch) -> None: + # A discovery read carries no key and bills nothing, so the possibly-in- + # flight class a run needs an opt-in for is retried here by default. + import comfy_sdk.model_catalog as catalog + + monkeypatch.setattr(catalog.time, "sleep", lambda _s: None) + with Comfy() as client: + real = client._low.raw_request + calls: list[int] = [] + + def flaky(*args: Any, **kwargs: Any) -> httpx.Response: + calls.append(1) + if len(calls) == 1: + raise httpx.ReadTimeout("slow") + return real(*args, **kwargs) + + client._low.raw_request = flaky # type: ignore[method-assign] + assert client.models.schema(MODEL).document == server.state.schema_document + assert len(calls) == 2 + + +def test_no_retry_still_means_one_discovery_attempt(server) -> None: + with Comfy(retry=NO_RETRY) as client: + calls: list[int] = [] + + def flaky(*args: Any, **kwargs: Any) -> httpx.Response: + calls.append(1) + raise httpx.ReadTimeout("slow") + + client._low.raw_request = flaky # type: ignore[method-assign] + with pytest.raises(httpx.ReadTimeout): + client.models.schema(MODEL) + assert len(calls) == 1 diff --git a/tests/test_router_spec_contract.py b/tests/test_router_spec_contract.py index efda9ed..1170412 100644 --- a/tests/test_router_spec_contract.py +++ b/tests/test_router_spec_contract.py @@ -12,7 +12,9 @@ ``post.operationId`` is ``runRouterModel``, and the ``servers[0].url`` it is addressed against -- compared against :data:`comfy_low.transport._MODEL_RUN_PATH_TEMPLATE` and - :data:`comfy_sdk.COMFY_ROUTER_BASE_URL`. + :data:`comfy_sdk.COMFY_ROUTER_BASE_URL` -- and, the same way, the two + discovery routes ``models.list()`` / ``models.schema()`` read, by their + ``get.operationId`` (``listRouterModels``, ``getRouterModelInputSchema``). Neither is generated, so a Router spec sync is the moment they can drift. The failures guarded against are a sync landing a new bucket that then reaches @@ -36,7 +38,12 @@ import pytest import yaml -from comfy_low.transport import _MODEL_RUN_PATH_TEMPLATE +from comfy_low.transport import ( + _MODEL_CATALOG_PATH, + _MODEL_RUN_PATH_TEMPLATE, + _MODEL_SCHEMA_PATH_TEMPLATE, + model_catalog_path, +) from comfy_sdk import COMFY_ROUTER_BASE_URL from comfy_sdk.models import _run_result from comfy_sdk.router_exceptions import ( @@ -265,6 +272,51 @@ def test_run_path_matches_vendored_spec() -> None: ) +def _declared_get_paths(operation_id: str) -> list[str]: + """Every path whose ``get.operationId`` is ``operation_id``, in spec order.""" + doc = yaml.safe_load(ROUTER_SPEC.read_text(encoding="utf-8")) + return [ + path + for path, item in (doc.get("paths") or {}).items() + if isinstance(item, dict) + and isinstance(item.get("get"), dict) + and item["get"].get("operationId") == operation_id + ] + + +@pytest.mark.parametrize( + ("operation_id", "bound", "constant"), + [ + ("listRouterModels", _MODEL_CATALOG_PATH, "_MODEL_CATALOG_PATH"), + ("getRouterModelInputSchema", _MODEL_SCHEMA_PATH_TEMPLATE, "_MODEL_SCHEMA_PATH_TEMPLATE"), + ], +) +def test_the_discovery_routes_match_the_vendored_spec( + operation_id: str, bound: str, constant: str +) -> None: + # `models.list()` / `models.schema()` are hand-bound exactly as the run + # route is, so a sync that moves either route has to fail here rather than + # ship an SDK that GETs a path the server no longer serves. + declared = _declared_get_paths(operation_id) + assert declared == [bound], ( + f"the vendored spec declares {operation_id} at {declared} and the SDK reads " + f"{bound!r} -- update comfy_low.transport.{constant}" + ) + + +def test_the_catalog_query_parameters_are_the_ones_the_spec_declares() -> None: + # `model_catalog_path` sends `cursor` and `limit` by name; a sync renaming + # either would otherwise be silently ignored by the server. + doc = yaml.safe_load(ROUTER_SPEC.read_text(encoding="utf-8")) + params = doc["paths"][_MODEL_CATALOG_PATH]["get"]["parameters"] + shared = doc["components"]["parameters"] + names = { + shared[p["$ref"].rsplit("/", 1)[-1]]["name"] if "$ref" in p else p["name"] for p in params + } + assert names == {"cursor", "limit"} + assert model_catalog_path("c", 5) == f"{_MODEL_CATALOG_PATH}?cursor=c&limit=5" + + def test_the_bound_path_has_exactly_the_two_segments_the_binding_fills() -> None: # `model_run_request` fills `{provider}` and `{model}` by name; a sync that # renamed or added a template variable would silently KeyError at call time diff --git a/tests/test_sync_async_parity.py b/tests/test_sync_async_parity.py index f241b06..1db1caa 100644 --- a/tests/test_sync_async_parity.py +++ b/tests/test_sync_async_parity.py @@ -106,6 +106,10 @@ ("comfy_sdk.jobs.Job", "get_outputs"): ( "reads outputs already on the handle — no re-fetch, so nothing to await" ), + ("comfy_sdk.models.Models", "list"): ( + "builds a lazy AsyncModelList and fetches nothing; `async for` and its awaited " + "page() are the I/O, so `async for m in client.models.list()` needs no extra await" + ), } #: A name that encodes sync-vs-async instead of letting the client encode it.