diff --git a/README.md b/README.md index 623057e..ed7461e 100644 --- a/README.md +++ b/README.md @@ -203,9 +203,14 @@ 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. +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. ### Search across Services @@ -264,11 +269,22 @@ 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 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 @@ -278,7 +294,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`. @@ -340,10 +357,21 @@ 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. 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 +decimal or date-time strings by their meaning requires the referenced Filter Definition. + ## Errors and validation Each role exposes typed errors: @@ -357,6 +385,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/scripts/conformance_adapter.py b/scripts/conformance_adapter.py index bf1fb35..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,149 @@ 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"] if subject == "local-identifier": return is_local_resource_identifier(case["value"]) == case["valid"] if subject == "identity-comparison": @@ -315,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"])) @@ -378,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/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/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 4dca744..b2ce3a7 100644 --- a/src/offering_protocol/agent/capabilities.py +++ b/src/offering_protocol/agent/capabilities.py @@ -2,21 +2,26 @@ 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, ) +from offering_protocol.core.validation import _agent_body _MAXIMUM_CAPABILITY_PAGES = 16 _MAXIMUM_FILTERS = 1_024 @@ -122,32 +127,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 +148,117 @@ 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, 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. + + 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, frozenset(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, + existing: frozenset[str] = frozenset(), +) -> list[FilterDefinition]: + values = await _load_definitions( + client, reference, budget, parse_filter_definition_page, CapabilityKind.FILTERS, existing + ) + return cast("list[FilterDefinition]", values) -async def _load_sorts(client: ServiceClient, reference: str) -> list[SortDefinition]: - values: list[SortDefinition] = [] +async def _load_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, existing + ) + return cast("list[SortDefinition]", values) + + +async def _load_definitions( + client: ServiceClient, + reference: str, + budget: int, + parse: Callable[[bytes | str], Any], + kind: CapabilityKind, + existing: frozenset[str], +) -> list[Any]: + """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") + ) + + values: list[Any] = [] next_reference = reference visited: set[str] = set() for _ in range(_MAXIMUM_CAPABILITY_PAGES): @@ -217,11 +268,16 @@ 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 + fallback = ( + client._cache_fallbacks.filters + if kind is CapabilityKind.FILTERS + else client._cache_fallbacks.sorts ) - page = parse_sort_definition_page(body) + body = await client._linked_odp(target, fallback, parse_page) + page = parse_page(body) values.extend(page.items) + 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: raise AgentError("ODP capability source exceeded 16 pages") @@ -229,34 +285,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..fd6691c 100644 --- a/src/offering_protocol/agent/client.py +++ b/src/offering_protocol/agent/client.py @@ -6,14 +6,16 @@ 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 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,24 +29,26 @@ 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 ( - 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, _nesting_depth from offering_protocol.directory.transport import ( HttpRequest, HttpResponse, @@ -60,7 +64,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): @@ -87,6 +95,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: @@ -117,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: @@ -125,10 +137,22 @@ 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._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 @@ -137,8 +161,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" @@ -149,6 +177,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), @@ -184,10 +213,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) @@ -201,10 +227,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 @@ -239,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( @@ -345,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): @@ -384,8 +425,11 @@ async def _request_cached( maximum_bytes: int, 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: @@ -399,7 +443,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") @@ -415,7 +461,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( @@ -435,7 +481,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 +491,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: @@ -476,7 +524,27 @@ async def _supporting_json( accept: str, media_types: set[str], 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") @@ -484,29 +552,45 @@ 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: 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}) + 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") 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: @@ -518,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()) ) @@ -526,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, @@ -534,18 +618,22 @@ 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") - value = _decode_json_object(response.body) - if _cacheable("GET", response.headers, timedelta()): + if maximum_depth is not None: + _require_depth(response.body, maximum_depth, "ODP supporting document") + _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, @@ -554,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: @@ -584,37 +672,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=(",", ":")) - - -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() + return parse_problem_response_strict(_agent_body(data, "problem"), status) def _append_query(target: str, values: dict[str, str]) -> str: @@ -634,27 +708,57 @@ 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: - if len(response.body) > maximum_bytes: - raise AgentError("ODP response exceeds its byte limit") +def _consume(response: HttpResponse, maximum_bytes: int, maximum_depth: int) -> HttpResponse: 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) + raise ServiceRequestError(response.status, _problem_message(response), response.headers) + if len(response.body) > maximum_bytes: + 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}") + _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 _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 +803,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..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 ( @@ -17,6 +20,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): @@ -93,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) @@ -120,6 +127,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/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/__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..cca31a4 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, @@ -27,7 +29,9 @@ OfferingSearchRequest, Operation, Page, + PriceType, ProblemDetails, + RefinementGroup, ResourceIdentity, ServiceDocument, SortDefinition, @@ -274,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): @@ -286,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"} @@ -321,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) @@ -440,6 +462,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 +486,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 +520,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 +575,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 +613,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 +684,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]: @@ -571,7 +735,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) @@ -745,6 +915,52 @@ def _is_language_tag(value: str) -> bool: return not in_extension or len(subtags[-1]) > 1 +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) -> tuple[str, object]: + """Check structural duplicates without guessing the type of a string-valued Filter.""" + if isinstance(value, bool): + return "boolean", value + if isinstance(value, int | float): + return "number", Decimal(str(value)) + if isinstance(value, str): + return "string", value + return "other", 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..117c541 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,10 @@ _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})$" +) class DirectoryError(RuntimeError): @@ -98,7 +108,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 +204,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 +221,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[DirectoryService] = [] + issues: list[ServiceIssue] = [] + for index, entry in enumerate(raw["items"]): + try: + items.append(DirectoryService.model_validate(_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. + 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] + 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 +346,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/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/directory/transport.py b/src/offering_protocol/directory/transport.py index 02a0a93..20e40eb 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: @@ -18,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) @@ -30,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: ... @@ -47,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 ) @@ -76,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: @@ -93,9 +124,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 +141,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..ae28633 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 @@ -41,11 +43,31 @@ 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" _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 +91,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 +117,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 +263,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 +272,131 @@ 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=self._document.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, operation) + 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, 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= + # 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")) + 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 - 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=parameters.get("cursor"), + cursor=values.get("cursor"), + language=language, limit=limit, path=request.path, representation=representation, @@ -349,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 @@ -370,11 +440,69 @@ 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) + 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( + 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": 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,27 +554,147 @@ 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: + 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() 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 + 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 + 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}") + + +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 True - return any( - item.split(";", 1)[0].strip().lower() in {"*/*", MEDIA_TYPE} for item in value.split(",") + 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..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: @@ -205,8 +233,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 +275,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..9b7d60a 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, ) @@ -277,21 +276,24 @@ 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 = """{ "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" }] }""" @@ -448,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"), @@ -477,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"), ), ) @@ -523,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"), ), ) @@ -565,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 @@ -595,7 +605,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 +630,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) @@ -662,8 +675,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"{}") @@ -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"), ), ) @@ -834,7 +846,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( @@ -899,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 new file mode 100644 index 0000000..c24871a --- /dev/null +++ b/tests/test_agent_conformance.py @@ -0,0 +1,676 @@ +"""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, +) +from offering_protocol.core import ( + CapabilityLink, + CollectionSearchRequest, + FilterCapabilitySource, + FilterDefinition, + FilterOperator, + FilterType, + MissingPlacement, + SearchCapabilities, + SortCapabilitySource, + SortDefinition, + SortDirection, + SortKey, +) +from offering_protocol.core.validation import _nesting_depth +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_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.""" + 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"