From 2767f73b807cc74a0e482cc572862cbf0084e71d Mon Sep 17 00:00:00 2001 From: jabdikadyr Date: Thu, 17 Sep 2026 17:22:57 +0600 Subject: [PATCH 1/4] fix(sdk): close conformance gaps in the agent, directory and service --- README.md | 3 +- scripts/verify-consumer.sh | 46 +- src/offering_protocol/agent/capabilities.py | 222 ++++--- src/offering_protocol/agent/client.py | 116 +++- src/offering_protocol/agent/details.py | 4 + src/offering_protocol/core/__init__.py | 8 + src/offering_protocol/core/validation.py | 211 +++++- src/offering_protocol/directory/__init__.py | 2 + src/offering_protocol/directory/addresses.py | 83 +++ src/offering_protocol/directory/client.py | 161 +++-- src/offering_protocol/directory/models.py | 14 + src/offering_protocol/directory/transport.py | 9 +- src/offering_protocol/service/service.py | 316 +++++++-- .../service/static_catalog.py | 5 + tests/test_agent.py | 13 +- tests/test_agent_conformance.py | 536 +++++++++++++++ tests/test_core_conformance.py | 371 +++++++++++ tests/test_directory.py | 17 +- tests/test_directory_conformance.py | 519 +++++++++++++++ tests/test_edges.py | 42 +- tests/test_service_conformance.py | 608 ++++++++++++++++++ 21 files changed, 3047 insertions(+), 259 deletions(-) create mode 100644 src/offering_protocol/directory/addresses.py create mode 100644 tests/test_agent_conformance.py create mode 100644 tests/test_core_conformance.py create mode 100644 tests/test_directory_conformance.py create mode 100644 tests/test_service_conformance.py diff --git a/README.md b/README.md index 623057e..1b71d4f 100644 --- a/README.md +++ b/README.md @@ -278,7 +278,8 @@ network and credential-isolation requirements. Local HTTP development is disable Attribute Schema resolution accepts JSON Schema Draft 2020-12 and is limited to 256 KiB per document, 16 documents, eight reference levels, and one MiB for the complete graph. OpenAPI -documents are limited to one MiB. These are fixed SDK safety ceilings. Linked schema documents must +documents are limited to one MiB and 32 levels of JSON nesting. Every other ODP response is limited +to 16 levels of nesting, and the Service Document to eight. These are fixed SDK safety ceilings. Linked schema documents must use HTTPS. Cross-document schema composition uses `$ref`; `$dynamicRef` accepts only a fragment reference such as `#node`. diff --git a/scripts/verify-consumer.sh b/scripts/verify-consumer.sh index 36ac723..6639fb4 100755 --- a/scripts/verify-consumer.sh +++ b/scripts/verify-consumer.sh @@ -5,7 +5,51 @@ repository=$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd) consumer=$(mktemp -d) trap 'rm -rf "$consumer"' EXIT -python=${ODP_PYTHON:-python3} +# The consumer venv has to be built with an interpreter this package actually supports. Left to +# whatever `python3` happens to be on PATH, an older one fails much further downstream: pip filters +# out every release of a dependency that needs a newer Python and reports "could not find a version +# that satisfies jsonschema", which says nothing about the interpreter that caused it. +minimum=$(sed -n 's/^requires-python *= *">=\([0-9.]*\)"/\1/p' "$repository/pyproject.toml") +: "${minimum:?pyproject.toml does not declare requires-python}" + +supports_minimum() { + "$1" -c "import sys; raise SystemExit(0 if sys.version_info >= tuple( + int(part) for part in '$minimum'.split('.')) else 1)" 2>/dev/null +} + +describe() { + command -v "$1" >/dev/null 2>&1 && + "$1" -c 'import platform; print(platform.python_version())' 2>/dev/null || + echo "not found" +} + +python="" +if [[ -n "${ODP_PYTHON:-}" ]]; then + # An interpreter named on purpose is used or refused, never quietly swapped for another one. + if ! supports_minimum "$ODP_PYTHON"; then + echo "ODP_PYTHON=$ODP_PYTHON is Python $(describe "$ODP_PYTHON")," \ + "and offering-protocol requires >=$minimum." >&2 + exit 1 + fi + python="$ODP_PYTHON" +else + for candidate in python3 python3.14 python3.13 python3.12 python3.11; do + if command -v "$candidate" >/dev/null 2>&1 && supports_minimum "$candidate"; then + python="$candidate" + break + fi + done + if [[ -z "$python" ]] && command -v uv >/dev/null 2>&1; then + python=$(uv python find ">=$minimum" 2>/dev/null || true) + fi + if [[ -z "$python" ]]; then + echo "offering-protocol requires Python >=$minimum;" \ + "python3 is $(describe python3) and no newer interpreter was found." >&2 + echo "Install one, or set ODP_PYTHON to an interpreter that satisfies it." >&2 + exit 1 + fi +fi + "$python" -m venv "$consumer/.venv" source=${ODP_CONSUMER_SOURCE:-wheel} if [[ "$source" == "wheel" ]]; then diff --git a/src/offering_protocol/agent/capabilities.py b/src/offering_protocol/agent/capabilities.py index 4dca744..9840dcf 100644 --- a/src/offering_protocol/agent/capabilities.py +++ b/src/offering_protocol/agent/capabilities.py @@ -2,20 +2,24 @@ from __future__ import annotations -from collections.abc import Iterable, Mapping +from collections.abc import Awaitable, Callable, Sequence from dataclasses import dataclass, field from enum import StrEnum -from urllib.parse import urljoin, urlsplit +from typing import Any, cast from offering_protocol.agent.client import AgentError, ServiceClient from offering_protocol.core import ( Collection, + FilterCapabilitySource, FilterDefinition, Operation, + ReferenceError, SearchCapabilities, + SortCapabilitySource, SortDefinition, parse_filter_definition_page, parse_sort_definition_page, + resolve_continuation, ) _MAXIMUM_CAPABILITY_PAGES = 16 @@ -122,32 +126,17 @@ async def _add_filters( scope: CapabilityScope, capabilities: SearchCapabilities, ) -> None: - if capabilities.filters is None: - return - try: - values = ( - await _load_filters(client, capabilities.filters.linked.href) - if capabilities.filters.linked is not None - else capabilities.filters.inline - ) - except AgentError as error: - result.issues.append(CapabilityIssue(CapabilityKind.FILTERS, str(error), scope)) - return - duplicates = _duplicates((item.id for item in values), result.filters) - for identifier in duplicates: - result.filters.pop(identifier, None) - accepted = [item for item in values if item.id not in duplicates] - if len(result.filters) + len(accepted) > _MAXIMUM_FILTERS: - result.issues.append( - CapabilityIssue( - CapabilityKind.FILTERS, - "Effective filters exceed 1024 entries", - scope, - ) - ) - return - result.filters.update((item.id, item) for item in accepted) - _report_duplicates(duplicates, CapabilityKind.FILTERS, scope, result.issues) + await _add_source( + client, + result, + CapabilityKind.FILTERS, + scope, + capabilities.filters, + result.filters, + _MAXIMUM_FILTERS, + _load_filters, + None, + ) async def _add_sorts( @@ -158,56 +147,109 @@ async def _add_sorts( scope: CapabilityScope, capabilities: SearchCapabilities, ) -> None: - if capabilities.sorts is None: + await _add_source( + client, + result, + CapabilityKind.SORTS, + scope, + capabilities.sorts, + target, + _MAXIMUM_SORTS, + _load_sorts, + scopes, + ) + + +async def _add_source( + client: ServiceClient, + result: SearchCapabilityCatalog, + kind: CapabilityKind, + scope: CapabilityScope, + source: FilterCapabilitySource | SortCapabilitySource | None, + target: dict[str, Any], + maximum: int, + load: Callable[[ServiceClient, str, int], Awaitable[list[Any]]], + scopes: dict[str, CapabilityScope] | None, +) -> None: + """Merges one capability source into the effective catalog, or reports why it cannot be. + + FLT-55 makes a source atomic: every page is retrieved and the source's own rules are enforced + before any of its definitions is exposed. So the work below decides everything first and writes + to `target` only once the source has been accepted -- an invalid or oversized source leaves the + sources merged before it exactly as they were. + """ + if source is None: return try: - values = ( - await _load_sorts(client, capabilities.sorts.linked.href) - if capabilities.sorts.linked is not None - else capabilities.sorts.inline + values: Sequence[Any] = ( + await load(client, source.linked.href, maximum - len(target)) + if source.linked is not None + else list(source.inline) ) + # FLT-55: a source enforces its own uniqueness, so an identifier this source publishes + # twice makes the whole source unusable -- there is no basis for choosing between them. + identifiers: set[str] = set() + for value in values: + if value.id in identifiers: + raise AgentError(f"Duplicate {kind.value} identifier {value.id} within one source") + identifiers.add(value.id) + # An identifier two effective sources both publish is quarantined instead: neither copy + # wins, and nothing else about either source is affected. + shared = sorted(identifier for identifier in identifiers if identifier in target) + accepted = [value for value in values if value.id not in shared] + # FLT-62: the bound is checked before anything is written, so a source that overflows the + # effective catalog cannot take the earlier valid sources down with it. + if len(target) - len(shared) + len(accepted) > maximum: + raise AgentError(f"Effective {kind.value} exceed their limit") except AgentError as error: - result.issues.append(CapabilityIssue(CapabilityKind.SORTS, str(error), scope)) + result.issues.append(CapabilityIssue(kind, str(error), scope)) return - duplicates = _duplicates((item.id for item in values), target) - for identifier in duplicates: + for identifier in shared: target.pop(identifier, None) - scopes.pop(identifier, None) - accepted = [item for item in values if item.id not in duplicates] - if len(target) + len(accepted) > _MAXIMUM_SORTS: + if scopes is not None: + scopes.pop(identifier, None) + for value in accepted: + target[value.id] = value + if scopes is not None: + scopes[value.id] = scope + if shared: result.issues.append( - CapabilityIssue(CapabilityKind.SORTS, "Effective sorts exceed 128 entries", scope) + CapabilityIssue(kind, f"Duplicate {kind.value}: {', '.join(shared)}", scope) ) - return - target.update((item.id, item) for item in accepted) - scopes.update((item.id, scope) for item in accepted) - _report_duplicates(duplicates, CapabilityKind.SORTS, scope, result.issues) -async def _load_filters(client: ServiceClient, reference: str) -> list[FilterDefinition]: - values: list[FilterDefinition] = [] - next_reference = reference - visited: set[str] = set() - for _ in range(_MAXIMUM_CAPABILITY_PAGES): - if not next_reference: - return values - target = _resolve_reference(next_reference, client.service_origin) - if target in visited: - raise AgentError("ODP capability pagination loop detected") - visited.add(target) - body = await client._linked_odp( - target, client._cache_fallbacks.collection, parse_filter_definition_page - ) - page = parse_filter_definition_page(body) - values.extend(page.items) - next_reference = page.next - if next_reference: - raise AgentError("ODP capability source exceeded 16 pages") - return values +async def _load_filters( + client: ServiceClient, reference: str, budget: int = _MAXIMUM_FILTERS +) -> list[FilterDefinition]: + values = await _load_definitions( + client, reference, budget, parse_filter_definition_page, CapabilityKind.FILTERS + ) + return cast("list[FilterDefinition]", values) + + +async def _load_sorts( + client: ServiceClient, reference: str, budget: int = _MAXIMUM_SORTS +) -> list[SortDefinition]: + values = await _load_definitions( + client, reference, budget, parse_sort_definition_page, CapabilityKind.SORTS + ) + return cast("list[SortDefinition]", values) -async def _load_sorts(client: ServiceClient, reference: str) -> list[SortDefinition]: - values: list[SortDefinition] = [] +async def _load_definitions( + client: ServiceClient, + reference: str, + budget: int, + parse: Callable[[bytes | str], Any], + kind: CapabilityKind, +) -> list[Any]: + """Retrieves a complete linked source, one page at a time. + + The budget is what the effective catalog has left. FLT-58 asks the Agent to stop retrieving a + source once it cannot fit, so the budget is checked as each page arrives rather than after the + whole source has been buffered: a source that can never fit costs one page, not sixteen. + """ + values: list[Any] = [] next_reference = reference visited: set[str] = set() for _ in range(_MAXIMUM_CAPABILITY_PAGES): @@ -217,11 +259,11 @@ async def _load_sorts(client: ServiceClient, reference: str) -> list[SortDefinit if target in visited: raise AgentError("ODP capability pagination loop detected") visited.add(target) - body = await client._linked_odp( - target, client._cache_fallbacks.collection, parse_sort_definition_page - ) - page = parse_sort_definition_page(body) + body = await client._linked_odp(target, client._cache_fallbacks.collection, parse) + page = parse(body) values.extend(page.items) + if len(values) > budget: + raise AgentError(f"Effective {kind.value} exceed their limit") next_reference = page.next if next_reference: raise AgentError("ODP capability source exceeded 16 pages") @@ -229,34 +271,14 @@ async def _load_sorts(client: ServiceClient, reference: str) -> list[SortDefinit def _resolve_reference(reference: str, origin: str) -> str: - target = urljoin(f"{origin}/", reference) - parsed = urlsplit(target) - if parsed.scheme not in {"http", "https"} or parsed.hostname is None: - raise AgentError("ODP capability reference must use HTTP or HTTPS") - return target - - -def _duplicates(values: Iterable[str], existing: Mapping[str, object]) -> set[str]: - seen: set[str] = set() - duplicates: set[str] = set() - for value in values: - if value in seen or value in existing: - duplicates.add(value) - seen.add(value) - return duplicates - + """Resolves a linked capability reference, which stays on the Service that advertised it. -def _report_duplicates( - duplicates: set[str], - kind: CapabilityKind, - scope: CapabilityScope, - issues: list[CapabilityIssue], -) -> None: - if duplicates: - issues.append( - CapabilityIssue( - kind, - f"Duplicate {kind.value}: {', '.join(sorted(duplicates))}", - scope, - ) - ) + FLT-52 makes `href` a same-origin Resource Reference and FLT-53 puts every `next` under the + common continuation contract, so both are resolved the way a page continuation is. Without + that, a Service Document written by somebody else could send this Agent's ODP requests to a + host of its choosing. + """ + try: + return resolve_continuation(reference, origin) + except ReferenceError as error: + raise AgentError(str(error)) from error diff --git a/src/offering_protocol/agent/client.py b/src/offering_protocol/agent/client.py index 717a03d..12a1c54 100644 --- a/src/offering_protocol/agent/client.py +++ b/src/offering_protocol/agent/client.py @@ -6,7 +6,7 @@ import json from collections.abc import Mapping from dataclasses import dataclass, replace -from datetime import datetime, timedelta +from datetime import UTC, datetime, timedelta from email.utils import parsedate_to_datetime from enum import StrEnum from typing import TYPE_CHECKING, cast @@ -30,21 +30,21 @@ resolve_continuation, ) from offering_protocol.core import ( - parse_collection as parse_collection_strict, + parse_agent_collection as parse_agent_collection_strict, ) from offering_protocol.core import ( - parse_collection_page as parse_collection_page_strict, + parse_agent_collection_page as parse_agent_collection_page_strict, ) from offering_protocol.core import ( - parse_offering as parse_offering_strict, + parse_agent_offering as parse_agent_offering_strict, ) from offering_protocol.core import ( - parse_offering_page as parse_offering_page_strict, + parse_agent_offering_page as parse_agent_offering_page_strict, ) from offering_protocol.core import ( parse_problem_response as parse_problem_response_strict, ) -from offering_protocol.core.validation import _normalize_agent_response +from offering_protocol.core.validation import _agent_body from offering_protocol.directory.transport import ( HttpRequest, HttpResponse, @@ -60,7 +60,11 @@ MEDIA_TYPE = "application/odp+json" _MAXIMUM_DOCUMENT_BYTES = 65_536 _MAXIMUM_RESOURCE_BYTES = 524_288 +_MAXIMUM_PROBLEM_BYTES = 16_384 _MAXIMUM_REDIRECTS = 5 +# ERR-21: a Service Document nests no deeper than 8 containers, every other ODP document 16. +_MAXIMUM_DOCUMENT_DEPTH = 8 +_MAXIMUM_DEPTH = 16 class Freshness(StrEnum): @@ -149,6 +153,7 @@ async def inspect(self) -> Inspection: _MAXIMUM_DOCUMENT_BYTES, self._cache_fallbacks.service_document, parse_agent_service_document, + _MAXIMUM_DOCUMENT_DEPTH, ) return Inspection( document=parse_agent_service_document(response.body), @@ -384,6 +389,7 @@ async def _request_cached( maximum_bytes: int, fallback: timedelta, parser: object, + maximum_depth: int = _MAXIMUM_DEPTH, ) -> _FetchedResponse: key = self._cache_key(method, target, body) cached = self._cache.get(key) @@ -415,7 +421,7 @@ async def _request_cached( record = replace(cached, expires=expires, final_url=final_url, stored=now) self._cache.set(key, record) return _FetchedResponse(record.body, record.final_url, Freshness.REVALIDATED) - response = _consume(response, maximum_bytes) + response = _consume(response, maximum_bytes, maximum_depth) _invoke_parser(parser, response.body) if _cacheable(method, response.headers, fallback): self._cache.set( @@ -476,6 +482,7 @@ async def _supporting_json( accept: str, media_types: set[str], maximum_bytes: int, + maximum_depth: int | None = None, ) -> dict[str, object]: current = target if not _is_https_url(current): @@ -538,6 +545,8 @@ async def _supporting_json( content_type = response.headers.get("content-type", "").split(";", 1)[0].strip().lower() if content_type not in media_types: raise AgentError("ODP supporting document has an unsupported media type") + if maximum_depth is not None: + _require_depth(response.body, maximum_depth, "ODP supporting document") value = _decode_json_object(response.body) if _cacheable("GET", response.headers, timedelta()): self._cache.set( @@ -584,30 +593,23 @@ def _operation_parser(operation: Operation) -> object: def parse_collection(data: bytes | str) -> Collection: - return parse_collection_strict(_normalize_body(data, "collection")) + return parse_agent_collection_strict(data) def parse_offering(data: bytes | str) -> Offering: - return parse_offering_strict(_normalize_body(data, "offering")) + return parse_agent_offering_strict(data) def parse_collection_page(data: bytes | str) -> Page[Collection]: - return parse_collection_page_strict(_normalize_body(data, "collection-page")) + return parse_agent_collection_page_strict(data) def parse_offering_page(data: bytes | str) -> OfferingPage[Offering]: - return parse_offering_page_strict(_normalize_body(data, "offering-page")) + return parse_agent_offering_page_strict(data) def parse_problem_response(data: bytes | str, status: int) -> ProblemDetails: - return parse_problem_response_strict(_normalize_body(data, "problem"), status) - - -def _normalize_body(data: bytes | str, kind: str) -> str: - raw = json.loads(data) - if not isinstance(raw, dict): - return data.decode() if isinstance(data, bytes) else data - return json.dumps(_normalize_agent_response(raw, kind), separators=(",", ":")) + return parse_problem_response_strict(_agent_body(data, "problem"), status) def _encode(value: object) -> bytes: @@ -634,27 +636,78 @@ def _decode_json_object(data: bytes) -> dict[str, object]: value = json.loads(data) except (UnicodeDecodeError, json.JSONDecodeError) as error: raise AgentError(f"ODP supporting document is invalid JSON: {error}") from error + except RecursionError as error: + raise AgentError("ODP supporting document is nested too deeply") from error if not isinstance(value, dict) or any(not isinstance(key, str) for key in value): raise AgentError("ODP supporting document must be a JSON object") return cast(dict[str, object], value) -def _consume(response: HttpResponse, maximum_bytes: int) -> HttpResponse: +def _consume(response: HttpResponse, maximum_bytes: int, maximum_depth: int) -> HttpResponse: + if not 200 <= response.status < 300: + raise ServiceRequestError(response.status, _problem_message(response), response.headers) if len(response.body) > maximum_bytes: raise AgentError("ODP response exceeds its byte limit") - if not 200 <= response.status < 300: - try: - problem = parse_problem_response(response.body, response.status) - message = problem.detail or problem.title - except ValueError: - message = response.body.decode(errors="replace") - raise ServiceRequestError(response.status, message, response.headers) content_type = response.headers.get("content-type", "").split(";", 1)[0].strip().lower() if content_type != MEDIA_TYPE: raise AgentError(f"ODP response must use {MEDIA_TYPE}") + _require_depth(response.body, maximum_depth, "ODP response") return response +def _problem_message(response: HttpResponse) -> str: + """Describes a refused request, reading the body only within the Problem Details limit. + + ERR-21 budgets a Problem Details response at 16,384 bytes, so a larger body is not a Problem + Details document this Agent will read. The HTTP status still describes the failure, which is + why an oversized error body reports the status rather than a byte-limit error. + """ + if len(response.body) > _MAXIMUM_PROBLEM_BYTES: + return f"ODP request failed with HTTP {response.status}" + try: + problem = parse_problem_response(response.body, response.status) + except ValueError: + return response.body.decode(errors="replace") + return problem.detail or problem.title + + +def _require_depth(body: bytes, maximum: int, subject: str) -> None: + """Rejects a document nested deeper than ERR-21 allows. + + A malformed body is left alone here so the document parser reports it in its own words. Depth + is counted the way ERR-18 measures it, from the top-level value, and the walk keeps its own + stack: a recursive one would exhaust the interpreter on exactly the documents this limit exists + to refuse. + """ + try: + value = json.loads(body) + except (RecursionError, UnicodeDecodeError, ValueError): + return + if _nesting_depth(value) > maximum: + raise AgentError(f"{subject} exceeds its nesting-depth limit") + + +def _nesting_depth(value: object) -> int: + """Counts container nesting from the top-level value. + + `{}` and `{"a": 1}` are both depth 1 and `{"a": {"b": 1}}` is depth 2: a scalar is a value a + container holds, not a level of its own. + """ + maximum = 0 + pending: list[tuple[int, object]] = [(1, value)] + while pending: + depth, current = pending.pop() + if isinstance(current, dict): + children: list[object] = list(current.values()) + elif isinstance(current, list): + children = list(current) + else: + continue + maximum = max(maximum, depth) + pending.extend((depth + 1, child) for child in children) + return maximum + + def _traversal_bounds(options: TraversalOptions) -> tuple[int, int]: if not 1 <= options.max_items <= 10_000 or not 1 <= options.max_pages <= 16: raise AgentError("traversal exceeds 10000 items or 16 pages") @@ -699,8 +752,13 @@ def _expiration(headers: dict[str, str], fallback: timedelta, now: datetime) -> except (KeyError, ValueError): pass if "expires" in headers: + # RFC 9111 5.3: an `Expires` a cache cannot read -- "0" above all -- names a time in the + # past, so an unreadable one expires the entry instead of granting it the fallback + # lifetime. A date written with the "-0000" zone parses to a naive value, which cannot be + # compared with the aware clock this cache keeps, so it is read as the UTC it means. try: - return parsedate_to_datetime(headers["expires"]) + expires = parsedate_to_datetime(headers["expires"]) except (TypeError, ValueError): - pass + return now + return expires if expires.tzinfo is not None else expires.replace(tzinfo=UTC) return now + fallback diff --git a/src/offering_protocol/agent/details.py b/src/offering_protocol/agent/details.py index e2ffddc..142bd1a 100644 --- a/src/offering_protocol/agent/details.py +++ b/src/offering_protocol/agent/details.py @@ -17,6 +17,9 @@ ) _MAXIMUM_OPENAPI_BYTES = 1_048_576 +# OFR-73: an OpenAPI Action document nests deeper than an ODP document, so it has its own +# allowance rather than the 16 ERR-21 gives every other retrieved document. +_MAXIMUM_OPENAPI_DEPTH = 32 class OfferingIssueScope(StrEnum): @@ -120,6 +123,7 @@ async def resolve_action(client: ServiceClient, offering_id: str, action_id: str "application/vnd.oai.openapi+json;version=3.1, application/json;q=0.9", {"application/vnd.oai.openapi+json", "application/json"}, _MAXIMUM_OPENAPI_BYTES, + _MAXIMUM_OPENAPI_DEPTH, ) version = document.get("openapi") if not isinstance(version, str) or not version.startswith("3.1."): diff --git a/src/offering_protocol/core/__init__.py b/src/offering_protocol/core/__init__.py index 65cb8d6..cc72ae2 100644 --- a/src/offering_protocol/core/__init__.py +++ b/src/offering_protocol/core/__init__.py @@ -16,6 +16,10 @@ from offering_protocol.core.validation import ( OdpValidationError, ValidationIssue, + parse_agent_collection, + parse_agent_collection_page, + parse_agent_offering, + parse_agent_offering_page, parse_agent_service_document, parse_collection, parse_collection_page, @@ -45,6 +49,10 @@ "is_local_resource_identifier", "operation_method", "operation_path", + "parse_agent_collection", + "parse_agent_collection_page", + "parse_agent_offering", + "parse_agent_offering_page", "parse_agent_service_document", "parse_collection", "parse_collection_page", diff --git a/src/offering_protocol/core/validation.py b/src/offering_protocol/core/validation.py index 5e8c825..7a10a50 100644 --- a/src/offering_protocol/core/validation.py +++ b/src/offering_protocol/core/validation.py @@ -27,7 +27,9 @@ OfferingSearchRequest, Operation, Page, + PriceType, ProblemDetails, + RefinementGroup, ResourceIdentity, ServiceDocument, SortDefinition, @@ -440,6 +442,23 @@ def _filter_agent_protocol_category( def parse_collection(data: bytes | str) -> Collection: + value = _read_collection(data) + _raise_refinement("Collection", _collection_issues(value)) + return value + + +def parse_agent_collection(data: bytes | str) -> Collection: + """Reads a Collection an Agent received. + + ROLE-03: discovery metadata an Agent can still use is not withheld over a defect it can work + around, so the invariants `_collection_issues` states -- which describe a hierarchy the Agent + simply does not walk -- do not refuse the document here. A Service, which MUST NOT publish one, + goes through `parse_collection`. + """ + return _read_collection(_agent_body(data, "collection")) + + +def _read_collection(data: bytes | str) -> Collection: value = _parse(data, "collection.schema.json", "Collection", Collection) _validate_representation( value.language, value.localizations, [image.src for image in value.images] @@ -447,7 +466,33 @@ def parse_collection(data: bytes | str) -> Collection: return value +def _collection_issues(value: Collection) -> list[ValidationIssue]: + """The Collection rules a JSON Schema cannot state: they compare one member against another.""" + # COL-20: a Collection naming itself as a parent is a one-node cycle, so anything walking the + # hierarchy upwards from it would never reach a root. + if value.id in value.parent_ids: + return [_issue("/parent_ids", "self-parent", "must not name the Collection itself")] + return [] + + def parse_offering(data: bytes | str) -> Offering: + value = _read_offering(data) + _raise_refinement("Offering", _offering_issues(value)) + return value + + +def parse_agent_offering(data: bytes | str) -> Offering: + """Reads an Offering an Agent received. + + ROLE-03: a defect an Agent can describe to its caller is a note about that Offering rather than + a reason to discard it, so the invariants `_offering_issues` states are left for the Agent to + report against the Actions or price they concern. A Service, which MUST NOT publish one, goes + through `parse_offering`. + """ + return _read_offering(_agent_body(data, "offering")) + + +def _read_offering(data: bytes | str) -> Offering: value = _parse(data, "offering.schema.json", "Offering", Offering) _validate_representation( value.language, value.localizations, [image.src for image in value.images] @@ -455,6 +500,29 @@ def parse_offering(data: bytes | str) -> Offering: return value +def _offering_issues(value: Offering) -> list[ValidationIssue]: + """The Offering rules a JSON Schema cannot state: they compare one member against another.""" + issues: list[ValidationIssue] = [] + # OFR-57: an Action identifier is unique within its Offering, so a repeat leaves a caller + # unable to say which Action it meant. + identifiers = [action.id for action in value.actions] + if len(identifiers) != len(set(identifiers)): + issues.append( + _issue("/actions", "unique-action-id", "must contain unique Action identifiers") + ) + # OFR-49: a range whose minimum is above its maximum describes no price at all. + price = value.price + if ( + price is not None + and price.price_type is PriceType.RANGE + and _compare_decimals(price.minimum, price.maximum) > 0 + ): + issues.append( + _issue("/price/minimum", "price-range", "must be less than or equal to maximum") + ) + return issues + + def parse_problem_details(data: bytes | str) -> ProblemDetails: value = _parse(data, "problem-details.schema.json", "Problem Details", ProblemDetails) expected_type = "https://offeringprotocol.org/problems/" + value.code.lower().replace("_", "-") @@ -487,7 +555,37 @@ def parse_collection_page(data: bytes | str) -> Page[Collection]: return value +def parse_agent_collection_page(data: bytes | str) -> Page[Collection]: + """Reads a page of Collections an Agent received, item by item, tolerantly.""" + value = _parse( + _agent_body(data, "collection-page"), + "page-envelope.schema.json", + "Collection page", + Page[Collection], + ) + for item in value.items: + _read_collection(_embedded_json(item, value.odp_version)) + return value + + def parse_offering_page(data: bytes | str) -> OfferingPage[Offering]: + value = _read_offering_page(data) + _raise_refinement("Offering page", _refinement_issues(value.refinements)) + for item in value.items: + _raise_refinement("Offering", _offering_issues(item)) + return value + + +def parse_agent_offering_page(data: bytes | str) -> OfferingPage[Offering]: + """Reads a page of Offerings an Agent received. + + A Refinement Group the Agent cannot use is no reason to discard the Offering results beside it, + so the invariants `_refinement_issues` states do not refuse the page here. + """ + return _read_offering_page(_agent_body(data, "offering-page")) + + +def _read_offering_page(data: bytes | str) -> OfferingPage[Offering]: value = _parse( data, "offering-search-response.schema.json", @@ -495,10 +593,40 @@ def parse_offering_page(data: bytes | str) -> OfferingPage[Offering]: OfferingPage[Offering], ) for item in value.items: - parse_offering(_embedded_json(item, value.odp_version)) + _read_offering(_embedded_json(item, value.odp_version)) return value +def _refinement_issues(groups: list[RefinementGroup]) -> list[ValidationIssue]: + """The Refinement rules a JSON Schema cannot state: they compare one member against another.""" + issues: list[ValidationIssue] = [] + # FLT-30: `filter_id` is unique among the returned groups, so a repeat leaves an Agent unable + # to say which group belongs to that Filter Definition. + identifiers = [group.filter_id for group in groups] + if len(identifiers) != len(set(identifiers)): + issues.append( + _issue( + "/refinements", + "unique-filter-id", + "must contain unique Refinement Group identifiers", + ) + ) + # FLT-32: bucket values are unique within a group. The schema's `uniqueItems` compares whole + # buckets, so it passes two buckets that name one value with differing counts -- exactly the + # case that leaves an Agent with two counts for the same candidate and no way to choose. + for index, group in enumerate(groups): + keys = [_bucket_key(bucket.value) for bucket in group.values] + if len(keys) != len(set(keys)): + issues.append( + _issue( + f"/refinements/{index}/values", + "unique-bucket-value", + "must contain unique Refinement Bucket values", + ) + ) + return issues + + def parse_collection_search_request(data: bytes | str) -> CollectionSearchRequest: return _parse( data, @@ -536,14 +664,30 @@ def parse_filter_definition(data: bytes | str) -> FilterDefinition: "contains an operator incompatible with the Filter type", ) ) - if value.filter_type is FilterType.BOOLEAN and value.unit is not None: - issues.append(_issue("/unit", "unit-type", "must not appear on a boolean Filter")) + # FLT-10: only a numeric Filter carries a unit. A unit on a string, date, date-time or boolean + # Filter describes a dimension its values do not have, so a caller reading it would convert or + # label values that were never quantities. + if value.unit is not None and value.filter_type not in { + FilterType.DECIMAL, + FilterType.INTEGER, + FilterType.NUMBER, + }: + issues.append(_issue("/unit", "unit-type", "must not appear on a non-numeric Filter")) _raise_refinement("Filter Definition", issues) return value def parse_sort_definition(data: bytes | str) -> SortDefinition: - return _parse(data, "sort-definition.schema.json", "Sort Definition", SortDefinition) + value = _parse(data, "sort-definition.schema.json", "Sort Definition", SortDefinition) + # FLT-41: every `filter_id` in a Sort Definition is distinct. Ordering by one Filter twice + # cannot change the order, so a repeat describes a recipe that does not mean what it says. + identifiers = [key.filter_id for key in value.keys] + if len(identifiers) != len(set(identifiers)): + _raise_refinement( + "Sort Definition", + [_issue("/keys", "unique-filter-id", "must order by each Filter at most once")], + ) + return value def parse_filter_definition_page(data: bytes | str) -> Page[FilterDefinition]: @@ -745,6 +889,65 @@ def _is_language_tag(value: str) -> bool: return not in_extension or len(subtags[-1]) > 1 +_DECIMAL = re.compile(r"^-?(?:0|[1-9][0-9]*)(?:\.[0-9]+)?$") + + +def _compare_decimals(left: str, right: str) -> int: + """Orders two ODP monetary values, which are decimal strings rather than JSON numbers (OFR-48). + + Decimal equality and ordering are numeric rather than lexical, so `1.0` equals `1.00` and + `9.00` sits below `10.00`. The comparison is done digit by digit: parsing to a float would lose + the precision the decimal form exists to keep. + """ + left_whole, left_fraction = _split_decimal(left) + right_whole, right_fraction = _split_decimal(right) + for ours, theirs in ( + (len(left_whole), len(right_whole)), + (left_whole, right_whole), + (left_fraction, right_fraction), + ): + if ours != theirs: + return 1 if ours > theirs else -1 # type: ignore[operator] + return 0 + + +def _split_decimal(value: str) -> tuple[str, str]: + whole, _, fraction = value.partition(".") + return whole.lstrip("0"), fraction.rstrip("0") + + +def _bucket_key(value: object) -> str: + """Compares two bucket values the way the referenced Filter Definition would. + + A response does not carry its Filter Definitions, so the type behind a JSON string -- `string`, + `decimal`, `date` or `date-time` -- is not known here. Every one of those compares two strings + exactly except `decimal`, whose equality is numeric, so a string that can only be a decimal is + reduced to one spelling per value. JSON numbers compare numerically. + """ + if isinstance(value, bool): + return f"b{value}" + if isinstance(value, int | float): + return f"n{float(value)}" + if isinstance(value, str): + if _DECIMAL.match(value): + whole, fraction = _split_decimal(value.removeprefix("-")) + sign = "-" if value.startswith("-") and (whole or fraction) else "" + return f"d{sign}{whole}.{fraction}" + return f"s{value}" + return f"o{json.dumps(value, separators=(',', ':'), sort_keys=True)}" + + +def _agent_body(data: bytes | str, kind: str) -> str: + """Applies the Agent's forward-compatibility filtering, leaving a body it can then validate.""" + try: + raw = json.loads(data) + except (UnicodeDecodeError, json.JSONDecodeError): + return data.decode(errors="replace") if isinstance(data, bytes) else data + if not isinstance(raw, dict): + return json.dumps(raw, separators=(",", ":")) + return json.dumps(_normalize_agent_response(raw, kind), separators=(",", ":")) + + def _issue(path: str, keyword: str, message: str) -> ValidationIssue: return ValidationIssue(path=path, keyword=keyword, message=message) diff --git a/src/offering_protocol/directory/__init__.py b/src/offering_protocol/directory/__init__.py index aa04856..0dfe585 100644 --- a/src/offering_protocol/directory/__init__.py +++ b/src/offering_protocol/directory/__init__.py @@ -23,6 +23,7 @@ SearchRequest, SearchResponse, ServiceFilters, + ServiceIssue, ServiceReference, ServiceResult, SuggestionRequest, @@ -60,6 +61,7 @@ "SearchRequest", "SearchResponse", "ServiceFilters", + "ServiceIssue", "ServiceReference", "ServiceResult", "SuggestionRequest", diff --git a/src/offering_protocol/directory/addresses.py b/src/offering_protocol/directory/addresses.py new file mode 100644 index 0000000..1167f46 --- /dev/null +++ b/src/offering_protocol/directory/addresses.py @@ -0,0 +1,83 @@ +"""SEC-08: the addresses the public internet does not route. + +A Directory result, a Service Origin, and every reference inside a Service's documents are written +by somebody else. Judging a resolved address against the IANA special-purpose registries is how an +SDK keeps a name a third party controls from naming an address its own network treats as internal. + +`ipaddress.is_global` is close but not complete: it reports `64:ff9b::a9fe:a9fe` as global, and that +address reaches link-local 169.254.169.254 through a NAT64 gateway -- the cloud metadata endpoint. +It also passes `5f00::/16`, `2620:4f:8000::/48`, `fec0::/10`, `ff00::/8`, and four IPv4 ranges. The +table below is the registry itself, so the answer does not depend on which ranges a standard library +happens to know about. +""" + +from __future__ import annotations + +from ipaddress import ( + IPv4Address, + IPv4Network, + IPv6Address, + IPv6Network, + ip_address, + ip_network, +) + +IPAddress = IPv4Address | IPv6Address + +# RFC 6890 and its successors: the addresses the public internet does not route. +_NON_PUBLIC: tuple[IPv4Network | IPv6Network, ...] = tuple( + ip_network(prefix) + for prefix in ( + "0.0.0.0/8", + "10.0.0.0/8", + "100.64.0.0/10", + "127.0.0.0/8", + "169.254.0.0/16", + "172.16.0.0/12", + "192.0.0.0/24", + "192.0.2.0/24", + "192.31.196.0/24", + "192.88.99.0/24", + "192.168.0.0/16", + "192.175.48.0/24", + "198.18.0.0/15", + "198.51.100.0/24", + "203.0.113.0/24", + "224.0.0.0/4", + "240.0.0.0/4", + "::/96", + # Each transition range embeds an IPv4 address, so a public-looking one can still be + # internal: `64:ff9b::a9fe:a9fe` and `2002:a9fe:a9fe::1` both reach 169.254.169.254. + "64:ff9b::/96", + "64:ff9b:1::/48", + "100::/64", + "2001::/32", + "2001:2::/48", + "2001:3::/32", + "2001:4:112::/48", + "2001:10::/28", + "2001:20::/28", + "2001:30::/28", + "2001:db8::/32", + "2002::/16", + "2620:4f:8000::/48", + "5f00::/16", + "fc00::/7", + "fe80::/10", + "fec0::/10", + "ff00::/8", + ) +) + + +def is_public(address: IPAddress) -> bool: + """True when the address falls in no special-purpose range, so the public internet routes it.""" + value = _unmap(address) + return not any(value in prefix for prefix in _NON_PUBLIC if prefix.version == value.version) + + +def _unmap(address: IPAddress) -> IPAddress: + """An IPv4-mapped IPv6 address is the IPv4 address it carries, and is judged as one.""" + if isinstance(address, IPv6Address) and address.ipv4_mapped is not None: + return ip_address(address.ipv4_mapped) + return address diff --git a/src/offering_protocol/directory/client.py b/src/offering_protocol/directory/client.py index 96a975d..f353474 100644 --- a/src/offering_protocol/directory/client.py +++ b/src/offering_protocol/directory/client.py @@ -3,15 +3,20 @@ from __future__ import annotations import json -from urllib.parse import urlencode, urljoin +import re +from ipaddress import ip_address +from urllib.parse import urlencode, urljoin, urlsplit from pydantic import ValidationError as ModelValidationError from offering_protocol.core import ( OdpValidationError, + ReferenceError, + ServiceDocument, derive_service_origin, parse_agent_service_document, ) +from offering_protocol.directory.addresses import is_public from offering_protocol.directory.models import ( DirectoryService, Environment, @@ -20,6 +25,7 @@ SearchPage, SearchRequest, SearchResponse, + ServiceIssue, SuggestionRequest, ) from offering_protocol.directory.results import parse_search_response @@ -32,6 +38,23 @@ _MAXIMUM_REDIRECTS = 5 _MAXIMUM_RESPONSE_BYTES = 524_288 +_MAXIMUM_ERROR_CHARACTERS = 2_048 +_RFC_3339 = re.compile( + r"^\d{4}-\d{2}-\d{2}[Tt]\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:[Zz]|[+-]\d{2}:\d{2})$" +) +#: Service Document members a Directory result may echo that this client does not validate. +#: +#: They are dropped rather than passed through, because a caller reading them off a record has no +#: way to tell they were never checked. `http` is the one that matters most: a caller could build +#: request URLs from an `endpoint_base` a Directory made up. +_UNVERIFIED_MEMBERS = ( + "branding", + "http", + "mcp", + "odp_version", + "payment_origins", + "search_capabilities", +) class DirectoryError(RuntimeError): @@ -98,7 +121,9 @@ def _continuation_url(self, next_reference: str) -> str: if not next_reference.strip(): raise DirectoryError("Directory continuation is empty") target = urljoin(f"{self.environment.origin}/", next_reference) - if derive_service_origin(target) != self.environment.origin: + # The continuation was written by the Directory, so a reference that is not a URL at all is + # a Directory failure to report rather than an exception from the URL parser. + if _origin_of(target, "Directory continuation") != self.environment.origin: raise DirectoryError("Directory continuation changed canonical origin") return target @@ -192,7 +217,7 @@ async def _request(self, method: str, target: str, body: bytes = b"") -> HttpRes if location is None: raise DirectoryError("Directory redirect omitted Location") next_target = urljoin(target, location) - if derive_service_origin(next_target) != derive_service_origin(target): + if _origin_of(next_target, "Directory redirect") != _origin_of(target, "Directory"): raise DirectoryError("Directory redirect changed origin") if response.status == 303 or (response.status in {301, 302} and method == "POST"): method = "GET" @@ -209,54 +234,111 @@ def _parse_mixed_response(body: bytes) -> SearchResponse: def _parse_search_page(body: bytes) -> SearchPage: + """Reads a Directory search page, keeping every record this client can read. + + ROLE-03: a Directory is a discovery aid, not an authority. One stale or nonconformant record + used to reject the page, which made every other Service in the result undiscoverable; it is now + dropped into `issues` and the rest of the page is handed back. + """ try: raw = json.loads(body) - if isinstance(raw, dict) and isinstance(raw.get("items"), list): - for item in raw["items"]: - if isinstance(item, dict) and "protocols" in item: - _normalize_service_protocols(item) - page = SearchPage.model_validate(raw) - except ( - UnicodeDecodeError, - json.JSONDecodeError, - ModelValidationError, - OdpValidationError, - ) as error: + except (UnicodeDecodeError, json.JSONDecodeError) as error: raise DirectoryError(f"invalid Directory response: {error}") from error - if len(page.items) > 100: + if not isinstance(raw, dict) or not isinstance(raw.get("items"), list): + raise DirectoryError("invalid Directory response: search page has no items") + if len(raw["items"]) > 100: raise DirectoryError("Directory search page exceeds 100 Services") + items: list[dict[str, object]] = [] + issues: list[ServiceIssue] = [] + for index, entry in enumerate(raw["items"]): + try: + items.append(_read_service(entry)) + except (ModelValidationError, OdpValidationError, ReferenceError, DirectoryError) as error: + issues.append(ServiceIssue(index=index, message=str(error))) + try: + page = SearchPage.model_validate({**raw, "items": items}) + except (ModelValidationError, OdpValidationError) as error: + raise DirectoryError(f"invalid Directory response: {error}") from error if page.facets is not None and any( facet.value.name.value != "tap" for facet in page.facets.trust ): raise DirectoryError("Directory trust facets are invalid") - for service in page.items: - if derive_service_origin(service.service_origin) != service.service_origin: - raise DirectoryError("Directory Service origin is not canonical") - return page + return page.model_copy(update={"issues": issues}) + + +def _read_service(entry: object) -> dict[str, object]: + """Validates one Directory record, or raises describing why it cannot be used. + + A Directory result echoes Service Document members, so those members are held to the Service + Document's own rules rather than merely to their JSON types -- a Directory that published + `"language": "not a tag"` would otherwise hand a caller a tag no language matcher can read. + """ + if not isinstance(entry, dict): + raise DirectoryError("Directory Service result is not an object") + item = dict(entry) + origin = item.get("service_origin") + if not isinstance(origin, str) or _origin_of(origin, "Directory Service origin") != origin: + raise DirectoryError("Directory Service origin is not a canonical HTTPS origin") + _require_public_origin(origin) + indexed_at = item.get("indexed_at") + # A value `datetime.fromisoformat` happens to accept is not an RFC 3339 timestamp, and a caller + # that compares or slices `indexed_at` needs the one shape. + if not isinstance(indexed_at, str) or not _RFC_3339.match(indexed_at): + raise DirectoryError("indexed_at must be an RFC 3339 date-time") + document = _validate_as_service_document(item) + # `protocols` is reinstated from the validated document only when something survived agent + # filtering, so a block naming nothing this ODP version knows does not pass straight through. + for member in (*_UNVERIFIED_MEMBERS, "protocols"): + item.pop(member, None) + if document.protocols is not None: + item["protocols"] = document.protocols.model_dump(mode="json", exclude_defaults=True) + return item -def _normalize_service_protocols(item: dict[str, object]) -> None: - candidate = { - "description": "Directory protocol validation", +def _validate_as_service_document(item: dict[str, object]) -> ServiceDocument: + """Holds the Service Document members a record echoes to the Service Document's own rules.""" + candidate: dict[str, object] = { "http": {"endpoint_base": "/"}, - "language": "en", - "localizations": ["en"], - "name": "Directory Service", "odp_version": "1.0", - "operations": [ - {"authentication": "not-required", "name": "get-offering"}, - {"authentication": "not-required", "name": "list-offerings"}, - ], - "protocols": item["protocols"], + "operations": item.get("operations"), } - document = parse_agent_service_document(json.dumps(candidate, separators=(",", ":"))) - if document.protocols is None: - item.pop("protocols") - else: - item["protocols"] = document.protocols.model_dump(mode="json", exclude_defaults=True) + for member in ("description", "keywords", "language", "localizations", "name", "protocols"): + if member in item: + candidate[member] = item[member] + for member in ("documentation_url", "status_url", "support_url", "website_url"): + value = item.get(member) + if value not in (None, ""): + candidate[member] = value + return parse_agent_service_document(json.dumps(candidate, separators=(",", ":"))) + + +def _require_public_origin(origin: str) -> None: + """A public Directory has no business pointing an Agent at a loopback or private host. + + The default transport refuses these too, but that guarantee should not depend on which + transport the consumer installed, nor on whether they enabled local development for a Service + of their own. + """ + host = urlsplit(origin).hostname or "" + try: + address = ip_address(host) + except ValueError: + return + if not is_public(address): + raise DirectoryError("Directory Service origin must not be a private or loopback host") + + +def _origin_of(value: str, subject: str) -> str: + """Derives a Service Origin, reporting a reference that is not one as a Directory failure.""" + try: + return derive_service_origin(value) + except ReferenceError as error: + raise DirectoryError(f"{subject} is not a usable origin: {error}") from error def _validate_search_request(request: SearchRequest) -> None: + # 0 means "the Directory decides"; anything else is a count, and a count below one asks for + # nothing while still being sent on the wire. if request.limit < 0 or request.limit > 100: raise DirectoryError("limit must be from 1 through 100") if request.query.strip() != request.query or len(request.query) > 512: @@ -277,14 +359,17 @@ def _validate_search_request(request: SearchRequest) -> None: def _consume_response(response: HttpResponse) -> HttpResponse: - if len(response.body) > _MAXIMUM_RESPONSE_BYTES: - raise DirectoryError("Directory response exceeds 524288 bytes") if not 200 <= response.status < 300: + # The status is what describes the failure. Reading the size first turned a refused request + # into a size complaint, and passing the whole body on made the error message as large as + # the response the Directory sent. raise DirectoryRequestError( response.status, - response.body.decode(errors="replace"), + response.body.decode(errors="replace")[:_MAXIMUM_ERROR_CHARACTERS], response.headers, ) + if len(response.body) > _MAXIMUM_RESPONSE_BYTES: + raise DirectoryError("Directory response exceeds 524288 bytes") content_type = response.headers.get("content-type", "").split(";", 1)[0].strip().lower() if content_type != "application/json": raise DirectoryError("Directory response must use application/json") diff --git a/src/offering_protocol/directory/models.py b/src/offering_protocol/directory/models.py index 37e42c3..31a5d82 100644 --- a/src/offering_protocol/directory/models.py +++ b/src/offering_protocol/directory/models.py @@ -142,9 +142,23 @@ class Facets(OdpModel): trust: list[Facet[TrustProtocol]] = Field(default_factory=list) +class ServiceIssue(OdpModel): + """A record the Directory published that this client would not hand back. + + ROLE-03: a Directory result is discovery metadata rather than authoritative Service data, so + one unusable record is a note about that record, not a reason to withhold every other Service + on the page. + """ + + index: int + message: str + + class SearchPage(OdpModel): facets: Facets | None = None + #: The records this client was able to read. Withheld records appear in `issues`. items: list[DirectoryService] + issues: list[ServiceIssue] = Field(default_factory=list) next: str = "" diff --git a/src/offering_protocol/directory/transport.py b/src/offering_protocol/directory/transport.py index 02a0a93..e13d01d 100644 --- a/src/offering_protocol/directory/transport.py +++ b/src/offering_protocol/directory/transport.py @@ -5,12 +5,14 @@ import asyncio import socket from dataclasses import dataclass, field -from ipaddress import IPv4Address, IPv6Address, ip_address +from ipaddress import ip_address from typing import Protocol from urllib.parse import urlsplit, urlunsplit import httpx +from offering_protocol.directory.addresses import IPAddress, is_public + @dataclass(frozen=True, slots=True) class HttpRequest: @@ -93,9 +95,6 @@ async def aclose(self) -> None: self._clients.clear() -IPAddress = IPv4Address | IPv6Address - - async def _pinned_target(url: str, allow_local_network: bool) -> tuple[str, str, str]: parsed = urlsplit(url) if parsed.hostname is None or parsed.username is not None or parsed.password is not None: @@ -113,7 +112,7 @@ async def _pinned_target(url: str, allow_local_network: bool) -> tuple[str, str, if local_hostname and allow_local_network: if any(not address.is_loopback for address in addresses): raise ValueError("ODP local-development host resolved outside the loopback network") - elif any(not address.is_global for address in addresses): + elif any(not is_public(address) for address in addresses): raise ValueError("ODP request host resolved to a non-public address") address = addresses[0] pinned_host = f"[{address}]" if address.version == 6 else str(address) diff --git a/src/offering_protocol/service/service.py b/src/offering_protocol/service/service.py index b2f1b91..ec6dfd9 100644 --- a/src/offering_protocol/service/service.py +++ b/src/offering_protocol/service/service.py @@ -2,6 +2,8 @@ from __future__ import annotations +import base64 +import hashlib import json from collections.abc import Callable from dataclasses import dataclass, field @@ -45,7 +47,26 @@ MEDIA_TYPE = "application/odp+json" PROBLEM_MEDIA_TYPE = "application/problem+json" _MAXIMUM_REQUEST_BYTES = 65_536 +_MAXIMUM_DOCUMENT_BYTES = 65_536 _MAXIMUM_RESOURCE_BYTES = 524_288 +#: Paths under the endpoint base that name an operation rather than a resource. +#: +#: Without this, `GET /offerings/search` reads as a request for an Offering called "search" and +#: answers 404, which tells an Agent the operation does not exist rather than that it uses POST. +_RESERVED_PATHS = {"/offerings/search": "POST", "/collections/search": "POST"} +#: RFC 9457: a title is a short summary of the problem *type* and does not change from one +#: occurrence to the next. The varying part of a failure belongs in `detail`. +_PROBLEM_TITLES = { + "CONTINUATION_UNAVAILABLE": "Continuation unavailable", + "INTERNAL_ERROR": "Internal error", + "INVALID_REQUEST": "Invalid request", + "METHOD_NOT_ALLOWED": "Method not allowed", + "NOT_ACCEPTABLE": "Not acceptable", + "NOT_FOUND": "Not found", + "PRECONDITION_FAILED": "Precondition failed", + "REQUEST_TOO_LARGE": "Request too large", + "UNSUPPORTED_MEDIA_TYPE": "Unsupported media type", +} Validated = TypeVar("Validated") @@ -69,11 +90,23 @@ class Response: class CatalogRequest: accept_language: str | None = None cursor: str | None = None + #: The language this Service selected for the response, by RFC 4647 Lookup over its + #: localizations. A Catalog answers in this tag rather than parsing `accept_language` itself. + language: str = "" limit: int = 0 path: str = "" representation: Representation = Representation.TERSE +@dataclass(frozen=True, slots=True) +class _Exchange: + """What building a response needs to know about the request it is answering.""" + + headers: dict[str, str] + language: str + method: str + + class ServiceError(RuntimeError): """Base error for Service integration failures.""" @@ -83,10 +116,13 @@ class CatalogError(ServiceError): class RequestError(ServiceError): - def __init__(self, status: int, code: str, message: str) -> None: + def __init__( + self, status: int, code: str, message: str, headers: dict[str, str] | None = None + ) -> None: super().__init__(message) self.status = status self.code = code + self.headers = headers or {} class Catalog(Protocol): @@ -226,7 +262,7 @@ async def handle(self, request: Request) -> Response: try: return await self._handle(request) except RequestError as error: - return _problem(error.status, error.code, str(error)) + return _problem(error.status, error.code, str(error), error.headers) except OdpValidationError as error: detail = "; ".join(f"{issue.path or '/'}: {issue.message}" for issue in error.issues) return _problem(400, "INVALID_REQUEST", detail) @@ -235,108 +271,126 @@ async def handle(self, request: Request) -> Response: async def _handle(self, request: Request) -> Response: headers = {name.lower(): value for name, value in request.headers.items()} - if not _accepts_odp(headers.get("accept")): - return _problem(406, "NOT_ACCEPTABLE", f"Accept must allow {MEDIA_TYPE}") + _require_accept(headers.get("accept")) + # RFC 9110 9.3.2: HEAD is GET without the body, so every resource answering GET answers + # HEAD. Routing on the effective method keeps the two from drifting apart. method = request.method.upper() + effective = "GET" if method == "HEAD" else method + language = _select_language( + headers.get("accept-language"), + self._document.language, + list(self._document.localizations), + ) + exchange = _Exchange(headers=headers, language=language, method=method) if request.path == "/.well-known/odp": - if method != "GET": - return _problem(405, "METHOD_NOT_ALLOWED", "The Service Document requires GET") - return _json_response(200, self._document, _MAXIMUM_REQUEST_BYTES) + _require_method(effective, ("GET",)) + return _json_response(self._document, _MAXIMUM_DOCUMENT_BYTES, exchange) if not request.path.startswith(self._endpoint_base): - return _problem(404, "NOT_FOUND", "ODP resource not found") + raise RequestError(404, "NOT_FOUND", "ODP resource not found") path = request.path[len(self._endpoint_base) :] - operation = _path_operation(method, path) + if path in _RESERVED_PATHS: + _require_method(effective, (_RESERVED_PATHS[path],)) + operation = _path_operation(effective, path) if operation is not None and operation not in { item.name for item in self._document.operations }: - return _problem(404, "NOT_FOUND", "ODP operation is not supported") - catalog_request = _catalog_request(request, headers) - if (method, path) == ("GET", "/offerings"): + raise RequestError(404, "NOT_FOUND", "ODP operation is not supported") + catalog_request = _catalog_request(request, headers, language) + if (effective, path) == ("GET", "/offerings"): offering_page = await self._catalog.list_offerings(catalog_request) return _json_response( - 200, _offering_page(offering_page, catalog_request.representation), _MAXIMUM_RESOURCE_BYTES, + exchange, ) - if (method, path) == ("POST", "/offerings/search"): + if (effective, path) == ("POST", "/offerings/search"): query = parse_offering_search_request(_search_body(request, headers)) offering_page = await self._catalog.search_offerings(query, catalog_request) return _json_response( - 200, _offering_page(offering_page, catalog_request.representation), _MAXIMUM_RESOURCE_BYTES, + exchange, ) - if (method, path) == ("GET", "/collections"): + if (effective, path) == ("GET", "/collections"): collection_page = await self._catalog.list_collections(catalog_request) return _json_response( - 200, _collection_page(collection_page, catalog_request.representation), _MAXIMUM_RESOURCE_BYTES, + exchange, ) - if (method, path) == ("POST", "/collections/search"): + if (effective, path) == ("POST", "/collections/search"): collection_query = parse_collection_search_request(_search_body(request, headers)) collection_page = await self._catalog.search_collections( collection_query, catalog_request ) return _json_response( - 200, _collection_page(collection_page, catalog_request.representation), _MAXIMUM_RESOURCE_BYTES, + exchange, ) - if method == "GET": - return await self._get_path(path, catalog_request) - return _problem(405, "METHOD_NOT_ALLOWED", "ODP operation uses a fixed HTTP method") + if effective == "GET": + return await self._get_path(path, catalog_request, exchange) + raise RequestError( + 405, "METHOD_NOT_ALLOWED", "ODP operation uses a fixed HTTP method", _allow(("GET",)) + ) - async def _get_path(self, path: str, request: CatalogRequest) -> Response: + async def _get_path(self, path: str, request: CatalogRequest, exchange: _Exchange) -> Response: if path.startswith("/offerings/"): - identifier = path.removeprefix("/offerings/") - if not is_local_resource_identifier(identifier): - return _problem(400, "INVALID_REQUEST", "Offering identifier is invalid") + identifier = _require_identifier(path.removeprefix("/offerings/"), "Offering") offering = await self._catalog.get_offering(identifier, request) if offering is None: - return _problem(404, "NOT_FOUND", "Offering not found") + raise RequestError(404, "NOT_FOUND", "Offering not found") if offering.id != identifier: raise ServiceError("Offering identifier does not match request path") return _json_response( - 200, - _offering(offering, request.representation), - _MAXIMUM_RESOURCE_BYTES, + _offering(offering, request.representation), _MAXIMUM_RESOURCE_BYTES, exchange ) if path.startswith("/collections/"): value = path.removeprefix("/collections/") if value.endswith("/offerings"): - identifier = value.removesuffix("/offerings") + # SVC-66 substitutes an identifier verbatim, so a path segment that is not a Local + # Resource Identifier names no resource this Service could ever hold -- and it is + # the Catalog that would otherwise have to decide what to do with it. + identifier = _require_identifier(value.removesuffix("/offerings"), "Collection") page = await self._catalog.list_collection_offerings(identifier, request) return _json_response( - 200, _offering_page(page, request.representation), _MAXIMUM_RESOURCE_BYTES, + exchange, ) - collection = await self._catalog.get_collection(value, request) + identifier = _require_identifier(value, "Collection") + collection = await self._catalog.get_collection(identifier, request) if collection is None: - return _problem(404, "NOT_FOUND", "Collection not found") - if collection.id != value: + raise RequestError(404, "NOT_FOUND", "Collection not found") + if collection.id != identifier: raise ServiceError("Collection identifier does not match request path") return _json_response( - 200, - _collection(collection, request.representation), - _MAXIMUM_RESOURCE_BYTES, + _collection(collection, request.representation), _MAXIMUM_RESOURCE_BYTES, exchange ) - return _problem(404, "NOT_FOUND", "ODP resource not found") - - -def _catalog_request(request: Request, headers: dict[str, str]) -> CatalogRequest: - parameters = dict(parse_qsl(request.query, keep_blank_values=True)) + raise RequestError(404, "NOT_FOUND", "ODP resource not found") + + +def _catalog_request(request: Request, headers: dict[str, str], language: str) -> CatalogRequest: + parameters = parse_qsl(request.query, keep_blank_values=True) + # SVC-73: a repeated `representation` is rejected rather than resolved. Collapsing repeats into + # a dict silently honoured whichever copy came last, so `representation=terse&representation= + # full` served a Full Representation to a request that also asked for a Terse one. + values: dict[str, str] = {} + for name, value in parameters: + if name in {"cursor", "limit", "representation"} and name in values: + raise RequestError(400, "INVALID_REQUEST", f"{name} must not be repeated") + values[name] = value try: - representation = Representation(parameters.get("representation", "terse")) - limit = int(parameters.get("limit", "0")) + representation = Representation(values.get("representation", "terse")) + limit = int(values.get("limit", "0")) except ValueError as error: raise RequestError(400, "INVALID_REQUEST", "query parameter is invalid") from error if not 0 <= limit <= 100: raise RequestError(400, "INVALID_REQUEST", "limit exceeds 100") return CatalogRequest( accept_language=headers.get("accept-language"), - cursor=parameters.get("cursor"), + cursor=values.get("cursor"), + language=language, limit=limit, path=request.path, representation=representation, @@ -370,11 +424,56 @@ def _path_operation(method: str, path: str) -> Operation | None: return None -def _json_response(status: int, value: object, maximum_bytes: int) -> Response: +def _json_response(value: object, maximum_bytes: int, exchange: _Exchange) -> Response: + """Serializes once, so the body can be measured, given a validator, and conditionally answered. + + SVC-60 and SVC-61 apply to every representation this Service serves, not only to the localized + ones: `Vary` is what stops a shared cache handing an English body to a French request, and the + entity tag is what lets an Agent revalidate instead of re-transferring (PAG-31). + """ body = _encode(value) if len(body) > maximum_bytes: raise ServiceError("response body is too large") - return Response(status, {"content-type": MEDIA_TYPE}, body) + etag = _entity_tag(exchange.language, body) + headers = { + "content-language": exchange.language, + "content-type": MEDIA_TYPE, + "etag": etag, + "vary": "Accept, Accept-Language", + } + if not _matches_entity_tag(exchange.headers.get("if-none-match"), etag): + return Response(200, headers, b"" if exchange.method == "HEAD" else body) + # RFC 9110 13.1.2: a matched `If-None-Match` is 304 for GET and HEAD, 412 for anything else. + if exchange.method in {"GET", "HEAD"}: + return Response(304, {name: headers[name] for name in ("etag", "vary")}, b"") + raise RequestError( + 412, "PRECONDITION_FAILED", "If-None-Match matched the current representation" + ) + + +def _require_method(method: str, allowed: tuple[str, ...]) -> None: + if method not in allowed: + raise RequestError( + 405, + "METHOD_NOT_ALLOWED", + f"ODP operation requires {' or '.join(allowed)}", + _allow(allowed), + ) + + +def _allow(allowed: tuple[str, ...]) -> dict[str, str]: + """RFC 9110 15.5.6 requires a 405 to name every method the resource supports. + + Every ODP operation that answers GET also answers HEAD, so HEAD travels with it. + """ + methods = [*allowed, "HEAD"] if "GET" in allowed else list(allowed) + return {"allow": ", ".join(methods)} + + +def _require_identifier(value: str, label: str) -> str: + if not is_local_resource_identifier(value): + raise RequestError(400, "INVALID_REQUEST", f"{label} identifier is invalid") + return value def _offering(value: Offering, representation: Representation) -> Offering: @@ -426,15 +525,17 @@ def _validated(parser: Callable[[bytes | str], Validated], value: object) -> Val raise ServiceError(f"Catalog returned an invalid ODP response: {error}") from error -def _problem(status: int, code: str, detail: str) -> Response: +def _problem( + status: int, code: str, detail: str, headers: dict[str, str] | None = None +) -> Response: value = ProblemDetails( code=code, detail=detail, status=status, - title=detail, + title=_PROBLEM_TITLES.get(code, code.replace("_", " ").capitalize()), type=f"https://offeringprotocol.org/problems/{code.lower().replace('_', '-')}", ) - return Response(status, {"content-type": PROBLEM_MEDIA_TYPE}, _encode(value)) + return Response(status, {"content-type": PROBLEM_MEDIA_TYPE, **(headers or {})}, _encode(value)) def _encode(value: object) -> bytes: @@ -444,9 +545,114 @@ def _encode(value: object) -> bytes: return json.dumps(value, separators=(",", ":")).encode() -def _accepts_odp(value: str | None) -> bool: +def _require_accept(value: str | None) -> None: + """MED-03/MED-04: an `Accept` that excludes ODP is a refusal, an absent one is not. + + `application/*` is a wildcard media range that covers this media type, and RFC 9110 12.4.2 + makes `q=0` an explicit statement that a range is *not* acceptable -- so an `Accept` naming ODP + with zero weight excludes it just as surely as one that never mentions it. + """ if value is None: - return True - return any( - item.split(";", 1)[0].strip().lower() in {"*/*", MEDIA_TYPE} for item in value.split(",") + return + for entry in value.split(","): + media_type = entry.split(";", 1)[0].strip().lower() + if media_type not in {"*/*", "application/*", MEDIA_TYPE}: + continue + if _quality_of(entry) > 0: + return + raise RequestError(406, "NOT_ACCEPTABLE", f"Accept must allow {MEDIA_TYPE}") + + +def _quality_of(entry: str) -> float: + """The weight of one `Accept` or `Accept-Language` entry, or 0 when it carries no usable one.""" + for parameter in entry.split(";")[1:]: + name, separator, value = parameter.partition("=") + if name.strip().lower() != "q" or not separator: + continue + try: + quality = float(value.strip()) + except ValueError: + return 0.0 + return quality if 0 <= quality <= 1 else 0.0 + return 1.0 + + +def _select_language(value: str | None, fallback: str, localizations: list[str]) -> str: + """SVC-58/59: RFC 4647 Lookup over the localizations this Service advertises. + + Returning the default when nothing matches is the point of SVC-59: a Service never refuses a + request over language, so the worst outcome of an unmatched range is the representation the + Service would have served anyway. + """ + if value is None: + return fallback + entries = [ + (entry.split(";", 1)[0].strip().lower(), _quality_of(entry)) for entry in value.split(",") + ] + entries = [(range_, quality) for range_, quality in entries if range_] + # RFC 9110 12.4.2: a range weighted zero is unacceptable. It carves tags out of the `*` + # residual below rather than competing for a match of its own. + refused = [range_ for range_, quality in entries if quality == 0 and range_ != "*"] + wanted = sorted( + ((range_, quality) for range_, quality in entries if quality > 0 and range_ != "*"), + key=lambda item: item[1], + reverse=True, ) + for range_, _ in wanted: + found = _lookup(range_, localizations) + if found is not None: + return found + # RFC 9110 12.5.4: `*` matches every tag no other range in the field matched, so it is the + # residual and can never outrank a range the caller named. + if not any(range_ == "*" and quality > 0 for range_, quality in entries): + return fallback + for tag in (fallback, *localizations): + if not any(_covers(range_, tag.lower()) for range_ in refused): + return tag + return fallback + + +def _lookup(range_: str, localizations: list[str]) -> str | None: + """RFC 4647 Lookup: truncate the range at its subtag boundaries until a tag matches exactly.""" + available = [(tag.lower(), tag) for tag in localizations] + candidate = range_ + while True: + for lowered, tag in available: + if lowered == candidate: + return tag + cut = candidate.rfind("-") + if cut < 0: + return None + candidate = candidate[:cut] + # A single-character subtag is an extension or private-use singleton; drop it with its + # parent rather than leaving a range that names a singleton and nothing else. + singleton = candidate.rfind("-") + if singleton >= 0 and len(candidate) - singleton == 2: + candidate = candidate[:singleton] + + +def _covers(range_: str, tag: str) -> bool: + return tag == range_ or tag.startswith(f"{range_}-") + + +def _entity_tag(language: str, body: bytes) -> str: + """SVC-61: a strong validator over the negotiated language and the exact bytes served. + + Hashing the body alone would give two language variants of an unlocalized body one validator, + which is the one thing an entity tag must not do. + """ + digest = hashlib.sha256(language.encode() + b"\x00" + body).digest() + return '"' + base64.urlsafe_b64encode(digest).decode().rstrip("=")[:27] + '"' + + +def _matches_entity_tag(value: str | None, etag: str) -> bool: + """RFC 9110 compares `If-None-Match` weakly, so `W/"x"` matches the strong `"x"` served here.""" + if value is None: + return False + for candidate in value.split(","): + candidate = candidate.strip() + if candidate.startswith("W/"): + candidate = candidate[2:] + if candidate in {"*", etag}: + return True + return False diff --git a/src/offering_protocol/service/static_catalog.py b/src/offering_protocol/service/static_catalog.py index 66f78c9..c3862f1 100644 --- a/src/offering_protocol/service/static_catalog.py +++ b/src/offering_protocol/service/static_catalog.py @@ -205,8 +205,12 @@ def _represent_collection(value: Collection, request: CatalogRequest, embedded: def _encode_cursor( request: CatalogRequest, limit: int, offset: int, continuation_key: bytes ) -> str: + # The cursor carries everything that decided what this page contained, so a continuation + # cannot quietly change variant part-way through a sequence -- including the language, because + # a page served in French is a different page from the same offsets served in English. value = { "expires": int(time.time()) + _CONTINUATION_LIFETIME_SECONDS, + "language": request.language, "limit": limit, "offset": offset, "path": request.path, @@ -243,6 +247,7 @@ def _decode_cursor(request: CatalogRequest, limit: int, continuation_key: bytes) if ( not isinstance(value, dict) or value.get("expires", 0) < int(time.time()) + or value.get("language") != request.language or value.get("limit") != limit or value.get("path") != request.path or value.get("representation") != request.representation.value diff --git a/tests/test_agent.py b/tests/test_agent.py index 35b9f28..1dc3f11 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -291,7 +291,7 @@ async def test_builds_agent_friendly_offering_details_without_invoking_action() DIRECTORY_PAGE = """{ "items":[{ "description":"Plants","indexed_at":"2026-08-25T00:00:00Z","language":"en", - "localizations":["en"],"name":"One","operations":[], + "localizations":["en"],"name":"One","operations":[{"authentication":"not-required","name":"get-offering"},{"authentication":"not-required","name":"list-offerings"}], "service_origin":"https://one.example" }] }""" @@ -595,7 +595,7 @@ async def test_capability_limits_duplicates_and_pagination_edges( CapabilityScope.SERVICE, SearchCapabilities(filters=FilterCapabilitySource(inline=[definition])), ) - assert "exceed 1024" in result.issues[-1].message + assert "Effective filters exceed their limit" in result.issues[-1].message sort = SortDefinition( description="Price", @@ -620,8 +620,11 @@ async def test_capability_limits_duplicates_and_pagination_edges( CapabilityScope.COLLECTION, SearchCapabilities(sorts=SortCapabilitySource(inline=[sort, sort])), ) - assert not target and not scopes - assert "Duplicate sorts" in result.issues[-1].message + # FLT-55: `[sort, sort]` repeats an identifier within one source, so that source is + # discarded whole and the sort merged from the earlier source is left exactly as it was. + assert target == {"price": sort} + assert scopes == {"price": CapabilityScope.SERVICE} + assert "within one source" in result.issues[-1].message with pytest.raises(AgentError): _resolve_reference("data:text/plain,x", client.service_origin) @@ -834,7 +837,7 @@ async def test_remaining_capability_and_cache_branches(monkeypatch: pytest.Monke CapabilityScope.SERVICE, SearchCapabilities(sorts=SortCapabilitySource(inline=[sort])), ) - assert "exceed 128" in result.issues[-1].message + assert "Effective sorts exceed their limit" in result.issues[-1].message monkeypatch.setattr("offering_protocol.agent.capabilities._MAXIMUM_CAPABILITY_PAGES", 1) filter_limit = ServiceClient( diff --git a/tests/test_agent_conformance.py b/tests/test_agent_conformance.py new file mode 100644 index 0000000..af7f976 --- /dev/null +++ b/tests/test_agent_conformance.py @@ -0,0 +1,536 @@ +"""Agent conformance: the rules an ODP Agent applies to what somebody else wrote. + +Every document reaching the Agent -- a Service Document, a capability page, a Problem Details +response, an OpenAPI document -- is written by the Service, so each test here states one rule the +Agent enforces on that input and shows the Agent refusing the input that breaks it. +""" + +from __future__ import annotations + +import json +from datetime import timedelta + +import pytest + +from helpers import SERVICE_DOCUMENT, QueueTransport, response +from offering_protocol.agent import AgentError, ServiceClient, ServiceRequestError +from offering_protocol.agent.cache import utc_now +from offering_protocol.agent.capabilities import ( + CapabilityKind, + CapabilityScope, + SearchCapabilityCatalog, + _add_filters, + _add_sorts, + _load_filters, + _load_sorts, +) +from offering_protocol.agent.client import ( + _MAXIMUM_DEPTH, + _MAXIMUM_DOCUMENT_DEPTH, + _consume, + _expiration, + _nesting_depth, +) +from offering_protocol.core import ( + CapabilityLink, + FilterCapabilitySource, + FilterDefinition, + FilterOperator, + FilterType, + MissingPlacement, + SearchCapabilities, + SortCapabilitySource, + SortDefinition, + SortDirection, + SortKey, +) +from offering_protocol.directory.transport import HttpResponse + +ORIGIN = "https://plants.example" + + +def _client(*replies: HttpResponse) -> ServiceClient: + return ServiceClient(ORIGIN, transport=QueueTransport(*replies)) + + +def _filter(identifier: str) -> FilterDefinition: + return FilterDefinition( + description="How heavy the plant is.", + id=identifier, + operators=[FilterOperator.EQUAL], + title="Weight", + type=FilterType.NUMBER, + ) + + +def _sort(identifier: str, filter_id: str = "weight") -> SortDefinition: + return SortDefinition( + description="Orders plants by weight.", + id=identifier, + keys=[ + SortKey( + direction=SortDirection.ASCENDING, + filter_id=filter_id, + missing=MissingPlacement.LAST, + ) + ], + title="Lightest first", + ) + + +def _filter_page(identifiers: list[str], next_reference: str = "") -> str: + page: dict[str, object] = { + "odp_version": "1.0", + "items": [ + json.loads(_filter(identifier).model_dump_json(by_alias=True, exclude_unset=True)) + for identifier in identifiers + ], + } + if next_reference: + page["next"] = next_reference + return json.dumps(page) + + +async def _merge_filters( + result: SearchCapabilityCatalog, + scope: CapabilityScope, + source: FilterCapabilitySource, + client: ServiceClient | None = None, +) -> None: + await _add_filters( + client, # type: ignore[arg-type] + result, + scope, + SearchCapabilities(filters=source), + ) + + +# -- linked capability sources stay on the Service origin ------------------------------------------ + + +@pytest.mark.asyncio +async def test_refuses_a_cross_origin_linked_source() -> None: + """FLT-52: `linked.href` is a same-origin Resource Reference. + + The Service Document that names it is written by the Service, so a cross-origin `href` would + let a Service point this Agent's ODP requests at a host of its choosing. + """ + client = _client(response(_filter_page([]))) + with pytest.raises(AgentError, match="Service origin"): + await _load_filters(client, "https://elsewhere.example/filters") + assert not client._transport.requests # type: ignore[attr-defined] + + +@pytest.mark.asyncio +async def test_refuses_a_cross_origin_continuation() -> None: + """FLT-53: a linked page's `next` obeys the common continuation contract. + + That contract keeps a continuation on the Service origin, so a page can no more redirect this + Agent off the Service than the advertisement that named the source could. + """ + client = _client(response(_filter_page(["weight"], "https://elsewhere.example/page-2"))) + with pytest.raises(AgentError, match="Service origin"): + await _load_filters(client, "/odp/filters") + assert [request.url for request in client._transport.requests] == [ # type: ignore[attr-defined] + f"{ORIGIN}/odp/filters" + ] + + +@pytest.mark.asyncio +async def test_follows_a_same_origin_source_and_its_continuations() -> None: + client = _client( + response(_filter_page(["weight"], "/odp/filters?page=2")), + response(_filter_page(["height"])), + ) + values = await _load_filters(client, "/odp/filters") + + assert [value.id for value in values] == ["weight", "height"] + + +@pytest.mark.asyncio +async def test_refuses_a_reference_that_is_not_an_odp_reference() -> None: + for reference in ("data:text/plain,x", "//elsewhere.example/filters", "filters"): + with pytest.raises(AgentError): + await _load_filters(_client(), reference) + + +# -- a source is atomic --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_a_repeat_within_one_source_discards_that_source() -> None: + """FLT-55: a source enforces its own uniqueness before any of it is exposed. + + Two definitions under one identifier in one source give the Agent no basis for choosing between + them, so the source is unusable rather than partly usable. + """ + result = SearchCapabilityCatalog() + await _merge_filters( + result, + CapabilityScope.SERVICE, + FilterCapabilitySource(inline=[_filter("weight"), _filter("weight"), _filter("height")]), + ) + + assert not result.filters + assert len(result.issues) == 1 + assert "within one source" in result.issues[0].message + assert result.issues[0].kind is CapabilityKind.FILTERS + + +@pytest.mark.asyncio +async def test_a_repeat_within_one_source_leaves_earlier_sources_alone() -> None: + result = SearchCapabilityCatalog() + await _merge_filters( + result, CapabilityScope.SERVICE, FilterCapabilitySource(inline=[_filter("weight")]) + ) + await _merge_filters( + result, + CapabilityScope.COLLECTION, + FilterCapabilitySource(inline=[_filter("height"), _filter("height")]), + ) + + assert sorted(result.filters) == ["weight"] + + +@pytest.mark.asyncio +async def test_an_identifier_two_sources_publish_is_quarantined() -> None: + """One identifier published by two effective sources is quarantined -- neither copy wins. + + That removes the identifier and nothing else, which is what tells this rule apart from the + within-source repeat above. + """ + result = SearchCapabilityCatalog() + await _merge_filters( + result, + CapabilityScope.SERVICE, + FilterCapabilitySource(inline=[_filter("weight"), _filter("service-only")]), + ) + await _merge_filters( + result, + CapabilityScope.COLLECTION, + FilterCapabilitySource(inline=[_filter("weight"), _filter("collection-only")]), + ) + + assert sorted(result.filters) == ["collection-only", "service-only"] + assert result.issues[-1].message == "Duplicate filters: weight" + assert result.issues[-1].scope is CapabilityScope.COLLECTION + + +# -- effective-catalog bounds --------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_a_source_that_overflows_the_bound_leaves_earlier_sources_intact() -> None: + """FLT-62: exceeding the bound invalidates the source that caused it, and only that source.""" + result = SearchCapabilityCatalog() + await _merge_filters( + result, CapabilityScope.SERVICE, FilterCapabilitySource(inline=[_filter("weight")]) + ) + # 1025 new definitions beside the one already merged cannot fit, even after the identifier + # both sources publish is quarantined. + overflowing = [_filter(f"f{index}") for index in range(1025)] + [_filter("weight")] + await _merge_filters( + result, CapabilityScope.COLLECTION, FilterCapabilitySource(inline=overflowing) + ) + + assert sorted(result.filters) == ["weight"] + assert "Effective filters exceed their limit" in result.issues[-1].message + + +@pytest.mark.asyncio +async def test_a_source_that_exactly_fills_the_bound_is_accepted() -> None: + result = SearchCapabilityCatalog() + await _merge_filters( + result, CapabilityScope.SERVICE, FilterCapabilitySource(inline=[_filter("weight")]) + ) + exact = [_filter(f"f{index}") for index in range(1023)] + await _merge_filters(result, CapabilityScope.COLLECTION, FilterCapabilitySource(inline=exact)) + + assert len(result.filters) == 1024 + assert not result.issues + + +@pytest.mark.asyncio +async def test_paging_stops_at_the_page_that_overflows_the_bound() -> None: + """FLT-58: a source that cannot fit costs one page rather than sixteen.""" + page = _filter_page([f"f{index}" for index in range(100)], "/odp/filters?page=next") + client = _client(*[response(page) for _ in range(16)]) + + with pytest.raises(AgentError, match="exceed their limit"): + await _load_filters(client, "/odp/filters", budget=50) + + assert len(client._transport.requests) == 1 # type: ignore[attr-defined] + + +@pytest.mark.asyncio +async def test_a_source_whose_last_page_still_offers_another_is_discarded() -> None: + """FLT-59: page 16 carrying `next` means page 17 is never retrieved.""" + pages = [ + response(_filter_page([f"p{index}"], f"/odp/filters?page={index + 1}")) + for index in range(16) + ] + client = _client(*pages) + + with pytest.raises(AgentError, match="16 pages"): + await _load_filters(client, "/odp/filters") + + assert len(client._transport.requests) == 16 # type: ignore[attr-defined] + + +@pytest.mark.asyncio +async def test_sorts_follow_the_same_source_rules() -> None: + result = SearchCapabilityCatalog() + target: dict[str, SortDefinition] = {} + scopes: dict[str, CapabilityScope] = {} + await _add_sorts( + _client(), + result, + target, + scopes, + CapabilityScope.SERVICE, + SearchCapabilities(sorts=SortCapabilitySource(inline=[_sort("light"), _sort("light")])), + ) + + assert not target and not scopes + assert "within one source" in result.issues[-1].message + + +@pytest.mark.asyncio +async def test_a_linked_sort_source_stays_on_the_service_origin() -> None: + result = SearchCapabilityCatalog() + await _add_sorts( + _client(), + result, + {}, + {}, + CapabilityScope.SERVICE, + SearchCapabilities( + sorts=SortCapabilitySource(linked=CapabilityLink(href="https://elsewhere.example/s")) + ), + ) + + assert "Service origin" in result.issues[-1].message + assert result.issues[-1].kind is CapabilityKind.SORTS + + +@pytest.mark.asyncio +async def test_reads_a_linked_source_into_the_catalog() -> None: + result = SearchCapabilityCatalog() + client = _client(response(_filter_page(["weight", "height"]))) + await _merge_filters( + result, + CapabilityScope.SERVICE, + FilterCapabilitySource(linked=CapabilityLink(href="/odp/filters")), + client, + ) + + assert sorted(result.filters) == ["height", "weight"] + assert not result.issues + + +@pytest.mark.asyncio +async def test_loading_sorts_returns_what_a_complete_source_advertised() -> None: + sorts = json.dumps( + { + "odp_version": "1.0", + "items": [ + json.loads(_sort("light").model_dump_json(by_alias=True, exclude_unset=True)) + ], + } + ) + values = await _load_sorts(_client(response(sorts)), "/odp/sorts") + + assert [value.id for value in values] == ["light"] + + +# -- nesting depth -------------------------------------------------------------------------------- + + +def test_measures_nesting_from_the_top_level_value() -> None: + """ERR-18: a scalar is a value a container holds, not a level of its own.""" + assert _nesting_depth(3) == 0 + assert _nesting_depth({}) == 1 + assert _nesting_depth({"a": 1}) == 1 + assert _nesting_depth({"a": {"b": 1}}) == 2 + assert _nesting_depth([[{"a": 1}]]) == 3 + assert _nesting_depth({"shallow": 1, "deep": {"a": {"b": 1}}}) == 3 + + +def _nested(levels: int) -> bytes: + value: dict[str, object] = {} + cursor = value + for _ in range(levels - 1): + child: dict[str, object] = {} + cursor["a"] = child + cursor = child + return json.dumps(value).encode() + + +def _odp(body: bytes, status: int = 200) -> HttpResponse: + return HttpResponse(status, {"content-type": "application/odp+json"}, body) + + +def test_refuses_a_response_nested_deeper_than_odp_allows() -> None: + """ERR-21: every ODP document except the Service Document nests no deeper than 16.""" + assert _consume(_odp(_nested(16)), 524_288, _MAXIMUM_DEPTH).status == 200 + with pytest.raises(AgentError, match="nesting-depth"): + _consume(_odp(_nested(17)), 524_288, _MAXIMUM_DEPTH) + + +def test_refuses_a_service_document_nested_deeper_than_its_own_limit() -> None: + """ERR-21: the Service Document has the tighter allowance of 8.""" + assert _consume(_odp(_nested(8)), 65_536, _MAXIMUM_DOCUMENT_DEPTH).status == 200 + with pytest.raises(AgentError, match="nesting-depth"): + _consume(_odp(_nested(9)), 65_536, _MAXIMUM_DOCUMENT_DEPTH) + + +@pytest.mark.asyncio +async def test_refuses_a_deeply_nested_service_document_over_http() -> None: + document = json.loads(SERVICE_DOCUMENT) + cursor = document.setdefault("branding", {}) + for _ in range(12): + cursor["a"] = {} + cursor = cursor["a"] + client = _client(response(json.dumps(document))) + + with pytest.raises(AgentError, match="nesting-depth"): + await client.inspect() + + +def test_leaves_a_malformed_body_to_the_document_parser() -> None: + """A body that is not JSON is reported by the parser, which can say what is wrong with it.""" + assert _consume(_odp(b"not json"), 524_288, _MAXIMUM_DEPTH).body == b"not json" + + +# -- refused requests ----------------------------------------------------------------------------- + + +def _problem(detail: str) -> bytes: + return json.dumps( + { + "type": "https://offeringprotocol.org/problems/not-found", + "title": "Not found", + "status": 404, + "code": "NOT_FOUND", + "detail": detail, + } + ).encode() + + +def test_reports_the_detail_a_problem_document_gave() -> None: + with pytest.raises(ServiceRequestError) as raised: + _consume(_odp(_problem("No such Offering."), 404), 524_288, _MAXIMUM_DEPTH) + + assert raised.value.status == 404 + assert "No such Offering." in str(raised.value) + + +def test_reads_an_error_body_only_within_the_problem_details_limit() -> None: + """ERR-21 budgets a Problem Details response at 16,384 bytes. + + A larger body is not a Problem Details document this Agent will read, so the status still + describes the failure rather than the size becoming the failure. + """ + oversized = b" " * 20_000 + _problem("No such Offering.") + with pytest.raises(ServiceRequestError) as raised: + _consume(_odp(oversized, 404), 524_288, _MAXIMUM_DEPTH) + + assert raised.value.status == 404 + assert "No such Offering." not in str(raised.value) + + +def test_reports_an_error_body_that_is_not_a_problem_document() -> None: + with pytest.raises(ServiceRequestError, match="upstream exploded"): + _consume(_odp(b"upstream exploded", 503), 524_288, _MAXIMUM_DEPTH) + + +def test_checks_the_status_before_the_representation_limits() -> None: + """A refused request is reported as refused, whatever the error body looked like.""" + with pytest.raises(ServiceRequestError): + _consume(HttpResponse(500, {"content-type": "text/html"}, b"

oops

"), 10, 16) + + +# -- freshness ------------------------------------------------------------------------------------ + + +def test_reads_an_expires_written_with_an_unknown_zone() -> None: + """RFC 5322's "-0000" means an unknown zone, which parses to a value carrying no zone at all. + + A cache that stored one would raise the next time it compared that value against its clock, so + the zone ODP means -- UTC -- is supplied here instead. + """ + now = utc_now() + expires = _expiration({"expires": "Thu, 01 Dec 2050 16:00:00 -0000"}, timedelta(), now) + + assert expires.tzinfo is not None + assert now < expires + + +def test_treats_an_unreadable_expires_as_already_expired() -> None: + """RFC 9111 5.3: an invalid `Expires`, "0" above all, names a time in the past.""" + now = utc_now() + + assert _expiration({"expires": "0"}, timedelta(hours=1), now) == now + assert _expiration({"expires": "not a date"}, timedelta(hours=1), now) == now + assert _expiration({"expires": "Wed, 21 Oct 2037 07:28:00 GMT"}, timedelta(), now) > now + + +@pytest.mark.asyncio +async def test_does_not_reuse_a_representation_whose_expires_was_unreadable() -> None: + """The whole point of the rule above: the next request revalidates rather than reusing.""" + headers = {"cache-control": "public", "etag": '"v1"', "expires": "0"} + client = _client( + response(SERVICE_DOCUMENT, headers=headers), + response(SERVICE_DOCUMENT, headers=headers), + ) + await client.inspect() + await client.inspect() + + requests = client._transport.requests # type: ignore[attr-defined] + assert len(requests) == 2 + assert requests[1].headers["if-none-match"] == '"v1"' + + +@pytest.mark.asyncio +async def test_quarantines_a_sort_two_sources_publish() -> None: + """The cross-source rule applies to sorts, including the scope bookkeeping behind them.""" + result = SearchCapabilityCatalog() + target: dict[str, SortDefinition] = {} + scopes: dict[str, CapabilityScope] = {} + for scope in (CapabilityScope.SERVICE, CapabilityScope.COLLECTION): + await _add_sorts( + _client(), + result, + target, + scopes, + scope, + SearchCapabilities( + sorts=SortCapabilitySource(inline=[_sort("light"), _sort(f"{scope.value}-only")]) + ), + ) + + assert sorted(target) == ["collection-only", "service-only"] + assert sorted(scopes) == ["collection-only", "service-only"] + assert result.issues[-1].message == "Duplicate sorts: light" + + +def test_refuses_a_supporting_document_nested_past_the_parser( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A document deep enough to exhaust the JSON parser is refused, not raised as a RecursionError. + + An Attribute Schema carries no nesting-depth ceiling of its own, so depth alone can stop the + parser before any limit is measured. What matters is that the failure reaches the caller as the + unusable document it describes rather than as an interpreter error, which is what is asserted + here -- the depth at which a given build of CPython actually gives up is not this SDK's to fix, + and pinning a test to it would make the test a property of the platform. + """ + from offering_protocol.agent.client import _decode_json_object + + def exhausted(*args: object, **kwargs: object) -> object: + raise RecursionError("maximum recursion depth exceeded while decoding a JSON object") + + monkeypatch.setattr(json, "loads", exhausted) + + with pytest.raises(AgentError, match="nested too deeply"): + _decode_json_object(b"[[[]]]") diff --git a/tests/test_core_conformance.py b/tests/test_core_conformance.py new file mode 100644 index 0000000..3d38e5e --- /dev/null +++ b/tests/test_core_conformance.py @@ -0,0 +1,371 @@ +"""Core conformance: the document rules a JSON Schema cannot state. + +The schemas carry every rule about one member in isolation. What is left to code is the rules that +compare one member against another -- a repeated identifier, a range that runs backwards, a unit on +something that has no dimension. Each test below states one of those, with a control alongside it so +a rejection can be read as a consequence of the one thing that changed. +""" + +from __future__ import annotations + +import json +from typing import Any + +import pytest + +from offering_protocol.core import ( + OdpValidationError, + parse_agent_collection, + parse_agent_offering, + parse_agent_offering_page, + parse_collection, + parse_collection_search_request, + parse_filter_definition, + parse_offering, + parse_offering_page, + parse_sort_definition, +) + +OFFERING: dict[str, Any] = {"id": "plant-1", "name": "Monstera", "odp_version": "1.0"} +COLLECTION: dict[str, Any] = {"id": "plants", "name": "Plants", "odp_version": "1.0"} + + +def amend(base: dict[str, Any], **changes: Any) -> str: + return json.dumps({**base, **changes}) + + +def action(identifier: str) -> dict[str, Any]: + return { + "authentication": "not-required", + "id": identifier, + "rel": "purchase", + "http": {"href": "https://plants.example/checkout", "method": "POST"}, + } + + +def keywords_of(error: OdpValidationError) -> list[str]: + return [issue.keyword for issue in error.issues] + + +def assert_rejected_for(body: str, parse: Any, keyword: str) -> None: + with pytest.raises(OdpValidationError) as raised: + parse(body) + assert keyword in keywords_of(raised.value), keywords_of(raised.value) + + +# -- Offerings ------------------------------------------------------------------ + + +def test_refuses_a_repeated_action_identifier() -> None: + """OFR-57: a repeat leaves a caller unable to say which Action it meant.""" + assert parse_offering(amend(OFFERING, actions=[action("buy"), action("rent")])).actions + assert_rejected_for( + amend(OFFERING, actions=[action("buy"), action("buy")]), parse_offering, "unique-action-id" + ) + + +def _range(minimum: str, maximum: str) -> str: + return amend( + OFFERING, + price={"type": "range", "currency": "USD", "minimum": minimum, "maximum": maximum}, + ) + + +def test_refuses_an_inverted_price_range() -> None: + """OFR-49: a range whose minimum is above its maximum describes no price at all.""" + assert parse_offering(_range("5.00", "99.00")).price is not None + assert_rejected_for(_range("99.00", "5.00"), parse_offering, "price-range") + assert_rejected_for(_range("10", "9.99"), parse_offering, "price-range") + + +def test_orders_a_price_range_numerically() -> None: + """OFR-48: a price is a decimal string, so its bounds order numerically, not lexically.""" + # Lexically "9.00" sorts after "10.00"; numerically it does not. + assert parse_offering(_range("9.00", "10.00")) + # Trailing and leading zeros do not change a value, so these bounds are equal. + # A leading zero is not an ODP decimal at all, so the schema refuses it before any comparison. + with pytest.raises(OdpValidationError): + parse_offering(_range("07", "7")) + for minimum, maximum in (("5.0", "5.00"), ("0", "0.000"), ("1.5", "1.50")): + assert parse_offering(_range(minimum, maximum)), (minimum, maximum) + assert_rejected_for(_range("1.51", "1.5"), parse_offering, "price-range") + + +def test_leaves_a_price_without_bounds_alone() -> None: + for price in ( + {"type": "free"}, + {"type": "quote"}, + {"type": "fixed", "amount": "39.00", "currency": "USD"}, + {"type": "metered", "amount": "0.10", "currency": "USD", "unit": "litre"}, + ): + assert parse_offering(amend(OFFERING, price=price)), price + + +# -- Collections ---------------------------------------------------------------- + + +def test_refuses_a_collection_that_parents_itself() -> None: + """COL-20: a one-node cycle, which nothing walking the hierarchy upwards escapes.""" + assert parse_collection(amend(COLLECTION, parent_ids=["garden"])).parent_ids + for parents in (["plants"], ["garden", "plants"]): + assert_rejected_for(amend(COLLECTION, parent_ids=parents), parse_collection, "self-parent") + + +def test_refuses_a_repeated_parent() -> None: + """COL-19: a Collection names each parent once.""" + with pytest.raises(OdpValidationError): + parse_collection(amend(COLLECTION, parent_ids=["garden", "garden"])) + + +# -- Filter Definitions ----------------------------------------------------------- + + +def _filter(filter_type: str, **changes: Any) -> str: + return json.dumps( + { + "id": "weight", + "title": "Weight", + "description": "How heavy the plant is.", + "type": filter_type, + "operators": ["eq"], + **changes, + } + ) + + +def test_refuses_a_unit_on_a_filter_that_measures_nothing() -> None: + """FLT-10: only a numeric Filter has a dimension for a unit to name.""" + unit = {"system": "ucum", "code": "kg"} + for filter_type in ("boolean", "date", "date-time", "string"): + assert_rejected_for(_filter(filter_type, unit=unit), parse_filter_definition, "unit-type") + for filter_type in ("decimal", "integer", "number"): + assert parse_filter_definition(_filter(filter_type, unit=unit)), filter_type + + +def test_accepts_every_filter_type_without_a_unit() -> None: + for filter_type in ( + "boolean", + "date", + "date-time", + "decimal", + "integer", + "number", + "string", + ): + assert parse_filter_definition(_filter(filter_type)), filter_type + + +def test_refuses_an_ordering_operator_on_an_unordered_type() -> None: + """FLT-11: an ordering operator on a type with no order cannot be evaluated.""" + for filter_type in ("boolean", "string"): + for operator in ("gt", "gte", "lt", "lte"): + assert_rejected_for( + _filter(filter_type, operators=["eq", operator]), + parse_filter_definition, + "operator-type", + ) + for filter_type in ("date", "date-time", "decimal", "integer", "number"): + assert parse_filter_definition(_filter(filter_type, operators=["gt", "lte"])), filter_type + + +# -- Sort Definitions ------------------------------------------------------------- + + +def _sort(*filter_ids: str) -> str: + return json.dumps( + { + "id": "cheapest", + "title": "Cheapest first", + "description": "Orders plants by price.", + "keys": [ + {"filter_id": filter_id, "direction": "ascending", "missing": "last"} + for filter_id in filter_ids + ], + } + ) + + +def test_refuses_a_recipe_that_orders_by_one_filter_twice() -> None: + """FLT-41: ordering by one Filter twice cannot change the order, so a repeat means nothing.""" + assert parse_sort_definition(_sort("price", "weight")).keys + assert_rejected_for(_sort("price", "price"), parse_sort_definition, "unique-filter-id") + assert_rejected_for( + _sort("price", "weight", "price"), parse_sort_definition, "unique-filter-id" + ) + + +# -- Refinements ------------------------------------------------------------------- + + +def _page(*groups: dict[str, Any]) -> str: + document: dict[str, Any] = {"odp_version": "1.0", "items": []} + if groups: + document["refinements"] = list(groups) + return json.dumps(document) + + +def _group(filter_id: str, *values: dict[str, Any]) -> dict[str, Any]: + return {"filter_id": filter_id, "values": list(values)} + + +def test_refuses_two_refinement_groups_for_one_filter() -> None: + """FLT-30: a repeat leaves an Agent unable to say which group belongs to that definition.""" + assert parse_offering_page( + _page( + _group("colour", {"value": "red", "count": 1}), + _group("size", {"value": "l", "count": 2}), + ) + ) + assert_rejected_for( + _page( + _group("colour", {"value": "red", "count": 1}), + _group("colour", {"value": "blue", "count": 2}), + ), + parse_offering_page, + "unique-filter-id", + ) + + +def test_refuses_a_repeated_bucket_value() -> None: + """FLT-32: `uniqueItems` compares whole buckets, so it passes one value with two counts.""" + for values in ( + ({"value": "green", "count": 4}, {"value": "green", "count": 2}), + ({"value": True, "count": 4}, {"value": True, "count": 2}), + ({"value": 3, "count": 4}, {"value": 3.0, "count": 2}), + ): + assert_rejected_for( + _page(_group("colour", *values)), parse_offering_page, "unique-bucket-value" + ) + + +def test_reads_two_spellings_of_one_decimal_as_one_bucket_value() -> None: + """FLT-32: decimal equality is numeric rather than lexical. + + So `1.0` and `1.00` name one value, and a group offering both hands a caller two counts for one + candidate with no way to choose between them. + """ + for values in ( + ({"value": "1.0", "count": 4}, {"value": "1.00", "count": 2}), + ({"value": "0", "count": 4}, {"value": "0.0", "count": 2}), + ({"value": "12", "count": 4}, {"value": "12.000", "count": 2}), + ({"value": "-1.5", "count": 4}, {"value": "-1.50", "count": 2}), + ): + assert_rejected_for( + _page(_group("weight", *values)), parse_offering_page, "unique-bucket-value" + ) + + +def test_keeps_bucket_values_that_differ_apart() -> None: + for values in ( + ({"value": "1.0", "count": 4}, {"value": "2.0", "count": 2}), + ({"value": "1.01", "count": 4}, {"value": "1.1", "count": 2}), + ({"value": "10", "count": 4}, {"value": "1.0", "count": 2}), + ({"value": "-1.0", "count": 4}, {"value": "1.0", "count": 2}), + # Neither is an ODP decimal -- a leading zero and a trailing period are not -- so both are + # compared as the strings they are. + ({"value": "01", "count": 4}, {"value": "1", "count": 2}), + ({"value": "1.", "count": 4}, {"value": "1", "count": 2}), + ({"value": True, "count": 4}, {"value": False, "count": 2}), + ({"value": 1, "count": 4}, {"value": "1", "count": 2}), + ): + assert parse_offering_page(_page(_group("weight", *values))), values + + +def test_reports_the_group_a_repeated_bucket_value_was_in() -> None: + with pytest.raises(OdpValidationError) as raised: + parse_offering_page( + _page( + _group("colour", {"value": "red", "count": 1}), + _group("size", {"value": "l", "count": 1}, {"value": "l", "count": 2}), + ) + ) + + assert [issue.path for issue in raised.value.issues] == ["/refinements/1/values"] + + +# -- what an Agent tolerates --------------------------------------------------------- + + +def test_hands_an_agent_a_document_a_service_must_not_publish() -> None: + """ROLE-03: an Agent describes a defect to its caller rather than discarding what it can use. + + The invariants above are what a Service must satisfy before publishing. On the Agent side they + are reported against the Action, hierarchy or group they concern, so the document survives. + """ + offering = amend(OFFERING, actions=[action("buy"), action("buy")]) + with pytest.raises(OdpValidationError): + parse_offering(offering) + assert len(parse_agent_offering(offering).actions) == 2 + + collection = amend(COLLECTION, parent_ids=["plants"]) + with pytest.raises(OdpValidationError): + parse_collection(collection) + assert parse_agent_collection(collection).parent_ids == ["plants"] + + page = _page( + _group("colour", {"value": "red", "count": 1}), + _group("colour", {"value": "blue", "count": 2}), + ) + with pytest.raises(OdpValidationError): + parse_offering_page(page) + assert len(parse_agent_offering_page(page).refinements) == 2 + + +def test_still_refuses_an_agent_document_it_cannot_read_at_all() -> None: + """Tolerance stops at defects that leave nothing to use.""" + for parse in (parse_agent_offering, parse_agent_collection): + with pytest.raises(OdpValidationError): + parse("not json") + with pytest.raises(OdpValidationError): + parse("{}") + with pytest.raises(OdpValidationError): + parse_agent_offering(amend(OFFERING, language="not a tag", localizations=["not a tag"])) + + +def test_filters_what_a_later_odp_version_added_before_reading_it() -> None: + """The Agent entry points normalize first, so a member ODP does not define never reaches the + schema.""" + unknown = amend(OFFERING, price={"type": "auction", "reserve": "10.00"}) + + with pytest.raises(OdpValidationError): + parse_offering(unknown) + assert parse_agent_offering(unknown).price is None + + +# -- a null parent asks a question an absent one does not ------------------------------- + + +def test_tells_a_null_parent_apart_from_an_absent_one() -> None: + """COL-06: a null `parent_id` selects root Collections; an absent one applies no constraint.""" + rooted = parse_collection_search_request('{"odp_version":"1.0","query":"x","parent_id":null}') + anywhere = parse_collection_search_request('{"odp_version":"1.0","query":"x"}') + named = parse_collection_search_request('{"odp_version":"1.0","parent_id":"garden"}') + + assert rooted.parent_id is None and "parent_id" in rooted.model_fields_set + assert anywhere.parent_id is None and "parent_id" not in anywhere.model_fields_set + assert named.parent_id == "garden" + + # And the difference survives the round trip, which is what a Service reads it back from. + assert rooted.to_dict()["parent_id"] is None + assert "parent_id" not in anywhere.to_dict() + + +def test_refuses_a_collection_search_that_asks_nothing() -> None: + with pytest.raises(OdpValidationError): + parse_collection_search_request('{"odp_version":"1.0"}') + + +def test_compares_any_bucket_value_the_model_can_hold() -> None: + """`RefinementBucket.value` is typed as any JSON value, so the comparison is total over one. + + The schema narrows what actually arrives to a scalar, and this keeps the two from drifting: if + that narrowing were ever relaxed, a composite value would still compare by what it contains + rather than by object identity. + """ + from offering_protocol.core.validation import _bucket_key + + assert _bucket_key({"a": 1, "b": 2}) == _bucket_key({"b": 2, "a": 1}) + assert _bucket_key([1, 2]) != _bucket_key([2, 1]) + assert _bucket_key(None) != _bucket_key("null") + assert _bucket_key(True) != _bucket_key(1) + assert _bucket_key("1.0") == _bucket_key("1.00") diff --git a/tests/test_directory.py b/tests/test_directory.py index 3377421..7b7a74d 100644 --- a/tests/test_directory.py +++ b/tests/test_directory.py @@ -32,7 +32,10 @@ "language":"en", "localizations":["en"], "name":"Indica Flowers", - "operations":[], + "operations":[ + {"authentication":"not-required","name":"get-offering"}, + {"authentication":"not-required","name":"list-offerings"} + ], "service_origin":"https://demo.inflowpay.ai" }] }""" @@ -120,11 +123,15 @@ async def test_search_filters_unknown_protocols_and_rejects_malformed_known() -> ).search_services(SearchRequest()) assert page.items[0].protocols is None + # ROLE-03: a descriptor bearing a recognized name stays subject to every rule for that + # descriptor, so this record is unusable -- but it is dropped and reported, not raised, because + # a Directory is a discovery aid and the other Services on the page remain findable. malformed = candidate.replace('"name":"mpp"', '"name":"mpp","extra":true') - with pytest.raises(DirectoryError): - await DirectoryClient( - transport=QueueTransport(response(malformed, content_type="application/json")) - ).search_services(SearchRequest()) + page = await DirectoryClient( + transport=QueueTransport(response(malformed, content_type="application/json")) + ).search_services(SearchRequest()) + assert not page.items + assert page.issues[0].index == 0 @pytest.mark.asyncio diff --git a/tests/test_directory_conformance.py b/tests/test_directory_conformance.py new file mode 100644 index 0000000..377a94a --- /dev/null +++ b/tests/test_directory_conformance.py @@ -0,0 +1,519 @@ +"""Directory conformance: what an Agent may believe about what a Directory tells it. + +ODP defines no directory wire format (ROLE-07, CNF-14), so almost nothing here is a wire-format +rule. What binds is ROLE-03: a Directory is a discovery aid indexing metadata it does not own, and +an Agent must not treat what it publishes as authoritative Service data. Every test below states one +consequence of that -- a record is validated before it is handed back, an unusable record costs only +itself, and nothing a Directory made up is passed off as a checked Service Document member. +""" + +from __future__ import annotations + +import json +from ipaddress import ip_address + +import pytest + +from helpers import QueueTransport, response +from offering_protocol.core.models import Protocol, TrustProtocol +from offering_protocol.directory import ( + DirectoryClient, + DirectoryError, + DirectoryRequestError, + Environment, + SearchPage, + SearchRequest, + ServiceFilters, + SuggestionRequest, +) +from offering_protocol.directory.addresses import is_public + +BASELINE_OPERATIONS = [ + {"authentication": "not-required", "name": "get-offering"}, + {"authentication": "not-required", "name": "list-offerings"}, +] +SERVICE: dict[str, object] = { + "description": "An AI-enabled plant store.", + "indexed_at": "2026-01-01T00:00:00Z", + "language": "en", + "localizations": ["en"], + "name": "Plants", + "operations": BASELINE_OPERATIONS, + "service_origin": "https://plants.example", +} + + +def amend(**changes: object) -> dict[str, object]: + return {**SERVICE, **changes} + + +def json_response(value: object, **kwargs: object) -> object: + body = value if isinstance(value, str) else json.dumps(value) + return response(body, content_type="application/json", **kwargs) # type: ignore[arg-type] + + +def client(*replies: object) -> DirectoryClient: + return DirectoryClient(Environment.PRODUCTION, transport=QueueTransport(*replies)) # type: ignore[arg-type] + + +async def read(*services: object, **page: object) -> SearchPage: + body = json_response({"items": list(services), **page}) + return await client(body).search_services(SearchRequest()) + + +# -- a record is validated before it is handed back --------------------------------------- + + +@pytest.mark.asyncio +async def test_reads_a_conformant_record() -> None: + page = await read(SERVICE) + + assert len(page.items) == 1 + assert not page.issues + assert page.items[0].service_origin == "https://plants.example" + assert page.items[0].name == "Plants" + + +@pytest.mark.asyncio +async def test_refuses_an_origin_that_is_not_a_canonical_https_origin() -> None: + """IDN-01: a Service is identified by its canonical origin, so two spellings are not one.""" + for origin in ( + "https://PLANTS.example", + "https://plants.example:443", + "https://plants.example/odp", + "https://plants.example?a=1", + "http://plants.example", + "not a url", + "", + ): + page = await read(amend(service_origin=origin)) + assert not page.items, origin + assert page.issues[0].index == 0, origin + + +@pytest.mark.asyncio +async def test_refuses_an_origin_that_is_not_a_string() -> None: + page = await read(amend(service_origin=None), amend(service_origin=["https://a.example"])) + + assert not page.items + assert len(page.issues) == 2 + + +@pytest.mark.asyncio +async def test_refuses_an_origin_on_a_private_or_loopback_host() -> None: + """A public Directory has no business pointing an Agent inside its own network. + + The default transport refuses these too, but a consumer who installed their own transport, or + who enabled local development for a Service of their own, would otherwise have no guard left. + """ + for origin in ( + "https://127.0.0.1", + "https://10.0.0.1", + "https://169.254.169.254", + "https://[::1]", + "https://[64:ff9b::a9fe:a9fe]", + ): + page = await read(amend(service_origin=origin)) + assert not page.items, origin + assert "private or loopback" in page.issues[0].message, origin + + +@pytest.mark.asyncio +async def test_refuses_an_indexed_at_that_is_not_an_rfc_3339_date_time() -> None: + """A caller compares or slices this value, so it needs one shape rather than whatever parsed.""" + for value in ("yesterday", "2026-01-01", "December 17, 1995", "2026-01-01T00:00:00", 17, None): + page = await read(amend(indexed_at=value)) + assert not page.items, value + assert "RFC 3339" in page.issues[0].message, value + + for value in ("2026-01-01T00:00:00Z", "2026-01-01t00:00:00.123z", "2026-01-01T00:00:00+02:00"): + assert (await read(amend(indexed_at=value))).items, value + + +@pytest.mark.asyncio +async def test_holds_echoed_service_document_members_to_the_service_document_rules() -> None: + """A record echoes Service-owned metadata. + + Echoing it does not lower the bar it has to clear: a Directory that published + `"language": "not a tag"` would otherwise hand a caller a tag no language matcher can read. + """ + changes: tuple[dict[str, object], ...] = ( + {"language": "not a tag"}, + {"localizations": ["fr"]}, + {"localizations": ["en", "EN"]}, + {"operations": []}, + {"operations": [{"authentication": "biometric", "name": "list-offerings"}]}, + {"name": ""}, + {"website_url": "javascript:alert(1)"}, + ) + for change in changes: + page = await read(amend(**change)) + assert not page.items, change + assert page.issues[0].index == 0, change + + +# -- one unusable record costs only itself (ROLE-03) --------------------------------------- + + +@pytest.mark.asyncio +async def test_keeps_the_services_beside_an_unusable_record() -> None: + """Rejecting the page made every other Service in the result undiscoverable.""" + page = await read( + SERVICE, + amend(service_origin="not a url", name="Broken"), + amend(service_origin="https://ferns.example", name="Ferns"), + ) + + assert [item.name for item in page.items] == ["Plants", "Ferns"] + assert [issue.index for issue in page.issues] == [1] + + +@pytest.mark.asyncio +async def test_reports_each_unusable_record_at_its_own_position() -> None: + page = await read(amend(indexed_at="yesterday"), SERVICE, amend(language="not a tag")) + + assert len(page.items) == 1 + assert [issue.index for issue in page.issues] == [0, 2] + assert all(issue.message for issue in page.issues) + + +@pytest.mark.asyncio +async def test_reports_an_entry_that_is_not_an_object() -> None: + page = await read("https://plants.example", 17, None) + + assert not page.items + assert len(page.issues) == 3 + + +# -- what a Directory says about protocols -------------------------------------------------- + + +@pytest.mark.asyncio +async def test_drops_a_protocol_this_odp_version_does_not_name() -> None: + page = await read( + amend(protocols={"payments": [{"authentication": "not-required", "name": "future"}]}) + ) + + assert page.items[0].protocols is None + assert not page.issues + + +@pytest.mark.asyncio +async def test_keeps_a_protocol_this_odp_version_names() -> None: + page = await read(amend(protocols={"trust": [{"name": "tap"}]})) + + protocols = page.items[0].protocols + assert protocols is not None + assert protocols.trust[0].name.value == "tap" + + +@pytest.mark.asyncio +async def test_refuses_a_recognized_descriptor_that_breaks_its_own_rules() -> None: + """SVC-51: filtering the unknown does not make what remains valid.""" + page = await read( + amend(protocols={"payments": [{"authentication": "not-required", "name": "mpp", "x": 1}]}) + ) + + assert not page.items + assert page.issues[0].index == 0 + + +# -- nothing unchecked is passed off as checked ----------------------------------------------- + + +@pytest.mark.asyncio +async def test_drops_service_document_members_it_does_not_validate() -> None: + """A caller reading these off a record cannot tell they were never checked. + + `http` is the one that matters: a caller could build request URLs from an `endpoint_base` the + Directory invented, which is exactly the authority ROLE-03 says a Directory does not have. + """ + page = await read( + amend( + branding={"icon": {"src": "/i.png"}, "logo": {"src": "/l.png"}}, + http={"endpoint_base": "/somewhere-else"}, + mcp=[{"type": "streamable-http", "url": "https://elsewhere.example/mcp"}], + odp_version="1.0", + payment_origins=["https://pay.example"], + search_capabilities={"filters": {"inline": []}}, + ) + ) + + assert not set(page.items[0].additional) + + +@pytest.mark.asyncio +async def test_keeps_the_directory_owned_members_it_was_given() -> None: + """Directory-owned signals are the Directory's to publish, and are passed through untouched.""" + page = await read(amend(ranking_score=0.94, verified=True)) + + assert page.items[0].additional == {"ranking_score": 0.94, "verified": True} + + +# -- page-level rules ------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_refuses_a_page_carrying_more_services_than_one_page_holds() -> None: + page = json_response({"items": [SERVICE] * 101}) + with pytest.raises(DirectoryError, match="100 Services"): + await client(page).search_services(SearchRequest()) + + +@pytest.mark.asyncio +async def test_refuses_a_body_that_is_not_a_search_page() -> None: + for body in ("not json", "[]", '{"facets":{}}', '{"items":"none"}'): + with pytest.raises(DirectoryError): + await client(json_response(body)).search_services(SearchRequest()) + + +@pytest.mark.asyncio +async def test_refuses_trust_facets_naming_another_protocol() -> None: + """`tap` is the only trust protocol this ODP version names, so anything else is unreadable.""" + facets = {"trust": [{"count": 2, "value": {"name": "mpp"}}]} + with pytest.raises(DirectoryError, match="trust facets"): + await read(SERVICE, facets=facets) + + page = await read(SERVICE, facets={"trust": [{"count": 2, "value": {"name": "tap"}}]}) + assert page.facets is not None + assert page.facets.trust[0].count == 2 + + +# -- continuations stay on the canonical Directory ----------------------------------------------- + + +@pytest.mark.asyncio +async def test_keeps_a_continuation_on_the_canonical_origin() -> None: + directory = client(json_response({"items": [SERVICE]})) + await directory.continue_search_services("/v1/services/search?cursor=2") + + assert directory._transport.requests[0].url == ( # type: ignore[attr-defined] + "https://api.inflowpay.ai/v1/services/search?cursor=2" + ) + + +@pytest.mark.asyncio +async def test_refuses_a_continuation_that_leaves_the_canonical_origin() -> None: + for reference in ( + "https://elsewhere.example/next", + "//elsewhere.example/next", + "https://user@api.inflowpay.ai/next", + "mailto:someone@example.com", + ): + with pytest.raises(DirectoryError): + await client().continue_search_services(reference) + + # A continuation format is the Directory's to define, so an ordinary relative reference is + # followed -- what is checked is where it lands, not how it was spelled. + directory = client(json_response({"items": []})) + await directory.continue_search_services("services?cursor=2") + assert directory._transport.requests[0].url.startswith( # type: ignore[attr-defined] + "https://api.inflowpay.ai/" + ) + + +@pytest.mark.asyncio +async def test_refuses_a_redirect_that_leaves_the_origin_or_points_nowhere() -> None: + for location in ("https://elsewhere.example/x", "mailto:someone@example.com"): + redirect = json_response("", status=302, headers={"location": location}) + with pytest.raises(DirectoryError): + await client(redirect).search_services(SearchRequest()) + + missing = json_response("", status=302) + with pytest.raises(DirectoryError, match="Location"): + await client(missing).search_services(SearchRequest()) + + +# -- requests this client will not send ------------------------------------------------------------ + + +@pytest.mark.asyncio +async def test_refuses_a_limit_that_is_not_a_count() -> None: + """Zero means the Directory decides; anything below one asks for nothing.""" + for limit in (-1, 101): + with pytest.raises(DirectoryError, match="1 through 100"): + await client().search_services(SearchRequest(limit=limit)) + for limit in (-1, 26): + with pytest.raises(DirectoryError, match="1 through 25"): + await client().suggest_services(SuggestionRequest(prefix="pl", limit=limit)) + + +@pytest.mark.asyncio +async def test_sends_the_limits_it_accepts() -> None: + directory = client(json_response({"items": []})) + await directory.search_services(SearchRequest(limit=100)) + assert json.loads(directory._transport.requests[0].body) == {"limit": 100} # type: ignore[attr-defined] + + directory = client(json_response({"items": []})) + await directory.suggest_services(SuggestionRequest(prefix="pl", limit=25)) + assert "limit=25" in directory._transport.requests[0].url # type: ignore[attr-defined] + + directory = client(json_response({"items": []})) + await directory.suggest_services(SuggestionRequest(prefix="pl")) + assert "limit" not in directory._transport.requests[0].url # type: ignore[attr-defined] + + +@pytest.mark.asyncio +async def test_refuses_a_query_or_keyword_set_it_cannot_send() -> None: + for request in ( + SearchRequest(query=" padded "), + SearchRequest(query="x" * 513), + SearchRequest(filters=ServiceFilters(keywords=["x"] * 33)), + SearchRequest(filters=ServiceFilters(keywords=[""])), + SearchRequest(filters=ServiceFilters(keywords=["x" * 65])), + ): + with pytest.raises(DirectoryError): + await client().search_services(request) + + +@pytest.mark.asyncio +async def test_refuses_a_trust_filter_that_is_not_the_one_descriptor_odp_names() -> None: + for trust in ([], [TrustProtocol(name=Protocol.TAP), TrustProtocol(name=Protocol.TAP)]): + with pytest.raises(DirectoryError, match="tap"): + await client().search_services(SearchRequest(filters=ServiceFilters(trust=trust))) + + directory = client(json_response({"items": []})) + await directory.search_services( + SearchRequest(filters=ServiceFilters(trust=[TrustProtocol(name=Protocol.TAP)])) + ) + assert directory._transport.requests # type: ignore[attr-defined] + + +# -- refused requests and oversized responses ----------------------------------------------------- + + +@pytest.mark.asyncio +async def test_reports_a_refused_request_by_its_status() -> None: + """A refused request is a status, not a size complaint, whatever the body looked like.""" + with pytest.raises(DirectoryRequestError) as raised: + await client(json_response("x" * 600_000, status=503)).search_services(SearchRequest()) + + assert raised.value.status == 503 + assert len(str(raised.value)) < 4_096 + + +@pytest.mark.asyncio +async def test_refuses_a_response_larger_than_one_page_may_be() -> None: + body = json.dumps({"items": [], "padding": "x" * 600_000}) + with pytest.raises(DirectoryError, match="524288"): + await client(json_response(body)).search_services(SearchRequest()) + + +@pytest.mark.asyncio +async def test_refuses_a_response_that_is_not_json() -> None: + with pytest.raises(DirectoryError, match="application/json"): + await client(response('{"items":[]}', content_type="text/html")).search_services( + SearchRequest() + ) + + +# -- suggestions are the Directory's index, not Offering-search values ---------------------------- + + +@pytest.mark.asyncio +async def test_reads_suggestions_the_directory_returned() -> None: + directory = client(json_response({"items": ["plants", "planters"]})) + + assert await directory.suggest_services(SuggestionRequest(prefix="plan")) == [ + "plants", + "planters", + ] + + +@pytest.mark.asyncio +async def test_refuses_suggestions_it_cannot_use() -> None: + for body in ( + '{"suggestions":[]}', + '["plants", 17]', + '["plants", ""]', + '["plants", " padded "]', + f'["{"x" * 129}"]', + json.dumps(["s"] * 26), + "not json", + ): + with pytest.raises(DirectoryError): + await client(json_response(body)).suggest_services(SuggestionRequest(prefix="plan")) + + +@pytest.mark.asyncio +async def test_refuses_a_prefix_it_cannot_send() -> None: + for prefix in ("", " ", "x" * 129): + with pytest.raises(DirectoryError, match="prefix"): + await client().suggest_services(SuggestionRequest(prefix=prefix)) + + +# -- addresses the public internet does not route ------------------------------------------------- + + +def test_judges_an_address_by_the_special_purpose_registries() -> None: + """SEC-08, and the reason this table exists rather than `ipaddress.is_global`. + + The IPv6 transition ranges each embed an IPv4 address: without them a name resolving to + `64:ff9b::a9fe:a9fe` reaches link-local 169.254.169.254 through a NAT64 gateway -- the cloud + metadata endpoint. The standard library reports that address, and five other registry ranges, + as globally routable. + """ + for value in ("8.8.8.8", "1.1.1.1", "2606:4700::1111", "::ffff:8.8.8.8"): + assert is_public(ip_address(value)), value + + for value in ( + "0.0.0.0", + "10.0.0.1", + "100.64.0.1", + "127.0.0.1", + "169.254.169.254", + "172.16.0.1", + "192.0.0.1", + "192.0.2.1", + "192.31.196.1", + "192.88.99.1", + "192.168.1.1", + "192.175.48.1", + "198.18.0.1", + "198.51.100.1", + "203.0.113.1", + "224.0.0.1", + "240.0.0.1", + "::", + "::1", + "::ffff:169.254.169.254", + "64:ff9b::a9fe:a9fe", + "64:ff9b:1::1", + "100::1", + "2001::1", + "2001:db8::1", + "2002:a9fe:a9fe::1", + "2620:4f:8000::1", + "5f00::1", + "fc00::1", + "fe80::1", + "fec0::1", + "ff02::1", + ): + assert not is_public(ip_address(value)), value + + +def test_reads_an_ipv4_mapped_address_as_the_address_it_carries() -> None: + assert is_public(ip_address("::ffff:8.8.8.8")) + assert not is_public(ip_address("::ffff:10.0.0.1")) + + +@pytest.mark.asyncio +async def test_refuses_a_page_whose_own_members_are_unreadable() -> None: + """Per-record tolerance does not extend to the page around the records. + + A record this client cannot read is one Service it cannot offer; a page it cannot read is a + result set of unknown shape, and there is nothing to hand back. + """ + for page in ({"items": [], "next": 17}, {"items": [], "facets": {"trust": "all"}}): + with pytest.raises(DirectoryError, match="invalid Directory response"): + await client(json_response(page)).search_services(SearchRequest()) + + +@pytest.mark.asyncio +async def test_accepts_an_origin_that_is_a_public_address_literal() -> None: + """The host check is about where the address routes, not about it being a name.""" + page = await read(amend(service_origin="https://8.8.8.8")) + + assert page.items[0].service_origin == "https://8.8.8.8" + assert not page.issues diff --git a/tests/test_edges.py b/tests/test_edges.py index 2dfaf90..36cf57e 100644 --- a/tests/test_edges.py +++ b/tests/test_edges.py @@ -19,7 +19,6 @@ _cacheable, _expiration, _has_freshness, - _normalize_body, _operation_parser, ) from offering_protocol.agent.client import ( @@ -53,7 +52,12 @@ ) from offering_protocol.core import __all__ as core_exports from offering_protocol.core.references import ReferenceError -from offering_protocol.core.validation import OdpValidationError, _is_language_tag, _parse +from offering_protocol.core.validation import ( + OdpValidationError, + _agent_body, + _is_language_tag, + _parse, +) from offering_protocol.directory import ( DirectoryClient, DirectoryError, @@ -74,6 +78,7 @@ ) from offering_protocol.service.service import ( _encode, + _Exchange, _json_response, _validate_collection_representation, _validate_offering_representation, @@ -185,7 +190,7 @@ async def test_directory_response_edge_cases_and_real_transport_adapter( service = ( '{"description":"Plants","indexed_at":"2026-01-01T00:00:00Z",' '"language":"en","localizations":["en"],"name":"Plant",' - '"operations":[],"service_origin":"https://plants.example"}' + '"operations":[{"authentication":"not-required","name":"get-offering"},{"authentication":"not-required","name":"list-offerings"}],"service_origin":"https://plants.example"}' ) with pytest.raises(DirectoryError): await DirectoryClient( @@ -196,17 +201,19 @@ async def test_directory_response_edge_cases_and_real_transport_adapter( ) ) ).search_services(SearchRequest()) - with pytest.raises(DirectoryError): - await DirectoryClient( - transport=QueueTransport( - response( - '{"items":[' - + service.replace("https://plants.example", "https://PLANTS.example") - + "]}", - content_type="application/json", - ) + # A non-canonical origin makes that record unusable, not the page. + non_canonical = await DirectoryClient( + transport=QueueTransport( + response( + '{"items":[' + + service.replace("https://plants.example", "https://PLANTS.example") + + "]}", + content_type="application/json", ) - ).search_services(SearchRequest()) + ) + ).search_services(SearchRequest()) + assert not non_canonical.items + assert "canonical" in non_canonical.issues[0].message async def handler(request: httpx.Request) -> httpx.Response: assert request.url.host == "93.184.216.34" @@ -297,7 +304,7 @@ async def list_offerings(self, request: CatalogRequest) -> OfferingPage[Offering await catalog.list_collection_offerings("missing", CatalogRequest()) assert _encode({"answer": 42}) == b'{"answer":42}' with pytest.raises(ServiceError): - _json_response(200, "x" * 10, 1) + _json_response("x" * 10, 1, _Exchange(headers={}, language="en", method="GET")) @pytest.mark.asyncio @@ -341,9 +348,12 @@ async def test_agent_collection_search_and_cache_header_edges() -> None: {"cache-control": "max-age=10", "age": "3"}, timedelta(), now ) == now + timedelta(seconds=7) assert _expiration({"expires": "Wed, 21 Oct 2037 07:28:00 GMT"}, timedelta(), now).year == 2037 - assert _expiration({"expires": "bad"}, timedelta(seconds=2), now) == now + timedelta(seconds=2) + # RFC 9111 5.3: an `Expires` the cache cannot read names a time in the past. + assert _expiration({"expires": "bad"}, timedelta(seconds=2), now) == now assert _operation_parser(Operation.GET_COLLECTION) is parse_agent_collection - assert _normalize_body(b"[]", "collection") == "[]" + # A body that is not a JSON object has no members to filter, so it reaches the schema as it + # stands and is refused there rather than during normalization. + assert _agent_body(b"[]", "collection") == "[]" @pytest.mark.asyncio diff --git a/tests/test_service_conformance.py b/tests/test_service_conformance.py new file mode 100644 index 0000000..91c95af --- /dev/null +++ b/tests/test_service_conformance.py @@ -0,0 +1,608 @@ +"""Service conformance: what a Service puts on the wire, beyond the documents it serves. + +The document rules live in core. What a Service owns is the exchange around them -- which variant it +selected and how it said so, the validator it issued, the methods each resource answers, and what it +refuses before a Catalog is ever asked. Each test states one of those rules. +""" + +from __future__ import annotations + +import json + +import pytest + +from offering_protocol.core import ( + Collection, + CollectionSearchRequest, + Offering, + OfferingPage, + OfferingSearchRequest, + Operation, + Page, +) +from offering_protocol.service import ( + MEDIA_TYPE, + PROBLEM_MEDIA_TYPE, + CatalogRequest, + Request, + Response, + Service, + ServiceBuilder, + StaticCatalog, + StaticCatalogOptions, +) + +TAGS = ["en", "en-GB", "fr", "de-CH", "zh-Hant"] + + +def _offering(identifier: str, *, collected: bool = False) -> Offering: + document: dict[str, object] = { + "id": identifier, + "name": identifier.title(), + "odp_version": "1.0", + "description": f"A plant called {identifier}.", + } + if collected: + document["collection_ids"] = ["plants"] + return Offering.model_validate(document) + + +def _catalog(count: int = 3) -> StaticCatalog: + return StaticCatalog( + StaticCatalogOptions( + collections=( + Collection.model_validate({"id": "plants", "name": "Plants", "odp_version": "1.0"}), + ), + offerings=tuple(_offering(f"p{index}", collected=True) for index in range(count)), + ) + ) + + +def service(*, localizations: list[str] | None = None, count: int = 3) -> Service: + builder = ServiceBuilder("Plants", "A plant store.", "en", "/odp") + if localizations is not None: + builder = builder.localizations(localizations) + return builder.build(_catalog(count)) + + +async def call(method: str, path: str, **kwargs: object) -> Response: + return await service().handle(Request(method=method, path=path, **kwargs)) # type: ignore[arg-type] + + +async def localized(accept_language: str, path: str = "/.well-known/odp") -> Response: + return await service(localizations=TAGS).handle( + Request(method="GET", path=path, headers={"accept-language": accept_language}) + ) + + +EVERY_PATH = ( + "/.well-known/odp", + "/odp/offerings", + "/odp/offerings/p0", + "/odp/collections", + "/odp/collections/plants", + "/odp/collections/plants/offerings", +) + + +# -- which variant it served --------------------------------------------------- + + +@pytest.mark.asyncio +async def test_selects_a_language_by_rfc_4647_lookup() -> None: + """SVC-58: Lookup walks a range down its own subtags and never sideways.""" + for accept, expected in ( + ("fr", "fr"), + ("FR", "fr"), + ("en-GB", "en-GB"), + # Lookup truncates: en-GB-oed has no match, en-GB does. + ("en-GB-oed", "en-GB"), + # de-CH-1901 falls back to de-CH, never sideways to another de-* tag. + ("de-CH-1901", "de-CH"), + ("zh-Hant-TW", "zh-Hant"), + # A single-character subtag is an extension singleton, dropped with its parent. + ("en-GB-a-bbb", "en-GB"), + ): + assert (await localized(accept)).headers["content-language"] == expected, accept + + +@pytest.mark.asyncio +async def test_never_refuses_a_request_over_language() -> None: + """SVC-59: nothing matching is not a reason to refuse, only to serve the default.""" + for accept in ("de-DE", "ja", "*", "ja, ko;q=0.5"): + reply = await localized(accept, "/odp/offerings") + assert reply.status == 200, accept + assert reply.headers["content-language"] == "en", accept + + +@pytest.mark.asyncio +async def test_honours_the_order_the_agent_asked_in() -> None: + for accept, expected in ( + ("fr;q=0.5, de-CH;q=0.9", "de-CH"), + ("de-CH;q=0.1, fr", "fr"), + ("fr, de-CH", "fr"), + # RFC 9110: a range weighted zero is not wanted at all. + ("fr;q=0, de-CH", "de-CH"), + ("fr;q=0, de-CH;q=0", "en"), + # `*` is the residual, so it cannot outrank a range the caller named. + ("fr, *;q=0.9", "fr"), + # A weight outside the grammar leaves the entry with nothing to honour. + ("fr;q=nonsense", "en"), + ): + assert (await localized(accept)).headers["content-language"] == expected, accept + + +@pytest.mark.asyncio +async def test_describes_the_variant_it_served_on_every_representation() -> None: + """SVC-60: `Vary` is what stops a shared cache handing one variant to another request.""" + for path in EVERY_PATH: + reply = await call("GET", path) + assert reply.headers["content-language"] == "en", path + assert reply.headers["vary"] == "Accept, Accept-Language", path + assert reply.headers["content-type"] == MEDIA_TYPE, path + + +@pytest.mark.asyncio +async def test_tells_the_catalog_which_variant_to_answer_in() -> None: + """A Catalog is handed the selected tag rather than being left to parse the field itself.""" + seen: list[CatalogRequest] = [] + + class Recording(StaticCatalog): + async def list_offerings(self, request: CatalogRequest) -> OfferingPage[Offering]: + seen.append(request) + return await super().list_offerings(request) + + built = ServiceBuilder("Plants", "A plant store.", "en", "/odp") + built = built.localizations(TAGS) + await built.build(Recording(StaticCatalogOptions(offerings=(_offering("p0"),)))).handle( + Request(method="GET", path="/odp/offerings", headers={"accept-language": "de-CH-1901, fr"}) + ) + + assert seen[0].language == "de-CH" + assert seen[0].accept_language == "de-CH-1901, fr" + + +# -- the validator it issued --------------------------------------------------- + + +@pytest.mark.asyncio +async def test_tags_every_representation_it_serves() -> None: + """SVC-61: a representation without a validator cannot be revalidated, only re-fetched.""" + for path in EVERY_PATH: + etag = (await call("GET", path)).headers["etag"] + assert etag.startswith('"') and etag.endswith('"'), path + assert len(etag) > 2, path + + +@pytest.mark.asyncio +async def test_gives_each_variant_a_tag_of_its_own() -> None: + """SVC-61 again: two languages sharing one tag is the one thing a validator must not do.""" + english = await localized("en") + french = await localized("fr") + + assert english.headers["etag"] != french.headers["etag"] + + +@pytest.mark.asyncio +async def test_gives_one_representation_one_tag() -> None: + first = await call("GET", "/odp/offerings") + second = await call("GET", "/odp/offerings") + + assert first.headers["etag"] == second.headers["etag"] + + +@pytest.mark.asyncio +async def test_distinguishes_terse_from_full() -> None: + """Two representations of one resource are two variants, and carry two validators. + + A Terse Offering omits the Actions a Full one carries, so the two differ by exactly the members + the representation rules say they should. + """ + actionable = Offering.model_validate( + { + "id": "p0", + "name": "P0", + "odp_version": "1.0", + "actions": [ + { + "authentication": "not-required", + "id": "buy", + "rel": "purchase", + "http": {"href": "https://plants.example/checkout", "method": "POST"}, + } + ], + } + ) + built = ServiceBuilder("Plants", "A plant store.", "en", "/odp").build( + StaticCatalog(StaticCatalogOptions(offerings=(actionable,))) + ) + terse = await built.handle(Request(method="GET", path="/odp/offerings/p0")) + full = await built.handle( + Request(method="GET", path="/odp/offerings/p0", query="representation=full") + ) + + assert b"actions" not in terse.body + assert b"actions" in full.body + assert terse.headers["etag"] != full.headers["etag"] + + +@pytest.mark.asyncio +async def test_answers_a_conditional_request_that_still_matches() -> None: + """PAG-31: a validator that still matches means the Agent already holds this representation.""" + built = service() + first = await built.handle(Request(method="GET", path="/odp/offerings")) + etag = first.headers["etag"] + + second = await built.handle( + Request(method="GET", path="/odp/offerings", headers={"if-none-match": etag}) + ) + + assert second.status == 304 + assert second.body == b"" + assert second.headers["etag"] == etag + assert "content-type" not in second.headers + + +@pytest.mark.asyncio +async def test_reads_every_form_a_conditional_field_takes() -> None: + """RFC 9110 compares `If-None-Match` weakly, so `W/"x"` matches the strong `"x"` served here.""" + built = service() + etag = (await built.handle(Request(method="GET", path="/odp/offerings"))).headers["etag"] + + for value in (etag, f"W/{etag}", f'"other", {etag}', "*"): + reply = await built.handle( + Request(method="GET", path="/odp/offerings", headers={"if-none-match": value}) + ) + assert reply.status == 304, value + + +@pytest.mark.asyncio +async def test_serves_a_conditional_request_that_no_longer_matches() -> None: + reply = await call("GET", "/odp/offerings", headers={"if-none-match": '"something-else"'}) + + assert reply.status == 200 + assert reply.body + + +@pytest.mark.asyncio +async def test_refuses_a_matched_precondition_on_a_method_that_is_not_a_retrieval() -> None: + """RFC 9110 13.1.2: a matched `If-None-Match` is 304 for GET and HEAD, 412 for anything else.""" + + class Searchable(StaticCatalog): + def operations(self) -> list[Operation]: + return [*super().operations(), Operation.SEARCH_OFFERINGS] + + async def search_offerings( + self, query: OfferingSearchRequest, request: CatalogRequest + ) -> OfferingPage[Offering]: + del query + return await self.list_offerings(request) + + built = ServiceBuilder("Plants", "A plant store.", "en", "/odp").build( + Searchable(StaticCatalogOptions(offerings=(_offering("p0"),))) + ) + reply = await built.handle( + Request( + method="POST", + path="/odp/offerings/search", + body=b'{"odp_version":"1.0","query":"plant"}', + headers={"content-type": MEDIA_TYPE, "if-none-match": "*"}, + ) + ) + + assert reply.status == 412 + + +# -- the methods each resource answers ----------------------------------------- + + +@pytest.mark.asyncio +async def test_answers_head_wherever_it_answers_get() -> None: + """RFC 9110 9.3.2: HEAD is GET without the body, so refusing it breaks every cache probe.""" + for path in EVERY_PATH: + get = await call("GET", path) + head = await call("HEAD", path) + + assert head.status == 200, path + assert head.body == b"", path + assert head.headers["etag"] == get.headers["etag"], path + + +@pytest.mark.asyncio +async def test_answers_a_conditional_head() -> None: + built = service() + etag = (await built.handle(Request(method="GET", path="/odp/offerings"))).headers["etag"] + reply = await built.handle( + Request(method="HEAD", path="/odp/offerings", headers={"if-none-match": etag}) + ) + + assert reply.status == 304 + assert reply.body == b"" + + +@pytest.mark.asyncio +async def test_names_the_methods_a_refused_one_should_have_been() -> None: + """RFC 9110 15.5.6: a 405 that does not say what is allowed leaves the caller guessing.""" + for method, path, allow in ( + ("POST", "/.well-known/odp", "GET, HEAD"), + ("DELETE", "/odp/offerings", "GET, HEAD"), + ("PUT", "/odp/collections", "GET, HEAD"), + ): + reply = await call(method, path) + assert reply.status == 405, path + assert reply.headers["allow"] == allow, path + + +@pytest.mark.asyncio +async def test_says_a_search_path_wants_post_rather_than_that_it_does_not_exist() -> None: + """A reserved path names an operation, not a resource. + + Reading `/offerings/search` as a request for an Offering called "search" answers 404, which + tells an Agent the operation is unavailable rather than that it uses a different method. + """ + for path in ("/odp/offerings/search", "/odp/collections/search"): + reply = await call("GET", path) + assert reply.status == 405, path + assert reply.headers["allow"] == "POST", path + + +# -- what it refuses before a Catalog is asked ---------------------------------- + + +@pytest.mark.asyncio +async def test_refuses_a_repeated_representation() -> None: + """SVC-73: collapsing repeats honoured whichever copy came last. + + `representation=terse&representation=full` then served a Full Representation to a request that + had also asked for a Terse one. + """ + for query in ( + "representation=terse&representation=full", + "representation=full&representation=full", + "limit=1&limit=2", + "cursor=a&cursor=b", + ): + reply = await call("GET", "/odp/offerings", query=query) + assert reply.status == 400, query + assert b"must not be repeated" in reply.body, query + + +@pytest.mark.asyncio +async def test_refuses_a_representation_or_limit_it_cannot_honour() -> None: + for query in ("representation=sideways", "limit=101", "limit=-1", "limit=many"): + assert (await call("GET", "/odp/offerings", query=query)).status == 400, query + + for query in ("representation=full", "representation=terse", "limit=100", "limit=0"): + assert (await call("GET", "/odp/offerings", query=query)).status == 200, query + + +@pytest.mark.asyncio +async def test_refuses_an_identifier_that_names_no_resource_it_could_hold() -> None: + """SVC-66 substitutes an identifier verbatim, so a segment that is not an LRI names nothing. + + Letting one through makes the Catalog decide what `..` means, which is not a question a Catalog + should have to answer. + """ + for path in ( + "/odp/offerings/..", + "/odp/offerings/a%2Fb", + "/odp/collections/..", + "/odp/collections/.", + "/odp/collections//offerings", + "/odp/collections/a b/offerings", + ): + reply = await call("GET", path) + assert reply.status == 400, path + assert b"identifier is invalid" in reply.body, path + + +@pytest.mark.asyncio +async def test_refuses_an_accept_that_excludes_odp() -> None: + """MED-04, and RFC 9110 12.4.2: naming ODP with zero weight excludes it just as surely.""" + for accept in ("text/html", "application/json", "application/odp+json;q=0", "*/*;q=0"): + reply = await call("GET", "/odp/offerings", headers={"accept": accept}) + assert reply.status == 406, accept + assert reply.headers["content-type"] == PROBLEM_MEDIA_TYPE, accept + + +@pytest.mark.asyncio +async def test_serves_an_accept_that_allows_odp() -> None: + """MED-03: a wildcard media range covers this media type, and `application/*` is one.""" + for accept in ( + "application/odp+json", + "application/*", + "*/*", + "text/html;q=0.9, application/odp+json", + "application/odp+json;q=0.5", + ): + assert (await call("GET", "/odp/offerings", headers={"accept": accept})).status == 200, ( + accept + ) + + +@pytest.mark.asyncio +async def test_refuses_a_search_body_it_cannot_read() -> None: + """MED-06: a request body for an ODP operation carries the ODP media type.""" + + class Searchable(StaticCatalog): + def operations(self) -> list[Operation]: + return [*super().operations(), Operation.SEARCH_COLLECTIONS] + + async def search_collections( + self, query: CollectionSearchRequest, request: CatalogRequest + ) -> Page[Collection]: + del query + return await self.list_collections(request) + + built = ServiceBuilder("Plants", "A plant store.", "en", "/odp").build( + Searchable( + StaticCatalogOptions( + collections=( + Collection.model_validate( + {"id": "plants", "name": "Plants", "odp_version": "1.0"} + ), + ), + offerings=(_offering("p0"),), + ) + ) + ) + body = b'{"odp_version":"1.0","query":"plants"}' + + for content_type, status in ( + (MEDIA_TYPE, 200), + ("application/json", 415), + ("", 415), + ("text/plain", 415), + ): + reply = await built.handle( + Request( + method="POST", + path="/odp/collections/search", + body=body, + headers={"content-type": content_type}, + ) + ) + assert reply.status == status, content_type + + oversized = await built.handle( + Request( + method="POST", + path="/odp/collections/search", + body=b" " * 70_000, + headers={"content-type": MEDIA_TYPE}, + ) + ) + assert oversized.status == 413 + + +@pytest.mark.asyncio +async def test_refuses_an_operation_it_does_not_advertise() -> None: + """ROLE-01: an unadvertised operation does not exist, whatever the path looks like.""" + reply = await call("GET", "/odp/collections") + assert reply.status == 200 + + without = ServiceBuilder("Plants", "A plant store.", "en", "/odp").build( + StaticCatalog(StaticCatalogOptions(offerings=(_offering("p0"),))) + ) + assert (await without.handle(Request(method="GET", path="/odp/collections"))).status == 404 + + +@pytest.mark.asyncio +async def test_refuses_a_path_outside_its_endpoint_base() -> None: + for path in ("/elsewhere", "/odp/baskets", "/odp"): + assert (await call("GET", path)).status == 404, path + + +# -- the problems it reports ----------------------------------------------------- + + +@pytest.mark.asyncio +async def test_reports_a_problem_whose_title_names_the_type_not_the_occurrence() -> None: + """RFC 9457: a title summarises the problem *type* and does not vary between occurrences. + + Setting it to the detail made every 404 a different "type" to anything grouping by title. + """ + missing = json.loads((await call("GET", "/odp/offerings/absent")).body) + invalid = json.loads((await call("GET", "/odp/offerings/..")).body) + + assert missing["title"] == "Not found" + assert missing["detail"] == "Offering not found" + assert missing["type"] == "https://offeringprotocol.org/problems/not-found" + assert invalid["title"] == "Invalid request" + assert invalid["title"] != invalid["detail"] + + +@pytest.mark.asyncio +async def test_reports_every_problem_with_the_problem_media_type() -> None: + for method, path in (("GET", "/odp/offerings/absent"), ("DELETE", "/odp/offerings")): + reply = await call(method, path) + assert reply.headers["content-type"] == PROBLEM_MEDIA_TYPE, path + assert json.loads(reply.body)["status"] == reply.status, path + + +# -- pagination -------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_keeps_a_continuation_in_the_variant_that_produced_it() -> None: + """A page served in one language is a different page from the same offsets served in another. + + The cursor carries the selected language for the same reason it carries the representation: a + continuation must not quietly change variant part-way through a sequence. + """ + built = service(localizations=TAGS, count=5) + first = await built.handle( + Request( + method="GET", + path="/odp/offerings", + query="limit=2", + headers={"accept-language": "fr"}, + ) + ) + query = json.loads(first.body)["next"].split("?", 1)[1] + + same = await built.handle( + Request(method="GET", path="/odp/offerings", query=query, headers={"accept-language": "fr"}) + ) + other = await built.handle( + Request(method="GET", path="/odp/offerings", query=query, headers={"accept-language": "en"}) + ) + + assert same.status == 200 + assert other.status == 410 + + +@pytest.mark.asyncio +async def test_walks_a_sequence_to_its_end() -> None: + built = service(count=5) + seen: list[str] = [] + query = "limit=2" + while True: + reply = await built.handle(Request(method="GET", path="/odp/offerings", query=query)) + page = json.loads(reply.body) + seen.extend(item["id"] for item in page["items"]) + if not page.get("next"): + break + query = page["next"].split("?", 1)[1] + + assert seen == ["p0", "p1", "p2", "p3", "p4"] + + +@pytest.mark.asyncio +async def test_refuses_a_continuation_it_did_not_issue() -> None: + for cursor in ("nonsense", "a.b", "YQ.Yg"): + reply = await call("GET", "/odp/offerings", query=f"cursor={cursor}") + assert reply.status == 410, cursor + assert json.loads(reply.body)["code"] == "CONTINUATION_UNAVAILABLE", cursor + + +@pytest.mark.asyncio +async def test_reads_a_media_range_carrying_parameters_other_than_weight() -> None: + """A parameter that is not `q` says nothing about whether the range is acceptable.""" + for accept in ("application/odp+json;charset=utf-8", "application/odp+json;version=1;q=0.8"): + assert (await call("GET", "/odp/offerings", headers={"accept": accept})).status == 200, ( + accept + ) + + +@pytest.mark.asyncio +async def test_serves_the_first_variant_the_residual_range_does_not_exclude() -> None: + """RFC 9110 12.5.4: `*` matches every tag no other range matched, minus the ones refused. + + So `*` beside `en;q=0` asks for anything except English, and the Service answers with the first + localization that survives rather than with its default. A refused range covers its subtags the + way a basic range does, so refusing `en` refuses `en-GB` with it. + """ + assert (await localized("*, en;q=0")).headers["content-language"] == "fr" + assert (await localized("*, fr;q=0, de-CH;q=0")).headers["content-language"] == "en" + assert (await localized("*, en;q=0, fr;q=0")).headers["content-language"] == "de-CH" + + +@pytest.mark.asyncio +async def test_falls_back_when_the_residual_range_excludes_everything() -> None: + """Refusing every variant still is not a reason to refuse the request (SVC-59).""" + refused = ", ".join(f"{tag};q=0" for tag in TAGS) + reply = await localized(f"*, {refused}") + + assert reply.status == 200 + assert reply.headers["content-language"] == "en" From ae7961a94b9e3028046e724b66825ecf397295ea Mon Sep 17 00:00:00 2001 From: Nas Kavian Date: Tue, 22 Sep 2026 00:01:33 -0400 Subject: [PATCH 2/4] fix(sdk): correct discovery and response negotiation --- README.md | 11 ++- scripts/conformance_adapter.py | 3 + src/offering_protocol/agent/client.py | 30 +++--- src/offering_protocol/core/validation.py | 35 +++---- src/offering_protocol/directory/client.py | 5 +- src/offering_protocol/service/service.py | 40 ++++++-- tests/test_agent.py | 3 - tests/test_agent_conformance.py | 73 ++++++++++++++ tests/test_core_conformance.py | 30 ++++-- tests/test_directory_conformance.py | 24 +++++ tests/test_service_conformance.py | 112 +++++++++++++++++++++- 11 files changed, 302 insertions(+), 64 deletions(-) diff --git a/README.md b/README.md index 1b71d4f..dc1ac59 100644 --- a/README.md +++ b/README.md @@ -341,10 +341,19 @@ its Actions can advertise enrollment, payment, and trust protocols, but ODP does credentials, invoke Actions, submit payments, or implement trust protocols. Applications compose the appropriate protocol clients around an Action resolved through ODP. -`parse_service_document` is the strict current-version Service parser. Agent inspection and +`parse_service_document` validates Service metadata against the supported ODP major version. +Compatible minor versions such as `1.7` are accepted without rewriting the received version; +SDK-generated documents use `1.0`. Agent inspection and Directory results filter unrecognized enrollment, payment, and trust descriptors while retaining strict validation for recognized descriptors. +Individual Offering and Collection GETs default to full representations; list and search operations +default to terse items. A Catalog receives the requested language in `CatalogRequest.language` and +declares the language it actually returns on each resource. Static catalogs do not translate content. + +Refinement parsing detects duplicate JSON values without guessing the type of a string. Comparing +decimal or date-time strings by their meaning requires the referenced Filter Definition. + ## Errors and validation Each role exposes typed errors: diff --git a/scripts/conformance_adapter.py b/scripts/conformance_adapter.py index bf1fb35..beaebea 100755 --- a/scripts/conformance_adapter.py +++ b/scripts/conformance_adapter.py @@ -242,6 +242,9 @@ async def evaluate_errors_limits(case: dict[str, Any]) -> bool | None: async def evaluate_case(subject: str, case: dict[str, Any], role: str) -> bool | None: + if subject == "protocol-version": + document = {"odp_version": case["received"], "id": "item", "name": "Item"} + return succeeds(lambda: parse_offering(json.dumps(document))) == case["compatible"] if subject == "local-identifier": return is_local_resource_identifier(case["value"]) == case["valid"] if subject == "identity-comparison": diff --git a/src/offering_protocol/agent/client.py b/src/offering_protocol/agent/client.py index 12a1c54..d47af3a 100644 --- a/src/offering_protocol/agent/client.py +++ b/src/offering_protocol/agent/client.py @@ -189,10 +189,7 @@ async def list_collections( self, representation: Representation = Representation.TERSE, limit: int = 0 ) -> Page[Collection]: body = await self._get_page(Operation.LIST_COLLECTIONS, None, representation, limit) - page = parse_collection_page(body) - for item in page.items: - parse_collection(_encode(item)) - return page + return parse_collection_page(body) async def get_collection(self, identifier: str) -> Collection: body = await self._get_page(Operation.GET_COLLECTION, identifier, Representation.FULL, 0) @@ -206,10 +203,7 @@ async def search_collections( body = await self._post_search( Operation.SEARCH_COLLECTIONS, request.to_dict(), representation ) - page = parse_collection_page(body) - for item in page.items: - parse_collection(_encode(item)) - return page + return parse_collection_page(body) async def list_offerings( self, representation: Representation = Representation.TERSE, limit: int = 0 @@ -498,7 +492,11 @@ async def _supporting_json( conditional["if-none-match"] = cached.etag if cached.last_modified: conditional["if-modified-since"] = cached.last_modified + visited: set[str] = set() for redirects in range(_MAXIMUM_REDIRECTS + 1): + if current in visited: + raise AgentError("ODP supporting document contains a redirect loop") + visited.add(current) try: response = await self._supporting_transport.send( HttpRequest("GET", current, {"accept": accept, **conditional}) @@ -511,9 +509,14 @@ async def _supporting_json( location = response.headers.get("location") if location is None: raise AgentError("ODP supporting document redirect omitted Location") - current = urljoin(current, location) - if not _is_https_url(current): + target_url = urljoin(current, location) + if not _is_https_url(target_url): raise AgentError("ODP supporting document redirect must use HTTPS") + if derive_service_origin(target_url) != derive_service_origin(current): + raise AgentError( + "ODP supporting document redirect must remain on the same origin" + ) + current = target_url continue if response.status == 304: if cached is None: @@ -612,13 +615,6 @@ def parse_problem_response(data: bytes | str, status: int) -> ProblemDetails: return parse_problem_response_strict(_agent_body(data, "problem"), status) -def _encode(value: object) -> bytes: - if not hasattr(value, "model_dump_json"): - raise TypeError("ODP model is not serializable") - encoded = value.model_dump_json(by_alias=True, exclude_unset=True) - return cast(str, encoded).encode() - - def _append_query(target: str, values: dict[str, str]) -> str: parts = urlsplit(target) query = dict(parse_qsl(parts.query, keep_blank_values=True)) diff --git a/src/offering_protocol/core/validation.py b/src/offering_protocol/core/validation.py index 7a10a50..fdec898 100644 --- a/src/offering_protocol/core/validation.py +++ b/src/offering_protocol/core/validation.py @@ -5,6 +5,7 @@ import json import re from dataclasses import dataclass, field +from decimal import Decimal from functools import lru_cache from importlib.resources import files from typing import Any, TypeVar @@ -17,6 +18,7 @@ from referencing import Registry, Resource from offering_protocol.core.models import ( + VERSION, Collection, CollectionSearchRequest, FilterDefinition, @@ -715,7 +717,13 @@ def validate_value(value: object, schema_name: str, document_type: str) -> None: validator = _validators().get(schema_name) if validator is None: raise RuntimeError(f"missing bundled schema {schema_name}") - issues = [_schema_issue(error) for error in validator.iter_errors(value)] + validation_value = value + if isinstance(value, dict): + version = value.get("odp_version") + if isinstance(version, str) and re.fullmatch(r"1\.(0|[1-9][0-9]*)", version): + # The bundled schemas describe 1.0; compatible minor versions use the same rules. + validation_value = {**value, "odp_version": VERSION} + issues = [_schema_issue(error) for error in validator.iter_errors(validation_value)] if issues: issues.sort(key=lambda issue: (issue.path, issue.keyword, issue.message)) raise OdpValidationError(document_type, issues) @@ -889,9 +897,6 @@ def _is_language_tag(value: str) -> bool: return not in_extension or len(subtags[-1]) > 1 -_DECIMAL = re.compile(r"^-?(?:0|[1-9][0-9]*)(?:\.[0-9]+)?$") - - def _compare_decimals(left: str, right: str) -> int: """Orders two ODP monetary values, which are decimal strings rather than JSON numbers (OFR-48). @@ -916,25 +921,15 @@ def _split_decimal(value: str) -> tuple[str, str]: return whole.lstrip("0"), fraction.rstrip("0") -def _bucket_key(value: object) -> str: - """Compares two bucket values the way the referenced Filter Definition would. - - A response does not carry its Filter Definitions, so the type behind a JSON string -- `string`, - `decimal`, `date` or `date-time` -- is not known here. Every one of those compares two strings - exactly except `decimal`, whose equality is numeric, so a string that can only be a decimal is - reduced to one spelling per value. JSON numbers compare numerically. - """ +def _bucket_key(value: object) -> tuple[str, object]: + """Check structural duplicates without guessing the type of a string-valued Filter.""" if isinstance(value, bool): - return f"b{value}" + return "boolean", value if isinstance(value, int | float): - return f"n{float(value)}" + return "number", Decimal(str(value)) if isinstance(value, str): - if _DECIMAL.match(value): - whole, fraction = _split_decimal(value.removeprefix("-")) - sign = "-" if value.startswith("-") and (whole or fraction) else "" - return f"d{sign}{whole}.{fraction}" - return f"s{value}" - return f"o{json.dumps(value, separators=(',', ':'), sort_keys=True)}" + return "string", value + return "other", json.dumps(value, separators=(",", ":"), sort_keys=True) def _agent_body(data: bytes | str, kind: str) -> str: diff --git a/src/offering_protocol/directory/client.py b/src/offering_protocol/directory/client.py index f353474..d523fef 100644 --- a/src/offering_protocol/directory/client.py +++ b/src/offering_protocol/directory/client.py @@ -248,11 +248,11 @@ def _parse_search_page(body: bytes) -> SearchPage: raise DirectoryError("invalid Directory response: search page has no items") if len(raw["items"]) > 100: raise DirectoryError("Directory search page exceeds 100 Services") - items: list[dict[str, object]] = [] + items: list[DirectoryService] = [] issues: list[ServiceIssue] = [] for index, entry in enumerate(raw["items"]): try: - items.append(_read_service(entry)) + items.append(DirectoryService.model_validate(_read_service(entry))) except (ModelValidationError, OdpValidationError, ReferenceError, DirectoryError) as error: issues.append(ServiceIssue(index=index, message=str(error))) try: @@ -292,6 +292,7 @@ def _read_service(entry: object) -> dict[str, object]: item.pop(member, None) if document.protocols is not None: item["protocols"] = document.protocols.model_dump(mode="json", exclude_defaults=True) + item["operations"] = [operation.model_dump(mode="json") for operation in document.operations] return item diff --git a/src/offering_protocol/service/service.py b/src/offering_protocol/service/service.py index ec6dfd9..a60be1d 100644 --- a/src/offering_protocol/service/service.py +++ b/src/offering_protocol/service/service.py @@ -281,7 +281,7 @@ async def _handle(self, request: Request) -> Response: self._document.language, list(self._document.localizations), ) - exchange = _Exchange(headers=headers, language=language, method=method) + exchange = _Exchange(headers=headers, language=self._document.language, method=method) if request.path == "/.well-known/odp": _require_method(effective, ("GET",)) return _json_response(self._document, _MAXIMUM_DOCUMENT_BYTES, exchange) @@ -295,7 +295,7 @@ async def _handle(self, request: Request) -> Response: item.name for item in self._document.operations }: raise RequestError(404, "NOT_FOUND", "ODP operation is not supported") - catalog_request = _catalog_request(request, headers, language) + catalog_request = _catalog_request(request, headers, language, operation) if (effective, path) == ("GET", "/offerings"): offering_page = await self._catalog.list_offerings(catalog_request) return _json_response( @@ -370,7 +370,9 @@ async def _get_path(self, path: str, request: CatalogRequest, exchange: _Exchang raise RequestError(404, "NOT_FOUND", "ODP resource not found") -def _catalog_request(request: Request, headers: dict[str, str], language: str) -> CatalogRequest: +def _catalog_request( + request: Request, headers: dict[str, str], language: str, operation: Operation | None +) -> CatalogRequest: parameters = parse_qsl(request.query, keep_blank_values=True) # SVC-73: a repeated `representation` is rejected rather than resolved. Collapsing repeats into # a dict silently honoured whichever copy came last, so `representation=terse&representation= @@ -381,7 +383,10 @@ def _catalog_request(request: Request, headers: dict[str, str], language: str) - raise RequestError(400, "INVALID_REQUEST", f"{name} must not be repeated") values[name] = value try: - representation = Representation(values.get("representation", "terse")) + default = ( + "full" if operation in {Operation.GET_COLLECTION, Operation.GET_OFFERING} else "terse" + ) + representation = Representation(values.get("representation", default)) limit = int(values.get("limit", "0")) except ValueError as error: raise RequestError(400, "INVALID_REQUEST", "query parameter is invalid") from error @@ -434,9 +439,20 @@ def _json_response(value: object, maximum_bytes: int, exchange: _Exchange) -> Re body = _encode(value) if len(body) > maximum_bytes: raise ServiceError("response body is too large") - etag = _entity_tag(exchange.language, body) + if isinstance(value, Page): + language = ( + ", ".join( + dict.fromkeys( + getattr(item, "language", "") or exchange.language for item in value.items + ) + ) + or exchange.language + ) + else: + language = getattr(value, "language", "") or exchange.language + etag = _entity_tag(language, body) headers = { - "content-language": exchange.language, + "content-language": language, "content-type": MEDIA_TYPE, "etag": etag, "vary": "Accept, Accept-Language", @@ -554,12 +570,20 @@ def _require_accept(value: str | None) -> None: """ if value is None: return + specificity = -1 + quality = 0.0 for entry in value.split(","): media_type = entry.split(";", 1)[0].strip().lower() if media_type not in {"*/*", "application/*", MEDIA_TYPE}: continue - if _quality_of(entry) > 0: - return + precision = {"*/*": 0, "application/*": 1, MEDIA_TYPE: 2}[media_type] + weight = _quality_of(entry) + if precision > specificity: + specificity, quality = precision, weight + elif precision == specificity: + quality = max(quality, weight) + if quality > 0: + return raise RequestError(406, "NOT_ACCEPTABLE", f"Accept must allow {MEDIA_TYPE}") diff --git a/tests/test_agent.py b/tests/test_agent.py index 1dc3f11..232b4e7 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -27,7 +27,6 @@ ) from offering_protocol.agent.client import ( _decode_json_object, - _encode, _expiration, _invoke_parser, ) @@ -665,8 +664,6 @@ async def test_agent_remaining_traversal_and_supporting_document_edges() -> None assert not (await client.list_collection_offerings("plants", limit=1)).items assert not await client.search_all_offerings(OfferingSearchRequest(query="plant")) - with pytest.raises(TypeError): - _encode({"not": "a model"}) with pytest.raises(AgentError): _invoke_parser(lambda _: (_ for _ in ()).throw(ValueError("bad")), b"{}") diff --git a/tests/test_agent_conformance.py b/tests/test_agent_conformance.py index af7f976..c2a134c 100644 --- a/tests/test_agent_conformance.py +++ b/tests/test_agent_conformance.py @@ -33,6 +33,7 @@ ) from offering_protocol.core import ( CapabilityLink, + CollectionSearchRequest, FilterCapabilitySource, FilterDefinition, FilterOperator, @@ -534,3 +535,75 @@ def exhausted(*args: object, **kwargs: object) -> object: with pytest.raises(AgentError, match="nested too deeply"): _decode_json_object(b"[[[]]]") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("search", [False, True]) +async def test_collection_page_inherits_version_without_requiring_it_on_items(search: bool) -> None: + document = json.loads(SERVICE_DOCUMENT) + operation = "search-collections" if search else "list-collections" + document["operations"].append({"name": operation, "authentication": "not-required"}) + transport = QueueTransport( + response(json.dumps(document)), + response('{"odp_version":"1.7","items":[{"id":"plants","name":"Plants"}]}'), + ) + client = ServiceClient("https://store.example", transport=transport) + page = ( + await client.search_collections(CollectionSearchRequest()) + if search + else await client.list_collections() + ) + assert page.odp_version == "1.7" + assert page.items[0].id == "plants" + assert "odp_version" not in page.items[0].model_fields_set + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "location", + [ + "https://other.example/schema.json", + "https://schemas.example:444/schema.json", + "/schema.json", + ], +) +async def test_supporting_redirects_refuse_origin_changes_and_loops(location: str) -> None: + offering = { + "odp_version": "1.0", + "id": "item", + "name": "Item", + "schema": {"url": "https://schemas.example/schema.json"}, + "attributes": {"colour": "green"}, + } + supporting = QueueTransport(response("", status=302, headers={"location": location})) + client = ServiceClient( + "https://store.example", + transport=QueueTransport(response(SERVICE_DOCUMENT), response(json.dumps(offering))), + supporting_transport=supporting, + ) + details = await client.get_offering_details("item") + assert len(supporting.requests) == 1 + assert len(details.issues) == 1 + assert "redirect" in details.issues[0].message + assert details.offering.name == "Item" + assert not details.offering.attributes + + +@pytest.mark.asyncio +async def test_supporting_redirects_allow_explicit_default_port() -> None: + supporting = QueueTransport( + response("", status=302, headers={"location": "https://schemas.example:443/final.json"}), + response("{}", content_type="application/json"), + ) + client = ServiceClient("https://store.example", supporting_transport=supporting) + assert ( + await client._supporting_json( + "https://schemas.example/start.json", + "schema", + "application/json", + {"application/json"}, + 100, + ) + == {} + ) + assert len(supporting.requests) == 2 diff --git a/tests/test_core_conformance.py b/tests/test_core_conformance.py index 3d38e5e..bfd3c2c 100644 --- a/tests/test_core_conformance.py +++ b/tests/test_core_conformance.py @@ -34,6 +34,22 @@ def amend(base: dict[str, Any], **changes: Any) -> str: return json.dumps({**base, **changes}) +@pytest.mark.parametrize("version", ["1.0", "1.1", "1.7", "1.9999999999999999999999999999"]) +def test_accepts_compatible_minor_versions_without_rewriting_models(version: str) -> None: + from offering_protocol.core import validate_value + + document = {**OFFERING, "odp_version": version} + validate_value(document, "offering.schema.json", "Offering") + assert document["odp_version"] == version + assert parse_offering(json.dumps(document)).odp_version == version + + +@pytest.mark.parametrize("version", ["2.0", "0.9", "1", "01.0", "1.01", "1.0.0", "1.0\n", None, 1]) +def test_rejects_incompatible_or_malformed_versions(version: object) -> None: + with pytest.raises(OdpValidationError): + parse_offering(amend(OFFERING, odp_version=version)) + + def action(identifier: str) -> dict[str, Any]: return { "authentication": "not-required", @@ -238,21 +254,14 @@ def test_refuses_a_repeated_bucket_value() -> None: ) -def test_reads_two_spellings_of_one_decimal_as_one_bucket_value() -> None: - """FLT-32: decimal equality is numeric rather than lexical. - - So `1.0` and `1.00` name one value, and a group offering both hands a caller two counts for one - candidate with no way to choose between them. - """ +def test_keeps_distinct_strings_without_a_filter_definition() -> None: for values in ( ({"value": "1.0", "count": 4}, {"value": "1.00", "count": 2}), ({"value": "0", "count": 4}, {"value": "0.0", "count": 2}), ({"value": "12", "count": 4}, {"value": "12.000", "count": 2}), ({"value": "-1.5", "count": 4}, {"value": "-1.50", "count": 2}), ): - assert_rejected_for( - _page(_group("weight", *values)), parse_offering_page, "unique-bucket-value" - ) + assert parse_offering_page(_page(_group("sku", *values))) def test_keeps_bucket_values_that_differ_apart() -> None: @@ -368,4 +377,5 @@ def test_compares_any_bucket_value_the_model_can_hold() -> None: assert _bucket_key([1, 2]) != _bucket_key([2, 1]) assert _bucket_key(None) != _bucket_key("null") assert _bucket_key(True) != _bucket_key(1) - assert _bucket_key("1.0") == _bucket_key("1.00") + assert _bucket_key("1.0") != _bucket_key("1.00") + assert _bucket_key(9007199254740992) != _bucket_key(9007199254740993) diff --git a/tests/test_directory_conformance.py b/tests/test_directory_conformance.py index 377a94a..9c54bb1 100644 --- a/tests/test_directory_conformance.py +++ b/tests/test_directory_conformance.py @@ -74,6 +74,30 @@ async def test_reads_a_conformant_record() -> None: assert page.items[0].name == "Plants" +@pytest.mark.asyncio +async def test_normalizes_unknown_operations_before_model_decoding() -> None: + page = await read( + SERVICE, + amend( + operations=[ + *BASELINE_OPERATIONS, + {"name": "future-operation", "authentication": "not-required"}, + ] + ), + ) + assert len(page.items) == 2 + assert not page.issues + assert len(page.items[1].operations) == 2 + + +@pytest.mark.asyncio +async def test_model_decode_failure_is_isolated_to_its_record() -> None: + page = await read(SERVICE, amend(website_url=None), SERVICE) + assert len(page.items) == 2 + assert len(page.issues) == 1 + assert page.issues[0].index == 1 + + @pytest.mark.asyncio async def test_refuses_an_origin_that_is_not_a_canonical_https_origin() -> None: """IDN-01: a Service is identified by its canonical origin, so two spellings are not one.""" diff --git a/tests/test_service_conformance.py b/tests/test_service_conformance.py index 91c95af..7d42e6d 100644 --- a/tests/test_service_conformance.py +++ b/tests/test_service_conformance.py @@ -69,8 +69,30 @@ async def call(method: str, path: str, **kwargs: object) -> Response: return await service().handle(Request(method=method, path=path, **kwargs)) # type: ignore[arg-type] -async def localized(accept_language: str, path: str = "/.well-known/odp") -> Response: - return await service(localizations=TAGS).handle( +async def localized(accept_language: str, path: str = "/odp/offerings") -> Response: + class TranslatedCatalog(StaticCatalog): + async def list_offerings(self, request: CatalogRequest) -> OfferingPage[Offering]: + page = await super().list_offerings(request) + return page.model_copy( + update={ + "items": [ + item.model_copy( + update={ + "language": request.language, + "name": f"Plant ({request.language})", + } + ) + for item in page.items + ] + } + ) + + built = ( + ServiceBuilder("Plants", "A plant store.", "en", "/odp") + .localizations(TAGS) + .build(TranslatedCatalog(StaticCatalogOptions(offerings=(_offering("p0"),)))) + ) + return await built.handle( Request(method="GET", path=path, headers={"accept-language": accept_language}) ) @@ -85,6 +107,86 @@ async def localized(accept_language: str, path: str = "/.well-known/odp") -> Res ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + "accept", + ["*/*, application/odp+json;q=0", "application/odp+json;q=0, */*", "application/*;q=0, */*"], +) +async def test_specific_media_exclusion_overrides_wildcard(accept: str) -> None: + result = await service().handle(Request("GET", "/.well-known/odp", headers={"accept": accept})) + assert result.status == 406 + + +@pytest.mark.asyncio +async def test_explicit_media_acceptance_overrides_wildcard_refusal() -> None: + result = await service().handle( + Request( + "GET", "/.well-known/odp", headers={"accept": "application/*;q=0, application/odp+json"} + ) + ) + assert result.status == 200 + + +@pytest.mark.asyncio +async def test_repeated_media_ranges_use_highest_quality() -> None: + result = await service().handle( + Request( + "GET", + "/.well-known/odp", + headers={"accept": "application/odp+json;q=0, application/odp+json;q=1"}, + ) + ) + assert result.status == 200 + + +@pytest.mark.asyncio +async def test_static_metadata_is_not_relabelled_as_a_translation() -> None: + for path in EVERY_PATH: + built = service(localizations=TAGS) + english = await built.handle(Request("GET", path, headers={"accept-language": "en"})) + french = await built.handle(Request("GET", path, headers={"accept-language": "fr"})) + assert french.headers["content-language"] == "en" + assert english.body == french.body + assert english.headers["etag"] == french.headers["etag"] + + +@pytest.mark.asyncio +async def test_individual_gets_default_to_full_and_lists_to_terse() -> None: + seen: list[CatalogRequest] = [] + + class RecordingCatalog(StaticCatalog): + async def get_collection( + self, identifier: str, request: CatalogRequest + ) -> Collection | None: + seen.append(request) + return await super().get_collection(identifier, request) + + built = ServiceBuilder("Plants", "Plants", "en", "/odp").build( + RecordingCatalog( + StaticCatalogOptions( + collections=(Collection(id="plants", name="Plants", odp_version="1.0"),) + ) + ) + ) + await built.handle(Request("GET", "/odp/collections/plants")) + assert seen[0].representation.value == "full" + for path in ("/odp/offerings/p0", "/odp/collections/plants"): + built = service() + default = await built.handle(Request("GET", path)) + full = await built.handle(Request("GET", path, query="representation=full")) + assert default.status == full.status == 200 + assert default.body == full.body + terse = await built.handle(Request("GET", path, query="representation=terse")) + assert terse.status == 200 + assert json.loads(terse.body)["odp_version"] == "1.0" + for path in ("/odp/offerings", "/odp/collections", "/odp/collections/plants/offerings"): + built = service() + default = await built.handle(Request("GET", path)) + terse = await built.handle(Request("GET", path, query="representation=terse")) + assert default.status == terse.status == 200 + assert default.body == terse.body + + # -- which variant it served --------------------------------------------------- @@ -216,12 +318,16 @@ async def test_distinguishes_terse_from_full() -> None: built = ServiceBuilder("Plants", "A plant store.", "en", "/odp").build( StaticCatalog(StaticCatalogOptions(offerings=(actionable,))) ) - terse = await built.handle(Request(method="GET", path="/odp/offerings/p0")) + terse = await built.handle( + Request(method="GET", path="/odp/offerings/p0", query="representation=terse") + ) full = await built.handle( Request(method="GET", path="/odp/offerings/p0", query="representation=full") ) + default = await built.handle(Request(method="GET", path="/odp/offerings/p0")) assert b"actions" not in terse.body + assert default.body == full.body assert b"actions" in full.body assert terse.headers["etag"] != full.headers["etag"] From c6f967f47bb6e8a7cbd2cc52d4f00ce41161b484 Mon Sep 17 00:00:00 2001 From: Nas Kavian Date: Tue, 22 Sep 2026 00:56:12 -0400 Subject: [PATCH 3/4] fix(sdk): isolate metadata requests and bound streamed responses --- README.md | 21 +- src/offering_protocol/agent/client.py | 48 ++- src/offering_protocol/directory/client.py | 16 +- src/offering_protocol/directory/transport.py | 41 +- src/offering_protocol/service/service.py | 5 + tests/test_agent.py | 28 +- tests/test_directory_conformance.py | 18 +- tests/test_service.py | 74 +++- tests/test_transport.py | 373 +++++++++++++++++++ 9 files changed, 571 insertions(+), 53 deletions(-) create mode 100644 tests/test_transport.py diff --git a/README.md b/README.md index dc1ac59..b3638d8 100644 --- a/README.md +++ b/README.md @@ -203,9 +203,10 @@ The Agent module also provides: - Conditional request and representation caching with injectable `Cache` and `Transport` protocols. Default fallback cache lifetimes are four hours for Service documents, one hour for Collections, -and five minutes for Offerings. HTTP cache directives take precedence. Provide distinct `transport` -and `supporting_transport` instances when protocol resources and linked schemas require different -credentials or network policy. +and five minutes for Offerings. HTTP cache directives take precedence. `ServiceClient` uses a +separate anonymous transport for linked schemas and OpenAPI documents, even when its primary +`transport` has authentication configured. An explicit `supporting_transport` override must also +send these requests anonymously; it must not share the primary transport's credentials or cookies. ### Search across Services @@ -269,6 +270,12 @@ must survive process restarts or share storage across workers. A custom `Transpo asynchronous `send()` and `aclose()` methods. Caller-provided caches and transports remain owned by the caller. +`HttpRequest.maximum_response_bytes` gives a custom transport the response budget. Enforce it while +reading, rather than buffering the complete response first. The built-in transport closes responses +on overflow, read failure, and cancellation. A successful response exceeding its budget raises +`TransportError` with `code="RESPONSE_LIMIT_EXCEEDED"`; `ServiceClient` preserves that code on +`AgentError`. Oversized error bodies are discarded while retaining the HTTP status and headers. + The built-in HTTP transport resolves and validates every destination before connecting, pins the connection to a validated public address, does not inherit proxy settings from the environment, and sends supporting-document requests without credentials. A custom transport must preserve those ODP @@ -348,7 +355,9 @@ Directory results filter unrecognized enrollment, payment, and trust descriptors strict validation for recognized descriptors. Individual Offering and Collection GETs default to full representations; list and search operations -default to terse items. A Catalog receives the requested language in `CatalogRequest.language` and +default to terse items. The Service handler writes `odp_version` on standalone resources and page +envelopes, omitting it from embedded page items without changing the Catalog's models. +A Catalog receives the requested language in `CatalogRequest.language` and declares the language it actually returns on each resource. Static catalogs do not translate content. Refinement parsing detects duplicate JSON values without guessing the type of a string. Comparing @@ -367,6 +376,10 @@ Protocol models preserve additive members in `model.additional` and round-trip t `model.to_dict()`. Parsing remains strict for normative constraints and fields that prohibit unknown members. +Directory records retain additional metadata such as branding and MCP endpoints. These are discovery +hints, not authorization or authoritative routing data. The default Agent factory uses the record's +`service_origin` and retrieves that Service's own document before making catalog requests. + Handle the narrowest error that the application can act upon and use the role's base error for the remaining failures: diff --git a/src/offering_protocol/agent/client.py b/src/offering_protocol/agent/client.py index d47af3a..f69ac2b 100644 --- a/src/offering_protocol/agent/client.py +++ b/src/offering_protocol/agent/client.py @@ -91,6 +91,10 @@ class TraversalOptions: class AgentError(RuntimeError): """Base error for Service discovery operations.""" + def __init__(self, message: str, *, code: str | None = None) -> None: + super().__init__(message) + self.code = code + class UnsupportedOperationError(AgentError): def __init__(self, operation: Operation) -> None: @@ -132,7 +136,12 @@ def __init__( self._cache_partition = cache_partition self._owns_transport = transport is None self._transport = transport or HttpxTransport(allow_local_network=allow_local_network) - self._supporting_transport = supporting_transport or self._transport + self._owns_supporting_transport = supporting_transport is None + self._supporting_transport = ( + supporting_transport + if supporting_transport is not None + else HttpxTransport(allow_local_network=allow_local_network) + ) async def __aenter__(self) -> ServiceClient: return self @@ -141,8 +150,12 @@ async def __aexit__(self, *args: object) -> None: await self.aclose() async def aclose(self) -> None: - if self._owns_transport: - await self._transport.aclose() + try: + if self._owns_transport: + await self._transport.aclose() + finally: + if self._owns_supporting_transport: + await self._supporting_transport.aclose() async def inspect(self) -> Inspection: requested_url = f"{self.service_origin}/.well-known/odp" @@ -399,7 +412,9 @@ async def _request_cached( headers["if-none-match"] = cached.etag if cached.last_modified: headers["if-modified-since"] = cached.last_modified - response, final_url = await self._request_raw(method, request_target, body, headers) + response, final_url = await self._request_raw( + method, request_target, body, headers, maximum_bytes + ) if response.status == 304: if cached is None: raise AgentError("ODP response returned 304 without a cached representation") @@ -435,7 +450,7 @@ async def _request_cached( return _FetchedResponse(response.body, final_url, Freshness.FETCHED) async def _request_raw( - self, method: str, target: str, body: bytes, conditional: dict[str, str] + self, method: str, target: str, body: bytes, conditional: dict[str, str], maximum_bytes: int ) -> tuple[HttpResponse, str]: redirect_origin = derive_service_origin(target) for redirects in range(_MAXIMUM_REDIRECTS + 1): @@ -445,9 +460,11 @@ async def _request_raw( if body: headers["content-type"] = MEDIA_TYPE try: - response = await self._transport.send(HttpRequest(method, target, headers, body)) + response = await self._transport.send( + HttpRequest(method, target, headers, body, maximum_response_bytes=maximum_bytes) + ) except TransportError as error: - raise AgentError(f"ODP Service request failed: {error}") from error + raise AgentError(f"ODP Service request failed: {error}", code=error.code) from error if response.status not in {301, 302, 303, 307, 308}: return response, target if redirects == _MAXIMUM_REDIRECTS: @@ -499,10 +516,17 @@ async def _supporting_json( visited.add(current) try: response = await self._supporting_transport.send( - HttpRequest("GET", current, {"accept": accept, **conditional}) + HttpRequest( + "GET", + current, + {"accept": accept, **conditional}, + maximum_response_bytes=maximum_bytes, + ) ) except TransportError as error: - raise AgentError(f"ODP supporting document request failed: {error}") from error + raise AgentError( + f"ODP supporting document request failed: {error}", code=error.code + ) from error if response.status in {301, 302, 303, 307, 308}: if redirects == _MAXIMUM_REDIRECTS: raise AgentError("ODP supporting document exceeded five redirects") @@ -544,7 +568,9 @@ async def _supporting_json( response.headers, ) if len(response.body) > maximum_bytes: - raise AgentError("ODP supporting document exceeds its byte limit") + raise AgentError( + "ODP supporting document exceeds its byte limit", code="RESPONSE_LIMIT_EXCEEDED" + ) content_type = response.headers.get("content-type", "").split(";", 1)[0].strip().lower() if content_type not in media_types: raise AgentError("ODP supporting document has an unsupported media type") @@ -643,7 +669,7 @@ def _consume(response: HttpResponse, maximum_bytes: int, maximum_depth: int) -> if not 200 <= response.status < 300: raise ServiceRequestError(response.status, _problem_message(response), response.headers) if len(response.body) > maximum_bytes: - raise AgentError("ODP response exceeds its byte limit") + raise AgentError("ODP response exceeds its byte limit", code="RESPONSE_LIMIT_EXCEEDED") content_type = response.headers.get("content-type", "").split(";", 1)[0].strip().lower() if content_type != MEDIA_TYPE: raise AgentError(f"ODP response must use {MEDIA_TYPE}") diff --git a/src/offering_protocol/directory/client.py b/src/offering_protocol/directory/client.py index d523fef..117c541 100644 --- a/src/offering_protocol/directory/client.py +++ b/src/offering_protocol/directory/client.py @@ -42,19 +42,6 @@ _RFC_3339 = re.compile( r"^\d{4}-\d{2}-\d{2}[Tt]\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:[Zz]|[+-]\d{2}:\d{2})$" ) -#: Service Document members a Directory result may echo that this client does not validate. -#: -#: They are dropped rather than passed through, because a caller reading them off a record has no -#: way to tell they were never checked. `http` is the one that matters most: a caller could build -#: request URLs from an `endpoint_base` a Directory made up. -_UNVERIFIED_MEMBERS = ( - "branding", - "http", - "mcp", - "odp_version", - "payment_origins", - "search_capabilities", -) class DirectoryError(RuntimeError): @@ -288,8 +275,7 @@ def _read_service(entry: object) -> dict[str, object]: document = _validate_as_service_document(item) # `protocols` is reinstated from the validated document only when something survived agent # filtering, so a block naming nothing this ODP version knows does not pass straight through. - for member in (*_UNVERIFIED_MEMBERS, "protocols"): - item.pop(member, None) + item.pop("protocols", None) if document.protocols is not None: item["protocols"] = document.protocols.model_dump(mode="json", exclude_defaults=True) item["operations"] = [operation.model_dump(mode="json") for operation in document.operations] diff --git a/src/offering_protocol/directory/transport.py b/src/offering_protocol/directory/transport.py index e13d01d..20e40eb 100644 --- a/src/offering_protocol/directory/transport.py +++ b/src/offering_protocol/directory/transport.py @@ -20,6 +20,7 @@ class HttpRequest: url: str headers: dict[str, str] = field(default_factory=dict) body: bytes = b"" + maximum_response_bytes: int = 524_288 @dataclass(frozen=True, slots=True) @@ -32,6 +33,10 @@ class HttpResponse: class TransportError(RuntimeError): """Raised when the HTTP transport cannot complete a request.""" + def __init__(self, message: str, *, code: str | None = None) -> None: + super().__init__(message) + self.code = code + class Transport(Protocol): async def send(self, request: HttpRequest) -> HttpResponse: ... @@ -49,6 +54,8 @@ def __init__( async def send(self, request: HttpRequest) -> HttpResponse: try: + if request.maximum_response_bytes <= 0: + raise ValueError("maximum_response_bytes must be positive") target, hostname, host_header = await _pinned_target( request.url, self._allow_local_network ) @@ -78,14 +85,36 @@ async def send(self, request: HttpRequest) -> HttpResponse: content=request.body, extensions={"sni_hostname": hostname}, ) - response = await client.send(outgoing, follow_redirects=False) + response = await client.send(outgoing, follow_redirects=False, stream=True) + try: + body = bytearray() + if request.method.upper() != "HEAD" and not 300 <= response.status_code < 400: + success = 200 <= response.status_code < 300 + maximum = ( + request.maximum_response_bytes + if success + else min(request.maximum_response_bytes, 16_384) + ) + async for chunk in response.aiter_bytes(): + if len(chunk) > maximum - len(body): + if success: + raise TransportError( + "HTTP response exceeds its byte limit", + code="RESPONSE_LIMIT_EXCEEDED", + ) + # Preserve the HTTP failure without exposing truncated error JSON. + body.clear() + break + body.extend(chunk) + return HttpResponse( + status=response.status_code, + headers={name.lower(): value for name, value in response.headers.items()}, + body=bytes(body), + ) + finally: + await response.aclose() except (httpx.HTTPError, OSError, ValueError) as error: raise TransportError(f"HTTP transport failed: {error}") from error - return HttpResponse( - status=response.status_code, - headers={name.lower(): value for name, value in response.headers.items()}, - body=response.content, - ) async def aclose(self) -> None: if self._client is not None: diff --git a/src/offering_protocol/service/service.py b/src/offering_protocol/service/service.py index a60be1d..7ab137c 100644 --- a/src/offering_protocol/service/service.py +++ b/src/offering_protocol/service/service.py @@ -555,6 +555,11 @@ def _problem( def _encode(value: object) -> bytes: + if isinstance(value, Page): + document = value.model_dump(mode="json", by_alias=True, exclude_unset=True) + for item in document["items"]: + item.pop("odp_version", None) + return json.dumps(document, separators=(",", ":"), ensure_ascii=False).encode() if hasattr(value, "model_dump_json"): encoded = value.model_dump_json(by_alias=True, exclude_unset=True) return cast(str, encoded).encode() diff --git a/tests/test_agent.py b/tests/test_agent.py index 232b4e7..c64e686 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -276,15 +276,18 @@ async def test_builds_agent_friendly_offering_details_without_invoking_action() transport = QueueTransport( response(SERVICE_DOCUMENT), response(ACTION_OFFERING), - response(schema, content_type="application/schema+json"), ) + supporting = QueueTransport(response(schema, content_type="application/schema+json")) details = await ServiceClient( - "https://demo.inflowpay.ai", transport=transport + "https://demo.inflowpay.ai", transport=transport, supporting_transport=supporting ).get_offering_details("rubber-plant") assert details.actions[0].http is not None assert details.actions[0].http.url == "https://demo.inflowpay.ai/actions/purchase" assert details.attribute_schema is not None - assert len(transport.requests) == 3 + assert len(transport.requests) == 2 + assert len(supporting.requests) == 1 + assert [request.maximum_response_bytes for request in transport.requests] == [65_536, 524_288] + assert supporting.requests[0].maximum_response_bytes == 262_144 DIRECTORY_PAGE = """{ @@ -447,6 +450,8 @@ async def test_resolves_http_action_request_schema() -> None: transport=QueueTransport( response(SERVICE_DOCUMENT), response(offering), + ), + supporting_transport=QueueTransport( response(schema, content_type="application/schema+json"), response(schema, content_type="application/schema+json"), response(schema, content_type="application/schema+json"), @@ -476,22 +481,26 @@ async def test_resolves_openapi_action_and_rejects_invalid_openapi() -> None: "openapi":"3.1.0","paths":{"/purchase":{"post":{ "operationId":"purchasePlant","responses":{}}}} }""" + supporting = QueueTransport(response(openapi, content_type="application/json")) client = ServiceClient( "https://demo.inflowpay.ai", transport=QueueTransport( response(document), response(offering), - response(openapi, content_type="application/json"), ), + supporting_transport=supporting, ) resolved = await client.resolve_action("plant", "purchase") assert resolved.operation == {"operationId": "purchasePlant", "responses": {}} + assert supporting.requests[0].maximum_response_bytes == 1_048_576 invalid_client = ServiceClient( "https://demo.inflowpay.ai", transport=QueueTransport( response(document), response(offering), + ), + supporting_transport=QueueTransport( response('{"openapi":"3.0.0"}', content_type="application/json"), ), ) @@ -522,6 +531,8 @@ async def test_offering_details_report_unusable_actions_and_attributes() -> None transport=QueueTransport( response(SERVICE_DOCUMENT), response(offering), + ), + supporting_transport=QueueTransport( response(schema, content_type="application/schema+json"), ), ) @@ -564,7 +575,7 @@ async def test_supporting_document_security_and_cache_edges() -> None: revalidating = ServiceClient( "https://demo.inflowpay.ai", cache=cache, - transport=QueueTransport(response(b"", status=304)), + supporting_transport=QueueTransport(response(b"", status=304)), ) assert await revalidating._supporting_json( "https://schemas.example/a", "schema", "application/json", {"application/json"}, 100 @@ -721,9 +732,8 @@ async def test_agent_remaining_traversal_and_supporting_document_edges() -> None async def test_action_resolution_boundaries(monkeypatch: pytest.MonkeyPatch) -> None: details = await ServiceClient( "https://demo.inflowpay.ai", - transport=QueueTransport( - response(SERVICE_DOCUMENT), response(ACTION_OFFERING), response(b"", status=500) - ), + transport=QueueTransport(response(SERVICE_DOCUMENT), response(ACTION_OFFERING)), + supporting_transport=QueueTransport(response(b"", status=500)), ).get_offering_details("rubber-plant") assert details.attribute_schema is None assert details.issues[0].scope.value == "attribute_schema" @@ -798,6 +808,8 @@ async def fake_details(client: ServiceClient, identifier: str) -> OfferingDetail ) ), response(offering), + ), + supporting_transport=QueueTransport( response(duplicate, content_type="application/json"), ), ) diff --git a/tests/test_directory_conformance.py b/tests/test_directory_conformance.py index 9c54bb1..33d8f2f 100644 --- a/tests/test_directory_conformance.py +++ b/tests/test_directory_conformance.py @@ -242,16 +242,11 @@ async def test_refuses_a_recognized_descriptor_that_breaks_its_own_rules() -> No assert page.issues[0].index == 0 -# -- nothing unchecked is passed off as checked ----------------------------------------------- +# -- Directory metadata is retained without becoming execution authority ---------------------- @pytest.mark.asyncio -async def test_drops_service_document_members_it_does_not_validate() -> None: - """A caller reading these off a record cannot tell they were never checked. - - `http` is the one that matters: a caller could build request URLs from an `endpoint_base` the - Directory invented, which is exactly the authority ROLE-03 says a Directory does not have. - """ +async def test_preserves_unverified_metadata_in_additional_members() -> None: page = await read( amend( branding={"icon": {"src": "/i.png"}, "logo": {"src": "/l.png"}}, @@ -263,7 +258,14 @@ async def test_drops_service_document_members_it_does_not_validate() -> None: ) ) - assert not set(page.items[0].additional) + assert page.items[0].additional == { + "branding": {"icon": {"src": "/i.png"}, "logo": {"src": "/l.png"}}, + "http": {"endpoint_base": "/somewhere-else"}, + "mcp": [{"type": "streamable-http", "url": "https://elsewhere.example/mcp"}], + "odp_version": "1.0", + "payment_origins": ["https://pay.example"], + "search_capabilities": {"filters": {"inline": []}}, + } @pytest.mark.asyncio diff --git a/tests/test_service.py b/tests/test_service.py index 60087c9..8d622d5 100644 --- a/tests/test_service.py +++ b/tests/test_service.py @@ -115,7 +115,7 @@ async def test_serves_document_offerings_and_collections() -> None: full_offerings = await service.handle( Request("GET", "/odp/offerings", query="representation=full") ) - assert json.loads(full_offerings.body)["items"][0]["odp_version"] == "1.0" + assert "odp_version" not in json.loads(full_offerings.body)["items"][0] full_collections = await service.handle( Request("GET", "/odp/collections", query="representation=full") ) @@ -180,6 +180,78 @@ async def search_collections( return await self.list_collections(request) +@pytest.mark.asyncio +@pytest.mark.parametrize("representation", ["full", "terse"]) +async def test_only_top_level_resources_emit_versions(representation: str) -> None: + service = _service(SearchCatalog(_catalog_options())) + for path in ( + "/odp/offerings", + "/odp/collections", + "/odp/collections/plants/offerings", + "/odp/offerings/search", + "/odp/collections/search", + ): + search = path.endswith("/search") + result = await service.handle( + Request( + "POST" if search else "GET", + path, + query=f"representation={representation}", + headers={"content-type": MEDIA_TYPE}, + body=b'{"odp_version":"1.0","query":"plant"}' if search else b"", + ) + ) + assert result.status == 200 + page = json.loads(result.body) + assert page["odp_version"] == "1.0" + assert page["items"] + assert all("odp_version" not in item for item in page["items"]) + for path in ("/odp/offerings/rubber-plant", "/odp/collections/plants"): + result = await service.handle( + Request("GET", path, query=f"representation={representation}") + ) + assert result.status == 200 + assert json.loads(result.body)["odp_version"] == "1.0" + + +@pytest.mark.asyncio +async def test_page_serialization_does_not_mutate_cached_catalog_models() -> None: + offering = Offering.model_validate( + { + "odp_version": "1.0", + "id": "item", + "name": "Item", + "custom_data": {"odp_version": "business-value"}, + } + ) + page = OfferingPage[Offering](odp_version="1.0", items=[offering]) + before = page.model_dump() + fields = offering.model_fields_set.copy() + + class CachedCatalog(Catalog): + def operations(self) -> list[Operation]: + return [Operation.LIST_OFFERINGS, Operation.GET_OFFERING] + + async def list_offerings(self, request: CatalogRequest) -> OfferingPage[Offering]: + return page + + async def get_offering(self, identifier: str, request: CatalogRequest) -> Offering: + return offering + + service = _service(CachedCatalog()) + for _ in range(2): + result = await service.handle(Request("GET", "/odp/offerings", query="representation=full")) + assert result.status == 200 + item = json.loads(result.body)["items"][0] + assert "odp_version" not in item + assert item["custom_data"] == {"odp_version": "business-value"} + standalone = await service.handle(Request("GET", "/odp/offerings/item")) + assert standalone.status == 200 + assert json.loads(standalone.body)["odp_version"] == "1.0" + assert page.model_dump() == before + assert offering.model_fields_set == fields + + @pytest.mark.asyncio async def test_routes_search_requests_with_fixed_media_type() -> None: service = _service(SearchCatalog(_catalog_options())) diff --git a/tests/test_transport.py b/tests/test_transport.py new file mode 100644 index 0000000..11f4bf7 --- /dev/null +++ b/tests/test_transport.py @@ -0,0 +1,373 @@ +from __future__ import annotations + +import asyncio +import gzip +import json +from collections.abc import AsyncIterator +from ipaddress import IPv4Address + +import httpx +import pytest + +from helpers import OFFERING_PAGE, SERVICE_DOCUMENT, QueueTransport, response +from offering_protocol.agent import ( + AgentError, + DefaultServiceClientFactory, + ServiceClient, + ServiceRequestError, +) +from offering_protocol.directory import ( + DirectoryClient, + HttpRequest, + HttpxTransport, + SearchRequest, + TransportError, +) + + +class BodyStream(httpx.AsyncByteStream): + def __init__(self, chunks: tuple[bytes, ...], failure: Exception | None = None) -> None: + self.chunks = chunks + self.failure = failure + self.reads = 0 + self.closed = False + + async def __aiter__(self) -> AsyncIterator[bytes]: + for chunk in self.chunks: + self.reads += 1 + yield chunk + if self.failure is not None: + raise self.failure + + async def aclose(self) -> None: + self.closed = True + + +@pytest.fixture +def public_dns(monkeypatch: pytest.MonkeyPatch) -> None: + async def resolve(hostname: str, port: int) -> tuple[IPv4Address, ...]: + return (IPv4Address("93.184.216.34"),) + + monkeypatch.setattr("offering_protocol.directory.transport._resolve_addresses", resolve) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("public_dns") +@pytest.mark.parametrize("budget, succeeds, reads", [(8, False, 3), (12, True, 3)]) +async def test_response_budget_is_enforced_on_actual_streamed_bytes( + budget: int, succeeds: bool, reads: int +) -> None: + body = BodyStream((b"abcd", b"efgh", b"ijkl")) + async with httpx.AsyncClient( + transport=httpx.MockTransport( + lambda _: httpx.Response(200, stream=body, headers={"content-length": "1"}) + ) + ) as client: + transport = HttpxTransport(client) + request = HttpRequest("GET", "https://service.example/", maximum_response_bytes=budget) + if succeeds: + assert (await transport.send(request)).body == b"abcdefghijkl" + else: + with pytest.raises(TransportError) as caught: + await transport.send(request) + assert caught.value.code == "RESPONSE_LIMIT_EXCEEDED" + assert body.reads == reads + assert body.closed + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("public_dns") +async def test_agent_stops_reading_before_the_complete_response() -> None: + body = BodyStream((b"x" * 4096,) * 256) + async with ( + httpx.AsyncClient( + transport=httpx.MockTransport( + lambda _: httpx.Response( + 200, stream=body, headers={"content-type": "application/odp+json"} + ) + ) + ) as client, + ServiceClient("https://service.example", transport=HttpxTransport(client)) as agent, + ): + with pytest.raises(AgentError) as caught: + await agent.inspect() + assert caught.value.code == "RESPONSE_LIMIT_EXCEEDED" + assert body.reads == 17 + assert body.closed + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("public_dns") +async def test_compression_does_not_bypass_the_decoded_byte_budget() -> None: + body = BodyStream((gzip.compress(b"x" * 4096),)) + async with httpx.AsyncClient( + transport=httpx.MockTransport( + lambda _: httpx.Response(200, stream=body, headers={"content-encoding": "gzip"}) + ) + ) as client: + with pytest.raises(TransportError) as caught: + await HttpxTransport(client).send( + HttpRequest("GET", "https://service.example/", maximum_response_bytes=64) + ) + assert caught.value.code == "RESPONSE_LIMIT_EXCEEDED" + assert body.closed + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("public_dns") +@pytest.mark.parametrize("method, status", [("HEAD", 200), ("GET", 302), ("GET", 304)]) +async def test_head_and_redirects_close_without_reading_unused_bodies( + method: str, status: int +) -> None: + body = BodyStream((b"unused",)) + async with httpx.AsyncClient( + transport=httpx.MockTransport( + lambda _: httpx.Response(status, stream=body, headers={"location": "/next"}) + ) + ) as client: + result = await HttpxTransport(client).send(HttpRequest(method, "https://service.example/")) + assert result.body == b"" + assert result.status == status + assert result.headers["location"] == "/next" + assert body.closed + assert body.reads == 0 + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("public_dns") +@pytest.mark.parametrize("size", [8, 32_768]) +async def test_error_status_survives_oversized_error_bodies(size: int) -> None: + body = BodyStream((b"x" * size, b"unread")) + async with ( + httpx.AsyncClient( + transport=httpx.MockTransport( + lambda _: httpx.Response(429, stream=body, headers={"retry-after": "5"}) + ) + ) as client, + ServiceClient("https://service.example", transport=HttpxTransport(client)) as agent, + ): + with pytest.raises(ServiceRequestError) as caught: + await agent.inspect() + assert caught.value.status == 429 + assert caught.value.headers["retry-after"] == "5" + assert caught.value.code is None + assert "x" * 100 not in str(caught.value) + assert body.closed + assert body.reads == (1 if size > 16_384 else 2) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("public_dns") +async def test_read_failure_closes_the_response() -> None: + body = BodyStream((b"partial",), httpx.ReadError("broken connection")) + async with httpx.AsyncClient( + transport=httpx.MockTransport(lambda _: httpx.Response(200, stream=body)) + ) as client: + with pytest.raises(TransportError, match="broken connection"): + await HttpxTransport(client).send(HttpRequest("GET", "https://service.example/")) + assert body.closed + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("public_dns") +async def test_cancellation_closes_the_response() -> None: + started = asyncio.Event() + finish = asyncio.Event() + + class BlockingStream(BodyStream): + async def __aiter__(self) -> AsyncIterator[bytes]: + started.set() + await finish.wait() + yield b"done" + + body = BlockingStream(()) + async with httpx.AsyncClient( + transport=httpx.MockTransport(lambda _: httpx.Response(200, stream=body)) + ) as client: + task = asyncio.create_task( + HttpxTransport(client).send(HttpRequest("GET", "https://service.example/")) + ) + await asyncio.wait_for(started.wait(), 5) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert body.closed + + +@pytest.mark.asyncio +async def test_invalid_response_budget_fails_before_network_access() -> None: + with pytest.raises(TransportError, match="must be positive"): + await HttpxTransport().send( + HttpRequest("GET", "https://service.example/", maximum_response_bytes=0) + ) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("public_dns") +@pytest.mark.parametrize("oversized", [False, True]) +async def test_default_supporting_transport_does_not_inherit_authentication( + monkeypatch: pytest.MonkeyPatch, + oversized: bool, +) -> None: + requests: list[httpx.Request] = [] + owned_clients: list[httpx.AsyncClient] = [] + client_type = httpx.AsyncClient + schema_stream = BodyStream((b"x" * 4096,) * 256) + + def serve(request: httpx.Request) -> httpx.Response: + requests.append(request) + if request.url.path == "/.well-known/odp": + document = json.loads(SERVICE_DOCUMENT) + media = "application/odp+json" + elif request.url.path.startswith("/odp/offerings/"): + document = { + "odp_version": "1.0", + "id": "item", + "name": "Item", + "schema": {"url": "https://schemas.example/root.json"}, + "attributes": {"colour": "green"}, + } + media = "application/odp+json" + else: + if oversized: + return httpx.Response( + 200, stream=schema_stream, headers={"content-type": "application/schema+json"} + ) + document = {"$schema": "https://json-schema.org/draft/2020-12/schema", "type": "object"} + if request.url.path == "/root.json": + document["$ref"] = "child.json" + media = "application/schema+json" + return httpx.Response( + 200, json=document, headers={"content-type": media, "set-cookie": "session=example"} + ) + + primary_client = client_type(auth=("example", "example"), transport=httpx.MockTransport(serve)) + + def anonymous_client(*, follow_redirects: bool, trust_env: bool) -> httpx.AsyncClient: + assert not trust_env + client = client_type( + follow_redirects=follow_redirects, + trust_env=trust_env, + transport=httpx.MockTransport(serve), + ) + owned_clients.append(client) + return client + + monkeypatch.setattr("offering_protocol.directory.transport.httpx.AsyncClient", anonymous_client) + try: + async with ServiceClient( + "https://store.example", transport=HttpxTransport(primary_client) + ) as agent: + details = await agent.get_offering_details("item") + assert details.offering.id == "item" + assert details.offering.name == "Item" + if oversized: + assert details.attribute_schema is None + assert details.offering.attributes == {} + assert len(details.issues) == 1 + assert details.issues[0].scope.value == "attribute_schema" + assert schema_stream.reads == 65 + assert schema_stream.closed + else: + assert not details.issues + assert details.attribute_schema is not None + assert not primary_client.is_closed + assert len(owned_clients) == 1 + assert owned_clients[0].is_closed + finally: + await primary_client.aclose() + assert len(requests) == (3 if oversized else 4) + assert all(request.headers.get("authorization") for request in requests[:2]) + assert all( + "authorization" not in request.headers and "cookie" not in request.headers + for request in requests[2:] + ) + assert all(request.headers["host"] == "schemas.example" for request in requests[2:]) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures("public_dns") +async def test_directory_metadata_does_not_override_live_service_routes( + monkeypatch: pytest.MonkeyPatch, +) -> None: + record = { + "description": "Example", + "indexed_at": "2026-01-01T00:00:00Z", + "language": "en", + "localizations": ["en"], + "name": "Example", + "operations": json.loads(SERVICE_DOCUMENT)["operations"], + "service_origin": "https://store.example", + "http": {"endpoint_base": "https://unverified.example/other"}, + "payment_origins": ["https://payments.example"], + "mcp": [{"url": "https://unverified.example/mcp", "transport": "streamable-http"}], + } + directory = DirectoryClient( + transport=QueueTransport( + response(json.dumps({"items": [record]}), content_type="application/json") + ) + ) + page = await directory.search_services(SearchRequest()) + assert not page.issues + assert page.items[0].additional["http"] == record["http"] + requests: list[httpx.Request] = [] + client_type = httpx.AsyncClient + + def serve(request: httpx.Request) -> httpx.Response: + requests.append(request) + body = SERVICE_DOCUMENT if request.url.path == "/.well-known/odp" else OFFERING_PAGE + return httpx.Response(200, content=body, headers={"content-type": "application/odp+json"}) + + monkeypatch.setattr( + "offering_protocol.directory.transport.httpx.AsyncClient", + lambda **kwargs: client_type(transport=httpx.MockTransport(serve), **kwargs), + ) + async with DefaultServiceClientFactory().create(page.items[0]) as client: + offerings = await client.list_offerings() + assert offerings.items[0].id == "rubber-plant" + assert [request.url.path for request in requests] == ["/.well-known/odp", "/odp/offerings"] + assert all(request.headers["host"] == "store.example" for request in requests) + + +@pytest.mark.asyncio +async def test_injected_transports_remain_caller_owned() -> None: + class OwnedTransport(QueueTransport): + closed = False + + async def aclose(self) -> None: + self.closed = True + + primary = OwnedTransport(response(SERVICE_DOCUMENT)) + supporting = OwnedTransport() + async with ServiceClient( + "https://service.example", transport=primary, supporting_transport=supporting + ) as client: + await client.inspect() + assert not primary.closed + assert not supporting.closed + + +@pytest.mark.asyncio +async def test_primary_close_failure_still_closes_owned_supporting_transport( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class ClosingTransport(QueueTransport): + def __init__(self, fail: bool) -> None: + super().__init__() + self.fail = fail + self.closed = False + + async def aclose(self) -> None: + self.closed = True + if self.fail: + raise RuntimeError("close failed") + + primary, supporting = ClosingTransport(True), ClosingTransport(False) + instances = iter((primary, supporting)) + monkeypatch.setattr( + "offering_protocol.agent.client.HttpxTransport", lambda **kwargs: next(instances) + ) + client = ServiceClient("https://service.example") + with pytest.raises(RuntimeError, match="close failed"): + await client.aclose() + assert primary.closed and supporting.closed From 70f8e60064d1bf049147ffb10b5e070847252d98 Mon Sep 17 00:00:00 2001 From: Nas Kavian Date: Tue, 22 Sep 2026 01:59:55 -0400 Subject: [PATCH 4/4] fix(sdk): correct discovery requests and cache isolation --- README.md | 11 +- scripts/conformance_adapter.py | 161 ++++++++- scripts/node_interoperability.py | 22 +- scripts/run-node-interoperability.sh | 34 +- src/offering_protocol/agent/cache.py | 4 + src/offering_protocol/agent/capabilities.py | 42 ++- src/offering_protocol/agent/client.py | 109 +++--- src/offering_protocol/agent/details.py | 6 +- src/offering_protocol/agent/schema.py | 66 ++-- src/offering_protocol/core/validation.py | 20 +- src/offering_protocol/directory/results.py | 5 +- src/offering_protocol/service/service.py | 17 +- .../service/static_catalog.py | 28 ++ tests/test_agent.py | 1 + tests/test_agent_conformance.py | 71 +++- tests/test_core.py | 2 +- tests/test_directory_mixed.py | 24 +- tests/test_review_regressions.py | 309 ++++++++++++++++++ tests/test_schema.py | 77 ++++- tests/test_service_conformance.py | 4 +- 20 files changed, 897 insertions(+), 116 deletions(-) create mode 100644 tests/test_review_regressions.py diff --git a/README.md b/README.md index b3638d8..ed7461e 100644 --- a/README.md +++ b/README.md @@ -203,7 +203,11 @@ The Agent module also provides: - Conditional request and representation caching with injectable `Cache` and `Transport` protocols. Default fallback cache lifetimes are four hours for Service documents, one hour for Collections, -and five minutes for Offerings. HTTP cache directives take precedence. `ServiceClient` uses a +five minutes for Offerings, zero for searches, one hour for Filter and Sort Definitions, and +24 hours for Attribute Schemas. Set each independently using `CacheFallbacks` (`service_document`, +`collection`, `offering`, `search`, `filters`, `sorts`, and `attribute_schema`). HTTP cache directives +take precedence. Continuations retain their originating operation's fallback; an unrecognized +continuation uses the search fallback. `ServiceClient` uses a separate anonymous transport for linked schemas and OpenAPI documents, even when its primary `transport` has authentication configured. An explicit `supporting_transport` override must also send these requests anonymously; it must not share the primary transport's credentials or cookies. @@ -265,6 +269,11 @@ invoke the resolved target. ### Caching and HTTP transport +Clients with a supplied `transport` use separate cache partitions by default. To share cached +responses between these clients, explicitly supply the same `cache_partition` only when they use +the same authentication context. Create a new client or select a new partition when changing +credentials. SDK-owned anonymous transports can share their anonymous partition. + `MemoryCache` is the default process-local cache. Implement the `Cache` protocol when representations must survive process restarts or share storage across workers. A custom `Transport` implements asynchronous `send()` and `aclose()` methods. Caller-provided caches and transports remain owned by diff --git a/scripts/conformance_adapter.py b/scripts/conformance_adapter.py index beaebea..a9f8117 100755 --- a/scripts/conformance_adapter.py +++ b/scripts/conformance_adapter.py @@ -9,6 +9,7 @@ from offering_protocol.agent import ServiceClient from offering_protocol.core import ( + Collection, CollectionSearchRequest, Offering, OfferingPage, @@ -38,7 +39,14 @@ ) from offering_protocol.core.validation import _normalize_agent_response from offering_protocol.directory.transport import HttpRequest, HttpResponse -from offering_protocol.service import CatalogRequest, Request, ServiceBuilder +from offering_protocol.service import ( + CatalogError, + CatalogRequest, + Request, + ServiceBuilder, + StaticCatalog, + StaticCatalogOptions, +) class MapTransport: @@ -241,7 +249,146 @@ async def evaluate_errors_limits(case: dict[str, Any]) -> bool | None: return (result.status == 200) == case["valid"] +def capability_definition(value: dict[str, Any], kind: str) -> dict[str, Any]: + defaults: dict[str, Any] = {"title": value["id"], "description": value["id"]} + if kind == "filters": + defaults.update(type="string", operators=["eq"]) + else: + defaults["keys"] = [{"filter_id": "region", "direction": "ascending", "missing": "last"}] + result = {**defaults, **value} + if kind == "sorts": + result["keys"] = [ + {"direction": "ascending", "missing": "last", **key} for key in result["keys"] + ] + return result + + +def capability_advertisement(value: dict[str, Any]) -> dict[str, Any]: + result = {} + for kind, source in value.items(): + if isinstance(source, list): + if not source: + continue + source = {"inline": source} + result[kind] = { + **source, + **( + {"inline": [capability_definition(item, kind) for item in source["inline"]]} + if "inline" in source + else {} + ), + } + return result + + +async def evaluate_capabilities(case: dict[str, Any]) -> bool | None: + operation = case["operation"] + origin = case.get("service_origin", "https://service.example") + document = service_document() + document["operations"] = [ + {"name": name, "authentication": "not-required"} + for name in { + "get-offering", + "list-offerings", + *case.get("operations", ["search-offerings"]), + } + ] + if operation == "validate-advertisement": + advertisement = capability_advertisement(case["advertisement"]) + document["search_capabilities"] = advertisement + try: + parse_service_document(json.dumps(document)) + for source in advertisement.values(): + if "linked" in source: + resolve_continuation(source["linked"]["href"], origin) + valid = True + except ValueError: + valid = False + return valid == case["valid"] + + documents = {} + if operation in {"validate-linked-source", "validate-linked-page-count"}: + reference = case.get("href", "/filters/0") + document["search_capabilities"] = {"filters": {"linked": {"href": reference}}} + if operation == "validate-linked-source": + pages = case["pages"] + else: + pages = [ + { + "odp_version": "1.0", + "items": [{"id": f"f{i}"}], + **({"next": f"/filters/{i + 1}"} if i + 1 < case["page_count"] else {}), + } + for i in range(case["page_count"]) + ] + for page in pages: + documents[resolve_continuation(reference, origin)] = response( + { + **page, + "items": [capability_definition(item, "filters") for item in page["items"]], + } + ) + reference = page.get("next", "") + elif operation == "merge-capabilities": + document["search_capabilities"] = capability_advertisement(case["service"]) + document["operations"].append({"name": "get-collection", "authentication": "not-required"}) + if "collection_id" in case: + documents[f"{origin}/odp/collections/{case['collection_id']}?representation=full"] = ( + response( + { + "odp_version": "1.0", + "id": case["collection_id"], + "name": "Collection", + "search_capabilities": capability_advertisement( + case["selected_collection"] + ), + } + ) + ) + else: + return None + documents[f"{origin}/.well-known/odp"] = response(document) + async with ServiceClient(origin, transport=MapTransport(documents)) as client: + if "collection_id" in case: + result = await client.get_collection_search_capabilities(case["collection_id"]) + else: + result = await client.get_offering_search_capabilities() + if operation == "merge-capabilities": + expected = case["expected"] + return ( + sorted(result.filters) == sorted(expected["filter_ids"]) + and sorted(result.sorts) == sorted(expected["sort_ids"]) + and len(result.issues) == len(expected["issues"]) + ) + return (not result.issues) == case["valid"] + + async def evaluate_case(subject: str, case: dict[str, Any], role: str) -> bool | None: + if subject == "search-capability-contract" and role == "agent": + return await evaluate_capabilities(case) + if subject == "collection-hierarchy" and role == "service": + if "chain_length" in case: + collections = [ + {"id": f"c{i}", **({"parent_ids": [f"c{i - 1}"]} if i else {})} + for i in range(case["chain_length"] + 1) + ] + else: + collections = case["collections"] + try: + StaticCatalog( + StaticCatalogOptions( + collections=tuple( + Collection.model_validate( + {"odp_version": "1.0", "name": value["id"], **value} + ) + for value in collections + ) + ) + ) + valid = True + except (CatalogError, ValueError): + valid = False + return valid == case["valid"] if subject == "protocol-version": document = {"odp_version": case["received"], "id": "item", "name": "Item"} return succeeds(lambda: parse_offering(json.dumps(document))) == case["compatible"] @@ -318,8 +465,13 @@ async def evaluate_case(subject: str, case: dict[str, Any], role: str) -> bool | ) return valid == case["valid"] if case.get("operation") == "validate-limit": - limit = case["limit"] - valid = isinstance(limit, int) and not isinstance(limit, bool) and 1 <= limit <= 100 + service = ServiceBuilder("Conformance", "Conformance", "en", "/odp").build( + EmptyCatalog() + ) + result = await service.handle( + Request("GET", "/odp/offerings", query=f"limit={case['limit']}") + ) + valid = result.status == 200 return valid == case["valid"] if case.get("operation") == "validate-next": valid = succeeds(lambda: resolve_continuation(case["next"], case["service_origin"])) @@ -381,7 +533,8 @@ async def evaluate(request: dict[str, Any]) -> dict[str, object]: return { "status": "skipped", "message": ( - f"No public Python operation maps {request['vector']['subject']}/{operation}" + f"Adapter does not exercise {request['vector']['subject']}/{operation}; " + "this is not a passing conformance result" ), } return ( diff --git a/scripts/node_interoperability.py b/scripts/node_interoperability.py index ec181d8..97bd391 100755 --- a/scripts/node_interoperability.py +++ b/scripts/node_interoperability.py @@ -5,7 +5,7 @@ import sys from offering_protocol.agent import ServiceClient -from offering_protocol.core import PriceType +from offering_protocol.core import OfferingSearchRequest, PriceType async def run(service_url: str) -> None: @@ -31,10 +31,24 @@ async def run(service_url: str) -> None: raise RuntimeError("download Action did not resolve to the Node.js Service") +async def search(service_url: str) -> None: + async with ServiceClient(service_url, allow_local_network=True) as client: + page = await client.search_offerings(OfferingSearchRequest(query="gpu", limit=2)) + if [item.id for item in page.items] != ["gpu-00000000", "gpu-00000001"]: + raise RuntimeError("Offering search did not match the Node.js marketplace") + if not page.next: + raise RuntimeError("Offering search did not return a continuation") + next_page = await client.continue_offerings(page.next) + if [item.id for item in next_page.items] != ["gpu-00000002", "gpu-00000003"]: + raise RuntimeError("Offering search continuation did not retain its query and limit") + if (await client.search_offerings(OfferingSearchRequest(query="plants"))).items: + raise RuntimeError("Offering search ignored the query") + + def main() -> None: - if len(sys.argv) != 2: - raise SystemExit("usage: node_interoperability.py SERVICE_URL") - asyncio.run(run(sys.argv[1])) + if len(sys.argv) != 3 or sys.argv[2] not in {"catalog", "search"}: + raise SystemExit("usage: node_interoperability.py SERVICE_URL catalog|search") + asyncio.run((run if sys.argv[2] == "catalog" else search)(sys.argv[1])) print("Python Agent interoperates with the Node.js example Service") diff --git a/scripts/run-node-interoperability.sh b/scripts/run-node-interoperability.sh index 852d051..2089ab4 100755 --- a/scripts/run-node-interoperability.sh +++ b/scripts/run-node-interoperability.sh @@ -8,18 +8,26 @@ log_file=${TMPDIR:-/tmp}/odp-node-interop.log package_manager=$(node -p 'require(require("node:path").resolve(process.argv[1])).packageManager' "$node_dir/package.json") corepack "$package_manager" --dir "$node_dir" build -HOST=127.0.0.1 PORT="$port" node "$node_dir/examples/odp-service-small/dist/index.js" >"$log_file" 2>&1 & -service_pid=$! -trap 'kill "$service_pid" 2>/dev/null || true' EXIT INT TERM +run_example() { + HOST=127.0.0.1 PORT="$port" node "$node_dir/examples/$1/dist/index.js" >"$log_file" 2>&1 & + service_pid=$! + trap 'kill "$service_pid" 2>/dev/null || true' EXIT INT TERM -attempt=0 -until curl --fail --silent --output /dev/null "$service_url/.well-known/odp"; do - attempt=$((attempt + 1)) - if [ "$attempt" -ge 50 ]; then - sed -n '1,120p' "$log_file" >&2 - exit 1 - fi - sleep 0.1 -done + attempt=0 + until curl --fail --silent --output /dev/null "$service_url/.well-known/odp"; do + attempt=$((attempt + 1)) + if [ "$attempt" -ge 50 ]; then + sed -n '1,120p' "$log_file" >&2 + exit 1 + fi + sleep 0.1 + done -uv run python scripts/node_interoperability.py "$service_url" + uv run python scripts/node_interoperability.py "$service_url" "$2" + kill "$service_pid" + wait "$service_pid" 2>/dev/null || true + trap - EXIT INT TERM +} + +run_example odp-service-small catalog +run_example odp-service-marketplace search diff --git a/src/offering_protocol/agent/cache.py b/src/offering_protocol/agent/cache.py index 43b9b88..9d4ba29 100644 --- a/src/offering_protocol/agent/cache.py +++ b/src/offering_protocol/agent/cache.py @@ -45,6 +45,10 @@ class CacheFallbacks: collection: timedelta = timedelta(hours=1) offering: timedelta = timedelta(minutes=5) service_document: timedelta = timedelta(hours=4) + search: timedelta = timedelta() + filters: timedelta = timedelta(hours=1) + sorts: timedelta = timedelta(hours=1) + attribute_schema: timedelta = timedelta(hours=24) def utc_now() -> datetime: diff --git a/src/offering_protocol/agent/capabilities.py b/src/offering_protocol/agent/capabilities.py index 9840dcf..b2ce3a7 100644 --- a/src/offering_protocol/agent/capabilities.py +++ b/src/offering_protocol/agent/capabilities.py @@ -21,6 +21,7 @@ parse_sort_definition_page, resolve_continuation, ) +from offering_protocol.core.validation import _agent_body _MAXIMUM_CAPABILITY_PAGES = 16 _MAXIMUM_FILTERS = 1_024 @@ -168,7 +169,7 @@ async def _add_source( source: FilterCapabilitySource | SortCapabilitySource | None, target: dict[str, Any], maximum: int, - load: Callable[[ServiceClient, str, int], Awaitable[list[Any]]], + load: Callable[[ServiceClient, str, int, frozenset[str]], Awaitable[list[Any]]], scopes: dict[str, CapabilityScope] | None, ) -> None: """Merges one capability source into the effective catalog, or reports why it cannot be. @@ -182,7 +183,7 @@ async def _add_source( return try: values: Sequence[Any] = ( - await load(client, source.linked.href, maximum - len(target)) + await load(client, source.linked.href, maximum, frozenset(target)) if source.linked is not None else list(source.inline) ) @@ -219,19 +220,25 @@ async def _add_source( async def _load_filters( - client: ServiceClient, reference: str, budget: int = _MAXIMUM_FILTERS + client: ServiceClient, + reference: str, + budget: int = _MAXIMUM_FILTERS, + existing: frozenset[str] = frozenset(), ) -> list[FilterDefinition]: values = await _load_definitions( - client, reference, budget, parse_filter_definition_page, CapabilityKind.FILTERS + client, reference, budget, parse_filter_definition_page, CapabilityKind.FILTERS, existing ) return cast("list[FilterDefinition]", values) async def _load_sorts( - client: ServiceClient, reference: str, budget: int = _MAXIMUM_SORTS + client: ServiceClient, + reference: str, + budget: int = _MAXIMUM_SORTS, + existing: frozenset[str] = frozenset(), ) -> list[SortDefinition]: values = await _load_definitions( - client, reference, budget, parse_sort_definition_page, CapabilityKind.SORTS + client, reference, budget, parse_sort_definition_page, CapabilityKind.SORTS, existing ) return cast("list[SortDefinition]", values) @@ -242,13 +249,15 @@ async def _load_definitions( budget: int, parse: Callable[[bytes | str], Any], kind: CapabilityKind, + existing: frozenset[str], ) -> list[Any]: - """Retrieves a complete linked source, one page at a time. + """Stop when new identifiers cannot fit even if all earlier identifiers are quarantined.""" + + def parse_page(body: bytes | str) -> Any: + return parse( + _agent_body(body, "filter-page" if kind is CapabilityKind.FILTERS else "sort-page") + ) - The budget is what the effective catalog has left. FLT-58 asks the Agent to stop retrieving a - source once it cannot fit, so the budget is checked as each page arrives rather than after the - whole source has been buffered: a source that can never fit costs one page, not sixteen. - """ values: list[Any] = [] next_reference = reference visited: set[str] = set() @@ -259,10 +268,15 @@ async def _load_definitions( if target in visited: raise AgentError("ODP capability pagination loop detected") visited.add(target) - body = await client._linked_odp(target, client._cache_fallbacks.collection, parse) - page = parse(body) + fallback = ( + client._cache_fallbacks.filters + if kind is CapabilityKind.FILTERS + else client._cache_fallbacks.sorts + ) + body = await client._linked_odp(target, fallback, parse_page) + page = parse_page(body) values.extend(page.items) - if len(values) > budget: + if len({value.id for value in values} - existing) > budget: raise AgentError(f"Effective {kind.value} exceed their limit") next_reference = page.next if next_reference: diff --git a/src/offering_protocol/agent/client.py b/src/offering_protocol/agent/client.py index f69ac2b..fd6691c 100644 --- a/src/offering_protocol/agent/client.py +++ b/src/offering_protocol/agent/client.py @@ -11,9 +11,11 @@ from enum import StrEnum from typing import TYPE_CHECKING, cast from urllib.parse import parse_qsl, urlencode, urljoin, urlsplit, urlunsplit +from uuid import uuid4 from offering_protocol.agent.cache import Cache, CacheFallbacks, CacheRecord, MemoryCache, utc_now from offering_protocol.core import ( + VERSION, Collection, CollectionSearchRequest, Offering, @@ -27,6 +29,8 @@ build_operation_url, derive_service_origin, parse_agent_service_document, + parse_collection_search_request, + parse_offering_search_request, resolve_continuation, ) from offering_protocol.core import ( @@ -44,7 +48,7 @@ from offering_protocol.core import ( parse_problem_response as parse_problem_response_strict, ) -from offering_protocol.core.validation import _agent_body +from offering_protocol.core.validation import _agent_body, _nesting_depth from offering_protocol.directory.transport import ( HttpRequest, HttpResponse, @@ -125,7 +129,7 @@ def __init__( allow_local_network: bool = False, cache: Cache | None = None, cache_fallbacks: CacheFallbacks | None = None, - cache_partition: str = "anonymous", + cache_partition: str | None = None, supporting_transport: Transport | None = None, transport: Transport | None = None, ) -> None: @@ -133,7 +137,14 @@ def __init__( self._accept_language = accept_language self._cache = cache or MemoryCache() self._cache_fallbacks = cache_fallbacks or CacheFallbacks() - self._cache_partition = cache_partition + self._continuation_fallbacks: dict[str, timedelta] = {} + self._cache_partition = ( + cache_partition + if cache_partition is not None + else str(uuid4()) + if transport is not None + else "anonymous" + ) self._owns_transport = transport is None self._transport = transport or HttpxTransport(allow_local_network=allow_local_network) self._owns_supporting_transport = supporting_transport is None @@ -251,26 +262,32 @@ async def search_offerings( async def continue_collections(self, next_reference: str) -> Page[Collection]: target = resolve_continuation(next_reference, self.service_origin) + fallback = self._continuation_fallbacks.get(target, self._cache_fallbacks.search) response = await self._request_cached( "GET", target, b"", _MAXIMUM_RESOURCE_BYTES, - self._cache_fallbacks.collection, + fallback, parse_collection_page, + cache_context=f"continuation:{fallback.total_seconds()}", ) + self._remember_continuation(response.body, fallback) return parse_collection_page(response.body) async def continue_offerings(self, next_reference: str) -> OfferingPage[Offering]: target = resolve_continuation(next_reference, self.service_origin) + fallback = self._continuation_fallbacks.get(target, self._cache_fallbacks.search) response = await self._request_cached( "GET", target, b"", _MAXIMUM_RESOURCE_BYTES, - self._cache_fallbacks.offering, + fallback, parse_offering_page, + cache_context=f"continuation:{fallback.total_seconds()}", ) + self._remember_continuation(response.body, fallback) return parse_offering_page(response.body) async def list_all_collections( @@ -357,31 +374,43 @@ async def _get_page( response = await self._request_cached( "GET", target, b"", _MAXIMUM_RESOURCE_BYTES, fallback, parser ) + if operation not in {Operation.GET_COLLECTION, Operation.GET_OFFERING}: + self._remember_continuation(response.body, fallback) return response.body async def _post_search( self, operation: Operation, value: Mapping[str, object], representation: Representation ) -> bytes: + body = json.dumps({"odp_version": VERSION, **value}, separators=(",", ":")).encode() + parser = ( + parse_collection_search_request + if operation is Operation.SEARCH_COLLECTIONS + else parse_offering_search_request + ) + parser(body) inspection = await self._require_operation(operation) target = build_operation_url( inspection.document.http.endpoint_base, operation, self.service_origin, None ) target = _append_query(target, {"representation": representation.value}) - fallback = ( - self._cache_fallbacks.collection - if operation is Operation.SEARCH_COLLECTIONS - else self._cache_fallbacks.offering - ) + fallback = self._cache_fallbacks.search response = await self._request_cached( "POST", target, - json.dumps(value, separators=(",", ":")).encode(), + body, _MAXIMUM_RESOURCE_BYTES, fallback, _operation_parser(operation), ) + self._remember_continuation(response.body, fallback) return response.body + def _remember_continuation(self, body: bytes, fallback: timedelta) -> None: + reference = json.loads(body).get("next") + if reference: + target = resolve_continuation(reference, self.service_origin) + self._continuation_fallbacks[target] = fallback + async def _require_operation(self, operation: Operation) -> Inspection: inspection = await self.inspect() if not any(item.name is operation for item in inspection.document.operations): @@ -397,8 +426,10 @@ async def _request_cached( fallback: timedelta, parser: object, maximum_depth: int = _MAXIMUM_DEPTH, + *, + cache_context: str = "", ) -> _FetchedResponse: - key = self._cache_key(method, target, body) + key = self._cache_key(method, target, body) + cache_context cached = self._cache.get(key) now = utc_now() if cached is not None and now < cached.expires: @@ -495,6 +526,25 @@ async def _supporting_json( maximum_bytes: int, maximum_depth: int | None = None, ) -> dict[str, object]: + response = await self._supporting_document( + target, resource_class, accept, media_types, maximum_bytes, maximum_depth + ) + return _decode_json_object(response.body) + + async def _supporting_document( + self, + target: str, + resource_class: str, + accept: str, + media_types: set[str], + maximum_bytes: int, + maximum_depth: int | None = None, + ) -> _FetchedResponse: + fallback = ( + self._cache_fallbacks.attribute_schema + if resource_class == "attribute-schema" + else timedelta() + ) current = target if not _is_https_url(current): raise AgentError("ODP supporting document URL must use HTTPS") @@ -502,7 +552,7 @@ async def _supporting_json( cached = self._cache.get(key) now = utc_now() if cached is not None and now < cached.expires: - return _decode_json_object(cached.body) + return _FetchedResponse(cached.body, cached.final_url, Freshness.FRESH) conditional: dict[str, str] = {} if cached is not None: if cached.etag: @@ -552,7 +602,7 @@ async def _supporting_json( else: lifetime = cached.expires - cached.stored expires = ( - _expiration(response.headers, timedelta(), now) + _expiration(response.headers, fallback, now) if _has_freshness(response.headers) else now + max(lifetime, timedelta()) ) @@ -560,7 +610,7 @@ async def _supporting_json( key, replace(cached, expires=expires, final_url=current, stored=now), ) - return _decode_json_object(cached.body) + return _FetchedResponse(cached.body, current, Freshness.REVALIDATED) if not 200 <= response.status < 300: raise ServiceRequestError( response.status, @@ -576,14 +626,14 @@ async def _supporting_json( raise AgentError("ODP supporting document has an unsupported media type") if maximum_depth is not None: _require_depth(response.body, maximum_depth, "ODP supporting document") - value = _decode_json_object(response.body) - if _cacheable("GET", response.headers, timedelta()): + _decode_json_object(response.body) + if _cacheable("GET", response.headers, fallback): self._cache.set( key, CacheRecord( body=response.body, etag=response.headers.get("etag"), - expires=_expiration(response.headers, timedelta(), now), + expires=_expiration(response.headers, fallback, now), final_url=current, last_modified=response.headers.get("last-modified"), status=response.status, @@ -592,7 +642,7 @@ async def _supporting_json( ) else: self._cache.delete(key) - return value + return _FetchedResponse(response.body, current, Freshness.FETCHED) raise AgentError("ODP supporting document exceeded its redirect limit") # pragma: no cover def _cache_key(self, method: str, target: str, body: bytes) -> str: @@ -709,27 +759,6 @@ def _require_depth(body: bytes, maximum: int, subject: str) -> None: raise AgentError(f"{subject} exceeds its nesting-depth limit") -def _nesting_depth(value: object) -> int: - """Counts container nesting from the top-level value. - - `{}` and `{"a": 1}` are both depth 1 and `{"a": {"b": 1}}` is depth 2: a scalar is a value a - container holds, not a level of its own. - """ - maximum = 0 - pending: list[tuple[int, object]] = [(1, value)] - while pending: - depth, current = pending.pop() - if isinstance(current, dict): - children: list[object] = list(current.values()) - elif isinstance(current, list): - children = list(current) - else: - continue - maximum = max(maximum, depth) - pending.extend((depth + 1, child) for child in children) - return maximum - - def _traversal_bounds(options: TraversalOptions) -> tuple[int, int]: if not 1 <= options.max_items <= 10_000 or not 1 <= options.max_pages <= 16: raise AgentError("traversal exceeds 10000 items or 16 pages") diff --git a/src/offering_protocol/agent/details.py b/src/offering_protocol/agent/details.py index 142bd1a..29ba669 100644 --- a/src/offering_protocol/agent/details.py +++ b/src/offering_protocol/agent/details.py @@ -6,6 +6,9 @@ from enum import StrEnum from urllib.parse import urljoin, urlsplit +from jsonschema.exceptions import SchemaError +from referencing.exceptions import Unresolvable + from offering_protocol.agent.client import AgentError, ServiceClient from offering_protocol.agent.schema import resolve_schema from offering_protocol.core import ( @@ -96,7 +99,8 @@ async def get_offering_details(client: ServiceClient, identifier: str) -> Offeri "Offering attributes do not match their Attribute Schema", ) ) - except (AgentError, ValueError) as error: + except (AgentError, ValueError, SchemaError, Unresolvable) as error: + attribute_schema = None offering = offering.model_copy(update={"attributes": {}}) issues.append(OfferingIssue(OfferingIssueScope.ATTRIBUTE_SCHEMA, str(error))) return OfferingDetails(tuple(actions), attribute_schema, tuple(issues), offering) diff --git a/src/offering_protocol/agent/schema.py b/src/offering_protocol/agent/schema.py index 2544718..aaf5831 100644 --- a/src/offering_protocol/agent/schema.py +++ b/src/offering_protocol/agent/schema.py @@ -2,22 +2,22 @@ from __future__ import annotations -import json from dataclasses import dataclass -from typing import Any, Protocol +from typing import Any, Protocol, cast from urllib.parse import urldefrag, urljoin, urlsplit -from jsonschema.validators import validator_for +from jsonschema.validators import Draft202012Validator, validator_for from referencing import Registry, Resource +from referencing.jsonschema import DRAFT202012 -from offering_protocol.agent.client import AgentError, ServiceClient +from offering_protocol.agent.client import AgentError, ServiceClient, _decode_json_object _DIALECT = "https://json-schema.org/draft/2020-12/schema" _MAXIMUM_DOCUMENT_BYTES = 262_144 _MAXIMUM_DOCUMENTS = 16 _MAXIMUM_DEPTH = 8 _MAXIMUM_GRAPH_BYTES = 1_048_576 -_STANDARD_VOCABULARY = "https://json-schema.org/draft/2020-12/vocab/" +_SUPPORTED_VOCABULARIES = frozenset(Draft202012Validator.META_SCHEMA["$vocabulary"]) class SchemaValidator(Protocol): @@ -33,6 +33,7 @@ class ResolvedSchema: async def resolve_schema(client: ServiceClient, target: str) -> ResolvedSchema: root_url = _document_url(target) documents: dict[str, dict[str, object]] = {} + retrievals: dict[str, str] = {} graph_bytes = 0 async def load(document_url: str, depth: int) -> None: @@ -43,28 +44,42 @@ async def load(document_url: str, depth: int) -> None: raise AgentError("ODP Attribute Schema graph exceeds 16 documents") if depth > _MAXIMUM_DEPTH: raise AgentError("ODP Attribute Schema graph exceeds eight reference levels") - document = await client._supporting_json( + response = await client._supporting_document( document_url, "attribute-schema", "application/schema+json", {"application/schema+json"}, _MAXIMUM_DOCUMENT_BYTES, ) + document = _decode_json_object(response.body) _require_schema(document) - encoded = json.dumps(document, separators=(",", ":")).encode() - graph_bytes += len(encoded) + graph_bytes += len(response.body) if graph_bytes > _MAXIMUM_GRAPH_BYTES: raise AgentError("ODP Attribute Schema graph exceeds its byte limit") documents[document_url] = document - for reference_url in _schema_references(document, document_url): + retrievals[document_url] = response.final_url + for reference_url in _schema_references(document, response.final_url): await load(reference_url, depth + 1) await load(root_url, 0) root = documents[root_url] registry: Registry[Any] = Registry().with_resources( - (url, Resource.from_contents(document)) for url, document in documents.items() + ( + url, + Resource.from_contents( + { + **document, + "$id": urljoin(retrievals[requested], cast("str", document.get("$id", ""))), + } + ), + ) + for requested, document in documents.items() + for url in {requested, retrievals[requested]} ) - validation_root = {"$id": root_url, **root} + validation_root = { + **root, + "$id": urljoin(retrievals[root_url], cast("str", root.get("$id", ""))), + } validator_type = validator_for(validation_root) validator_type.check_schema(validation_root) return ResolvedSchema(root, validator_type(validation_root, registry=registry)) @@ -86,12 +101,12 @@ def _document_url(value: str) -> str: def _require_schema(document: dict[str, object]) -> None: if document.get("$schema") != _DIALECT: raise AgentError("ODP Attribute Schema must declare JSON Schema Draft 2020-12") - pending: list[object] = [document] + validator_for(document).check_schema(document) + pending = [DRAFT202012.create_resource(document)] while pending: - value = pending.pop() - if isinstance(value, list): - pending.extend(value) - elif isinstance(value, dict): + resource = pending.pop() + value = resource.contents + if isinstance(value, dict): if "$dynamicRef" in value: reference = value["$dynamicRef"] if not isinstance(reference, str) or not reference.startswith("#"): @@ -101,23 +116,22 @@ def _require_schema(document: dict[str, object]) -> None: vocabulary = value.get("$vocabulary") if isinstance(vocabulary, dict): for uri, required in vocabulary.items(): - if required is True and not str(uri).startswith(_STANDARD_VOCABULARY): + if required is True and uri not in _SUPPORTED_VOCABULARIES: raise AgentError( f"ODP Attribute Schema requires unsupported vocabulary {uri}" ) - pending.extend(value.values()) + pending.extend(resource.subresources()) -def _schema_references(document: object, retrieval_url: str) -> tuple[str, ...]: +def _schema_references(document: Any, retrieval_url: str) -> tuple[str, ...]: references: list[str] = [] local_resources = {_document_url(retrieval_url)} - pending = [(document, retrieval_url)] + pending = [(DRAFT202012.create_resource(document), retrieval_url)] while pending: - value, inherited_base = pending.pop() - if isinstance(value, list): - pending.extend((child, inherited_base) for child in value) - elif isinstance(value, dict): - base = inherited_base + resource, inherited_base = pending.pop() + value = resource.contents + base = inherited_base + if isinstance(value, dict): identifier = value.get("$id") if isinstance(identifier, str): base = urljoin(inherited_base, identifier) @@ -125,5 +139,5 @@ def _schema_references(document: object, retrieval_url: str) -> tuple[str, ...]: reference = value.get("$ref") if isinstance(reference, str): references.append(_document_url(urljoin(base, reference))) - pending.extend((child, base) for keyword, child in value.items() if keyword != "$ref") + pending.extend((child, base) for child in resource.subresources()) return tuple(reference for reference in references if reference not in local_resources) diff --git a/src/offering_protocol/core/validation.py b/src/offering_protocol/core/validation.py index fdec898..cca31a4 100644 --- a/src/offering_protocol/core/validation.py +++ b/src/offering_protocol/core/validation.py @@ -278,6 +278,22 @@ def _filter_payment_options(value: dict[str, Any]) -> None: payment.pop("options", None) +def _nesting_depth(value: object) -> int: + maximum = 0 + pending: list[tuple[int, object]] = [(1, value)] + while pending: + depth, current = pending.pop() + if isinstance(current, dict): + children: list[object] = list(current.values()) + elif isinstance(current, list): + children = list(current) + else: + continue + maximum = max(maximum, depth) + pending.extend((depth + 1, child) for child in children) + return maximum + + def _normalize_branding(value: dict[str, Any]) -> None: branding = value.get("branding") if not isinstance(branding, dict): @@ -290,7 +306,8 @@ def _normalize_branding(value: dict[str, Any]) -> None: and isinstance(image.get("type"), str) and image["type"] not in {"image/png", "image/svg+xml", "image/webp"} ): - normalized.pop(member, None) + value.pop("branding", None) + return elif isinstance(image, dict): normalized[member] = { key: entry for key, entry in image.items() if key in {"src", "type"} @@ -325,6 +342,7 @@ def _normalize_offering(value: dict[str, Any]) -> None: schema = value.get("schema") if isinstance(schema, dict) and set(schema).difference({"url"}): value.pop("schema", None) + value.pop("attributes", None) price = value.get("price") if ( isinstance(price, dict) diff --git a/src/offering_protocol/directory/results.py b/src/offering_protocol/directory/results.py index 7d9aed2..e31e0c2 100644 --- a/src/offering_protocol/directory/results.py +++ b/src/offering_protocol/directory/results.py @@ -56,6 +56,7 @@ def _result(value: JsonValue) -> DirectoryResult: _text(service, "service_id", 128) _origin(service) _timestamp(service) + candidate = dict(service) for name in ( "branding", "http", @@ -64,9 +65,9 @@ def _result(value: JsonValue) -> DirectoryResult: "payment_origins", "search_capabilities", ): - service.pop(name, None) + candidate.pop(name, None) document = parse_agent_service_document( - json.dumps({**service, "odp_version": "1.0", "http": {"endpoint_base": "/"}}) + json.dumps({**candidate, "odp_version": "1.0", "http": {"endpoint_base": "/"}}) ) service["operations"] = [operation.to_dict() for operation in document.operations] if document.protocols is None: diff --git a/src/offering_protocol/service/service.py b/src/offering_protocol/service/service.py index 7ab137c..ae28633 100644 --- a/src/offering_protocol/service/service.py +++ b/src/offering_protocol/service/service.py @@ -43,6 +43,7 @@ parse_offering_search_request, parse_service_document, ) +from offering_protocol.core.validation import _nesting_depth MEDIA_TYPE = "application/odp+json" PROBLEM_MEDIA_TYPE = "application/problem+json" @@ -390,8 +391,8 @@ def _catalog_request( limit = int(values.get("limit", "0")) except ValueError as error: raise RequestError(400, "INVALID_REQUEST", "query parameter is invalid") from error - if not 0 <= limit <= 100: - raise RequestError(400, "INVALID_REQUEST", "limit exceeds 100") + if "limit" in values and not 1 <= limit <= 100: + raise RequestError(400, "INVALID_REQUEST", "limit must be between 1 and 100") return CatalogRequest( accept_language=headers.get("accept-language"), cursor=values.get("cursor"), @@ -408,6 +409,16 @@ def _search_body(request: Request, headers: dict[str, str]) -> bytes: content_type = headers.get("content-type", "").split(";", 1)[0] if content_type != MEDIA_TYPE: raise RequestError(415, "UNSUPPORTED_MEDIA_TYPE", f"Content-Type must be {MEDIA_TYPE}") + try: + value = json.loads(request.body) + except RecursionError as error: + raise RequestError( + 413, "REQUEST_TOO_LARGE", "request exceeds its nesting-depth limit" + ) from error + except (ValueError, UnicodeDecodeError): + return request.body + if _nesting_depth(value) > 16: + raise RequestError(413, "REQUEST_TOO_LARGE", "request exceeds its nesting-depth limit") return request.body @@ -439,6 +450,8 @@ def _json_response(value: object, maximum_bytes: int, exchange: _Exchange) -> Re body = _encode(value) if len(body) > maximum_bytes: raise ServiceError("response body is too large") + if _nesting_depth(json.loads(body)) > (8 if isinstance(value, ServiceDocument) else 16): + raise ServiceError("response body exceeds its nesting-depth limit") if isinstance(value, Page): language = ( ", ".join( diff --git a/src/offering_protocol/service/static_catalog.py b/src/offering_protocol/service/static_catalog.py index c3862f1..0995e17 100644 --- a/src/offering_protocol/service/static_catalog.py +++ b/src/offering_protocol/service/static_catalog.py @@ -48,6 +48,7 @@ def __init__(self, options: StaticCatalogOptions) -> None: except ValueError as error: raise CatalogError(f"Static Catalog contains an invalid resource: {error}") from error self._collection_by_id = _unique(self._collections, "Collection") + _validate_hierarchy(self._collection_by_id) self._offering_by_id = _unique(self._offerings, "Offering") for offering in self._offerings: if any( @@ -114,6 +115,33 @@ async def list_collection_offerings( Resource = TypeVar("Resource", Collection, Offering) +def _validate_hierarchy(collections: dict[str, Collection]) -> None: + depths: dict[str, int] = {} + visiting: set[str] = set() + + def depth(identifier: str) -> int: + if identifier in depths: + return depths[identifier] + if identifier not in collections: + raise CatalogError(f"Collection parent {identifier} does not exist") + if identifier in visiting: + raise CatalogError("Collection hierarchy contains a cycle") + if len(visiting) > 32: + raise CatalogError("Collection hierarchy exceeds 32 edges") + visiting.add(identifier) + result = max( + (depth(parent) + 1 for parent in collections[identifier].parent_ids), default=0 + ) + visiting.remove(identifier) + if result > 32: + raise CatalogError("Collection hierarchy exceeds 32 edges") + depths[identifier] = result + return result + + for identifier in collections: + depth(identifier) + + def _unique(values: tuple[Resource, ...], label: str) -> dict[str, Resource]: result: dict[str, Resource] = {} for value in values: diff --git a/tests/test_agent.py b/tests/test_agent.py index c64e686..9b7d60a 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -911,6 +911,7 @@ async def test_remaining_capability_and_cache_branches(monkeypatch: pytest.Monke conditional = ServiceClient( "https://demo.inflowpay.ai", cache=cache, + cache_partition="anonymous", transport=conditional_transport, ) assert (await conditional.inspect()).freshness is Freshness.REVALIDATED diff --git a/tests/test_agent_conformance.py b/tests/test_agent_conformance.py index c2a134c..c24871a 100644 --- a/tests/test_agent_conformance.py +++ b/tests/test_agent_conformance.py @@ -29,7 +29,6 @@ _MAXIMUM_DOCUMENT_DEPTH, _consume, _expiration, - _nesting_depth, ) from offering_protocol.core import ( CapabilityLink, @@ -45,6 +44,7 @@ SortDirection, SortKey, ) +from offering_protocol.core.validation import _nesting_depth from offering_protocol.directory.transport import HttpResponse ORIGIN = "https://plants.example" @@ -251,6 +251,73 @@ async def test_a_source_that_exactly_fills_the_bound_is_accepted() -> None: assert not result.issues +@pytest.mark.asyncio +async def test_linked_duplicate_is_quarantined_when_service_catalog_is_full() -> None: + document = json.loads(SERVICE_DOCUMENT) + document["operations"].extend( + [ + {"name": "search-offerings", "authentication": "not-required"}, + {"name": "get-collection", "authentication": "not-required"}, + ] + ) + document["search_capabilities"] = {"filters": {"linked": {"href": "/filters/0"}}} + collection = { + "odp_version": "1.0", + "id": "plants", + "name": "Plants", + "search_capabilities": {"filters": {"linked": {"href": "/collection-filters"}}}, + } + pages = [ + response( + _filter_page( + [f"f{index}" for index in range(start, min(start + 100, 1024))], + f"/filters/{start + 100}" if start + 100 < 1024 else "", + ) + ) + for start in range(0, 1024, 100) + ] + transport = QueueTransport( + response(json.dumps(document)), + response(json.dumps(collection)), + *pages, + response(_filter_page(["f0"])), + ) + async with ServiceClient(ORIGIN, transport=transport) as client: + result = await client.get_collection_search_capabilities("plants") + assert len(result.filters) == 1023 + assert "f0" not in result.filters + assert len(result.issues) == 1 + assert result.issues[0].message == "Duplicate filters: f0" + + +@pytest.mark.asyncio +async def test_linked_capabilities_skip_only_unknown_definitions() -> None: + document = json.loads(SERVICE_DOCUMENT) + document["operations"].append({"name": "search-offerings", "authentication": "not-required"}) + document["search_capabilities"] = { + "filters": {"linked": {"href": "/filters"}}, + "sorts": {"linked": {"href": "/sorts"}}, + } + good_filter = _filter("weight").to_dict() + good_sort = _sort("weight").to_dict() + filters = { + "odp_version": "1.0", + "items": [good_filter, {**good_filter, "id": "future", "type": "future"}], + } + sorts = { + "odp_version": "1.0", + "items": [good_sort, {**good_sort, "id": "future", "keys": [{"direction": "future"}]}], + } + transport = QueueTransport( + *(response(json.dumps(value)) for value in (document, filters, sorts)) + ) + async with ServiceClient(ORIGIN, transport=transport) as client: + result = await client.get_offering_search_capabilities() + assert list(result.filters) == ["weight"] + assert list(result.sorts) == ["weight"] + assert not result.issues + + @pytest.mark.asyncio async def test_paging_stops_at_the_page_that_overflows_the_bound() -> None: """FLT-58: a source that cannot fit costs one page rather than sixteen.""" @@ -549,7 +616,7 @@ async def test_collection_page_inherits_version_without_requiring_it_on_items(se ) client = ServiceClient("https://store.example", transport=transport) page = ( - await client.search_collections(CollectionSearchRequest()) + await client.search_collections(CollectionSearchRequest(query="plants")) if search else await client.list_collections() ) diff --git a/tests/test_core.py b/tests/test_core.py index 6bf171f..21df1ca 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -85,7 +85,7 @@ def test_normalizes_agent_response_capabilities() -> None: ) assert service["operations"] == [{"authentication": "not-required", "name": "list-offerings"}] assert service["mcp"] == [{"type": "streamable-http", "url": "/mcp"}] - assert service["branding"] == {"logo": {"src": "/logo", "type": "image/png"}} + assert "branding" not in service assert service["protocols"]["payments"][0]["options"] == ["inflow"] assert len(service["search_capabilities"]["filters"]["inline"]) == 1 assert len(service["search_capabilities"]["sorts"]["inline"]) == 1 diff --git a/tests/test_directory_mixed.py b/tests/test_directory_mixed.py index eaa9a6d..3990752 100644 --- a/tests/test_directory_mixed.py +++ b/tests/test_directory_mixed.py @@ -109,6 +109,28 @@ async def test_mixed_results_metadata_unknown_types_and_requests() -> None: } +@pytest.mark.asyncio +@pytest.mark.parametrize("kind", ["service", "collection"]) +async def test_mixed_results_preserve_descriptive_service_metadata(kind: str) -> None: + parent = service() + metadata: dict[str, JsonValue] = { + "mcp": [{"type": "streamable-http", "url": "/mcp"}], + "branding": {"icon": {"src": "/icon.png"}, "logo": {"src": "/logo.png"}}, + "payment_origins": ["https://payments.example"], + } + parent.update(metadata) + raw = item(kind) + raw["service"] = parent + result = await DirectoryClient(transport=transport_for({"items": [raw]})).search( + ResourceSearchRequest() + ) + assert not result.issues + entry = result.items[0] + assert isinstance(entry, (ServiceResult, CollectionResult)) + for key, value in metadata.items(): + assert entry.service.additional[key] == value + + @pytest.mark.asyncio @pytest.mark.parametrize( "path,value", @@ -166,7 +188,7 @@ async def test_optional_members_and_unverified_metadata() -> None: assert not result.issues parsed = result.items[0] assert isinstance(parsed, ServiceResult) - assert "http" not in parsed.service.additional + assert parsed.service.additional["http"] == {"endpoint_base": "https://untrusted.example/"} assert parsed.service.protocols is None assert parsed.available_through is not None and parsed.available_through.name is None parent = parsed.service.to_dict() diff --git a/tests/test_review_regressions.py b/tests/test_review_regressions.py new file mode 100644 index 0000000..79fd085 --- /dev/null +++ b/tests/test_review_regressions.py @@ -0,0 +1,309 @@ +from __future__ import annotations + +import json +from datetime import timedelta +from urllib.parse import urlsplit + +import httpx +import pytest + +from helpers import OFFERING, SERVICE_DOCUMENT, QueueTransport, response +from offering_protocol.agent import ServiceClient +from offering_protocol.agent.cache import CacheFallbacks, MemoryCache +from offering_protocol.core import ( + Collection, + CollectionSearchRequest, + OdpValidationError, + Offering, + OfferingSearchRequest, + parse_agent_service_document, +) +from offering_protocol.directory import HttpRequest, HttpResponse, HttpxTransport +from offering_protocol.service import ( + CatalogError, + Request, + Service, + StaticCatalog, + StaticCatalogOptions, +) +from test_service import SearchCatalog, _catalog_options, _service + + +class ServiceTransport: + def __init__(self, service: Service) -> None: + self.service = service + self.requests: list[HttpRequest] = [] + + async def send(self, request: HttpRequest) -> HttpResponse: + self.requests.append(request) + url = urlsplit(request.url) + result = await self.service.handle( + Request( + request.method, + url.path, + headers=request.headers, + query=url.query, + body=request.body, + ) + ) + return HttpResponse(result.status, result.headers, result.body) + + async def aclose(self) -> None: + pass + + +@pytest.mark.asyncio +async def test_search_requests_work_against_service_without_explicit_versions() -> None: + transport = ServiceTransport(_service(SearchCatalog(_catalog_options()))) + async with ServiceClient("https://service.example", transport=transport) as client: + assert (await client.search_collections(CollectionSearchRequest(query="plants"))).items + assert (await client.search_offerings(OfferingSearchRequest(query="plants"))).items + requests = [json.loads(item.body) for item in transport.requests if item.method == "POST"] + assert requests == [{"odp_version": "1.0", "query": "plants"}] * 2 + before = len(transport.requests) + with pytest.raises(OdpValidationError): + await client.search_collections(CollectionSearchRequest()) + with pytest.raises(OdpValidationError): + await client.search_offerings(OfferingSearchRequest(query="plants", limit=0)) + assert len(transport.requests) == before + + +@pytest.mark.asyncio +@pytest.mark.parametrize("partition", [None, "same-account"]) +async def test_httpx_authentication_isolated_unless_partition_explicit( + partition: str | None, +) -> None: + requests: list[httpx.Request] = [] + + def serve(request: httpx.Request) -> httpx.Response: + requests.append(request) + body = json.loads(SERVICE_DOCUMENT if request.url.path.endswith("odp") else OFFERING) + if "authorization" in request.headers: + body["name"] = "Private" + return httpx.Response(200, json=body, headers={"content-type": "application/odp+json"}) + + cache = MemoryCache() + async with ( + httpx.AsyncClient(auth=("user", "secret"), transport=httpx.MockTransport(serve)) as private, + httpx.AsyncClient(transport=httpx.MockTransport(serve)) as anonymous, + ServiceClient( + "https://127.0.0.1", + cache=cache, + cache_partition=partition, + transport=HttpxTransport(private, allow_local_network=True), + ) as first, + ServiceClient( + "https://127.0.0.1", + cache=cache, + cache_partition=partition, + transport=HttpxTransport(private if partition else anonymous, allow_local_network=True), + ) as second, + ): + assert (await first.get_offering("rubber-plant")).name == "Private" + assert (await second.get_offering("rubber-plant")).name == ( + "Private" if partition else "Rubber Plant" + ) + assert len(requests) == (2 if partition else 4) + if partition is None: + assert "authorization" not in requests[-1].headers + + +def test_unknown_branding_format_omits_whole_pair() -> None: + document = json.loads(SERVICE_DOCUMENT) + for unknown in ("icon", "logo"): + document["branding"] = { + "icon": {"src": "/icon.png", "type": "image/png"}, + "logo": {"src": "/logo.png", "type": "image/png"}, + } + document["branding"][unknown]["type"] = "image/future" + assert parse_agent_service_document(json.dumps(document)).branding is None + + document["branding"] = {"future": True} + assert parse_agent_service_document(json.dumps(document)).branding is None + document["branding"] = {"icon": None} + with pytest.raises(OdpValidationError): + parse_agent_service_document(json.dumps(document)) + document["branding"] = {"icon": {"src": "/icon.png"}, "logo": {"src": "/logo.png"}} + assert parse_agent_service_document(json.dumps(document)).branding is not None + + +@pytest.mark.asyncio +async def test_unknown_attribute_schema_metadata_does_not_reject_offering() -> None: + offering = json.loads(OFFERING) + offering.update( + schema={"url": "https://schemas.example/root", "future": True}, attributes={"x": 1} + ) + async with ServiceClient( + "https://service.example", + transport=QueueTransport(response(SERVICE_DOCUMENT), response(json.dumps(offering))), + ) as client: + result = await client.get_offering("rubber-plant") + assert result.name == "Rubber Plant" + assert result.schema_ is None and not result.attributes + + +@pytest.mark.asyncio +@pytest.mark.parametrize("collection", [False, True]) +@pytest.mark.parametrize("search", [False, True]) +async def test_continuations_retain_originating_cache_policy( + collection: bool, search: bool +) -> None: + kind = "collections" if collection else "offerings" + operation = f"{'search' if search else 'list'}-{kind}" + document = json.loads(SERVICE_DOCUMENT) + document["operations"] = [ + *document["operations"], + {"name": operation, "authentication": "not-required"}, + ] + if operation == "list-offerings": + document["operations"].pop() + next_reference = f"/odp/{kind}?cursor=opaque" + page = {"odp_version": "1.0", "items": [], "next": next_reference} + old = {"odp_version": "1.0", "items": [{"id": "item", "name": "Old"}]} + new = {"odp_version": "1.0", "items": [{"id": "item", "name": "New"}]} + transport = QueueTransport( + *(response(json.dumps(value)) for value in (document, page, old, new)) + ) + async with ServiceClient("https://service.example", transport=transport) as client: + if collection: + if search: + await client.search_collections(CollectionSearchRequest(query="plant")) + else: + await client.list_collections() + assert (await client.continue_collections(next_reference)).items[0].name == "Old" + name = (await client.continue_collections(next_reference)).items[0].name + else: + if search: + await client.search_offerings(OfferingSearchRequest(query="plant")) + else: + await client.list_offerings() + assert (await client.continue_offerings(next_reference)).items[0].name == "Old" + name = (await client.continue_offerings(next_reference)).items[0].name + assert name == ("New" if search else "Old") + assert len(transport.requests) == (4 if search else 3) + + +def test_static_catalog_validates_complete_hierarchy() -> None: + def build(parents: dict[str, list[str]]) -> StaticCatalog: + return StaticCatalog( + StaticCatalogOptions( + collections=tuple( + Collection.model_validate( + { + "id": key, + "name": key, + "odp_version": "1.0", + **({"parent_ids": value} if value else {}), + } + ) + for key, value in parents.items() + ) + ) + ) + + build({"root": [], "left": ["root"], "right": ["root"], "leaf": ["left", "right"]}) + with pytest.raises(CatalogError, match="does not exist"): + build({"leaf": ["missing"]}) + with pytest.raises(CatalogError, match="cycle"): + build({"a": ["b"], "b": ["a"]}) + for reverse in (False, True): + chain = {f"c{i}": [f"c{i - 1}"] if i else [] for i in range(33)} + build(dict(reversed(list(chain.items()))) if reverse else chain) + chain["c33"] = ["c32"] + with pytest.raises(CatalogError, match="32 edges"): + build(dict(reversed(list(chain.items()))) if reverse else chain) + + +@pytest.mark.asyncio +async def test_service_rejects_responses_deeper_than_resource_limit() -> None: + nested: dict[str, object] = {} + for _ in range(17): + nested = {"value": nested} + offering = Offering.model_validate({**json.loads(OFFERING), "custom_data": nested}) + service = _service(StaticCatalog(StaticCatalogOptions(offerings=(offering,)))) + reply = await service.handle(Request("GET", "/odp/offerings/rubber-plant")) + assert reply.status == 500 + + +def test_cache_classes_are_independently_configurable() -> None: + fallbacks = CacheFallbacks(search=timedelta(seconds=2), filters=timedelta(seconds=3)) + assert fallbacks.search.total_seconds() == 2 + assert fallbacks.filters.total_seconds() == 3 + assert fallbacks.sorts == timedelta(hours=1) + assert fallbacks.attribute_schema == timedelta(hours=24) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("depth,status", [(16, 200), (17, 413), (2000, 413)]) +async def test_service_enforces_request_nesting_limit(depth: int, status: int) -> None: + service = _service(SearchCatalog(_catalog_options())) + body = ( + b'{"odp_version":"1.0","query":"plant","future":' + + b"[" * (depth - 1) + + b"0" + + b"]" * (depth - 1) + + b"}" + ) + result = await service.handle( + Request( + "POST", + "/odp/offerings/search", + body=body, + headers={"content-type": "application/odp+json"}, + ) + ) + assert result.status == status + + +@pytest.mark.asyncio +async def test_search_continuation_does_not_reuse_list_fallback_entry() -> None: + document = json.loads(SERVICE_DOCUMENT) + document["operations"].append({"name": "search-offerings", "authentication": "not-required"}) + next_reference = "/odp/page?cursor=same" + page = {"odp_version": "1.0", "items": [], "next": next_reference} + old = {"odp_version": "1.0", "items": [{"id": "item", "name": "Old"}]} + new = {"odp_version": "1.0", "items": [{"id": "item", "name": "New"}]} + transport = QueueTransport( + *(response(json.dumps(value)) for value in (document, page, old, page, new)) + ) + async with ServiceClient("https://service.example", transport=transport) as client: + await client.list_offerings() + assert (await client.continue_offerings(next_reference)).items[0].name == "Old" + await client.search_offerings(OfferingSearchRequest(query="plant")) + assert (await client.continue_offerings(next_reference)).items[0].name == "New" + assert len(transport.requests) == 5 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body", [b"{", b"\xff"]) +async def test_malformed_search_body_remains_an_invalid_request(body: bytes) -> None: + result = await _service(SearchCatalog(_catalog_options())).handle( + Request( + "POST", + "/odp/offerings/search", + body=body, + headers={"content-type": "application/odp+json"}, + ) + ) + assert result.status == 400 + + +@pytest.mark.asyncio +async def test_json_decoder_recursion_failure_is_a_request_limit_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + service = _service(SearchCatalog(_catalog_options())) + + def fail(body: bytes) -> object: + raise RecursionError("decoder nesting limit") + + monkeypatch.setattr("offering_protocol.service.service.json.loads", fail) + result = await service.handle( + Request( + "POST", + "/odp/offerings/search", + body=b"[]", + headers={"content-type": "application/odp+json"}, + ) + ) + assert result.status == 413 diff --git a/tests/test_schema.py b/tests/test_schema.py index 82cc16f..3a9cac3 100644 --- a/tests/test_schema.py +++ b/tests/test_schema.py @@ -3,8 +3,9 @@ import json import pytest +from jsonschema.exceptions import SchemaError -from helpers import QueueTransport, response +from helpers import OFFERING, SERVICE_DOCUMENT, QueueTransport, response from offering_protocol.agent import AgentError, ServiceClient from offering_protocol.agent.schema import _document_url, _schema_references, resolve_schema @@ -132,7 +133,7 @@ async def test_schema_resolution_rejects_invalid_dialect_and_vocabulary() -> Non ) ), ) - with pytest.raises(AgentError, match="fragment-only reference"): + with pytest.raises(SchemaError if reference is None else AgentError): await resolve_schema( unsupported_dynamic_reference, "https://schemas.example/root.json", @@ -224,3 +225,75 @@ async def test_resolves_embedded_schema_resources() -> None: assert resolved.validator.is_valid({"plant": "rubber"}) assert not resolved.validator.is_valid({"plant": 4}) assert len(transport.requests) == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "schema", + [ + {"type": "bogus"}, + {"$ref": "#/$defs/missing"}, + {"properties": []}, + {"$vocabulary": {"https://json-schema.org/draft/2020-12/vocab/future": True}}, + {"$vocabulary": {"https://json-schema.org/draft/2020-12/vocab/format-assertion": True}}, + ], +) +async def test_invalid_schema_reports_issue_without_losing_offering( + schema: dict[str, object], +) -> None: + offering = json.loads(OFFERING) + offering.update(schema={"url": "https://schemas.example/root"}, attributes={"name": "plant"}) + async with ServiceClient( + "https://service.example", + transport=QueueTransport(response(SERVICE_DOCUMENT), response(json.dumps(offering))), + supporting_transport=QueueTransport( + response( + json.dumps({"$schema": DIALECT, **schema}), content_type="application/schema+json" + ) + ), + ) as client: + result = await client.get_offering_details("rubber-plant") + assert result.offering.name == "Rubber Plant" + assert not result.offering.attributes + assert result.attribute_schema is None + assert len(result.issues) == 1 + assert result.issues[0].scope == "attribute_schema" + + +@pytest.mark.asyncio +async def test_schema_literals_are_not_interpreted_as_schema_keywords() -> None: + literal = {"$ref": "https://unrelated.example/item", "$dynamicRef": "external.json"} + schema = { + "$schema": DIALECT, + "const": literal, + "examples": [literal], + "properties": {"value": True}, + "additionalProperties": True, + } + transport = QueueTransport(response(json.dumps(schema), content_type="application/schema+json")) + async with ServiceClient("https://service.example", supporting_transport=transport) as client: + resolved = await resolve_schema(client, "https://schemas.example/root") + assert resolved.validator.is_valid(literal) + assert len(transport.requests) == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("identifier", [None, "./root.json"]) +async def test_redirected_schema_uses_final_url_even_from_cache(identifier: str | None) -> None: + root = {"$schema": DIALECT, "$ref": "child.json", **({"$id": identifier} if identifier else {})} + child = {"$schema": DIALECT, "type": "string"} + transport = QueueTransport( + response("", status=302, headers={"location": "/v2/root.json"}), + response(json.dumps(root), content_type="application/schema+json"), + response(json.dumps(child), content_type="application/schema+json"), + ) + async with ServiceClient("https://service.example", supporting_transport=transport) as client: + for _ in range(2): + result = await resolve_schema(client, "https://schemas.example/root.json") + assert result.validator.is_valid("plant") + assert not result.validator.is_valid(1) + assert [request.url for request in transport.requests] == [ + "https://schemas.example/root.json", + "https://schemas.example/v2/root.json", + "https://schemas.example/v2/child.json", + ] diff --git a/tests/test_service_conformance.py b/tests/test_service_conformance.py index 7d42e6d..04a4b76 100644 --- a/tests/test_service_conformance.py +++ b/tests/test_service_conformance.py @@ -475,10 +475,10 @@ async def test_refuses_a_repeated_representation() -> None: @pytest.mark.asyncio async def test_refuses_a_representation_or_limit_it_cannot_honour() -> None: - for query in ("representation=sideways", "limit=101", "limit=-1", "limit=many"): + for query in ("representation=sideways", "limit=101", "limit=-1", "limit=many", "limit=0"): assert (await call("GET", "/odp/offerings", query=query)).status == 400, query - for query in ("representation=full", "representation=terse", "limit=100", "limit=0"): + for query in ("representation=full", "representation=terse", "limit=100"): assert (await call("GET", "/odp/offerings", query=query)).status == 200, query