diff --git a/README.md b/README.md index 3db196b..22e71a8 100644 --- a/README.md +++ b/README.md @@ -101,6 +101,45 @@ Use a separate unprotected application branch for public resources. For local de `allow_insecure_loopback=True` permits an HTTP `localhost` or loopback resource origin; production origins require HTTPS. +## Agent with a hosted identity Platform + +`PlatformIdentityProvider` lets an Agent use a remote AEP Platform for Service-scoped identity +custody and delegated assertion signing. It discovers the Platform, recovers an existing active +identity before provisioning one, caches discovery metadata according to HTTP cache directives, +and supplies the resulting signer directly to `Agent`. + +```python +import os + +from agent_enrollment_protocol.agent import ( + Agent, + AgentOptions, + PlatformIdentityProvider, + PlatformIdentityProviderOptions, +) + + +async def authentication_headers() -> dict[str, str]: + return {"Authorization": f"Bearer {os.environ['AEP_PLATFORM_ACCESS_TOKEN']}"} + + +async with PlatformIdentityProvider( + PlatformIdentityProviderOptions( + authentication_headers=authentication_headers, + platform_url="https://platform.example", + ) +) as identities: + async with Agent(AgentOptions(identity_provider=identities)) as agent: + result = await agent.service("https://service.example").enroll() +``` + +The Platform authentication callback is evaluated for each private request so applications can +refresh short-lived credentials. Supply `pending_sign_resolver` when the Platform can return +`202 Accepted` during delegated signing. The resolver receives the immutable retry interval and +opaque Platform context; returning updated context starts the next signing stage with a distinct +idempotency key. Without a resolver, pending signing raises `PlatformSignPendingError` for the +application to continue explicitly. + ## Hosted identity Platform `agent_enrollment_protocol.platform` implements discovery, Service-scoped Agent identity diff --git a/scripts/verify-consumer.sh b/scripts/verify-consumer.sh index 1138fe2..9fe9a3f 100755 --- a/scripts/verify-consumer.sh +++ b/scripts/verify-consumer.sh @@ -15,8 +15,18 @@ from importlib.metadata import version from agent_enrollment_protocol import adapters, agent, core, platform, service from agent_enrollment_protocol.adapters import AepAsgiApplication, AepAuthenticationMiddleware -from agent_enrollment_protocol.agent import Agent, AgentOptions, HttpxTransport, ServiceIdentity -from agent_enrollment_protocol.core import ClaimSupportEvaluation, evaluate_claim_support +from agent_enrollment_protocol.agent import ( + Agent, + AgentOptions, + HttpxTransport, + PlatformIdentityProvider, + ServiceIdentity, +) +from agent_enrollment_protocol.core import ( + AEP_PLATFORM_WELL_KNOWN_PATH, + ClaimSupportEvaluation, + evaluate_claim_support, +) from agent_enrollment_protocol.service import ( MemoryServiceCredentialStore, Service, @@ -35,6 +45,8 @@ assert service.__name__ == "agent_enrollment_protocol.service" assert Agent.__module__ == "agent_enrollment_protocol.agent.client" assert AgentOptions.__module__ == "agent_enrollment_protocol.agent.client" assert HttpxTransport.__module__ == "agent_enrollment_protocol.agent.transport" +assert PlatformIdentityProvider.__module__ == "agent_enrollment_protocol.agent.platform_provider" +assert AEP_PLATFORM_WELL_KNOWN_PATH == "/.well-known/aep-platform" assert ServiceIdentity.__module__ == "agent_enrollment_protocol.agent.types" assert Service.__module__ == "agent_enrollment_protocol.service.service" assert ServiceOptions.__module__ == "agent_enrollment_protocol.service.types" diff --git a/src/agent_enrollment_protocol/agent/__init__.py b/src/agent_enrollment_protocol/agent/__init__.py index 6798942..aaf086f 100644 --- a/src/agent_enrollment_protocol/agent/__init__.py +++ b/src/agent_enrollment_protocol/agent/__init__.py @@ -1,6 +1,20 @@ """Agent-side enrollment, credential, and authentication workflows.""" from .client import Agent, AgentOptions, ServiceSession +from .platform_provider import ( + MemoryPlatformDiscoveryCache, + PlatformAuthenticationHeaders, + PlatformCommandError, + PlatformContextProvider, + PlatformDiscoveryCache, + PlatformDiscoveryCacheEntry, + PlatformIdempotencyKeyFactory, + PlatformIdentityProvider, + PlatformIdentityProviderOptions, + PlatformPendingSign, + PlatformPendingSignResolver, + PlatformSignPendingError, +) from .stores import ( MemoryCredentialStore, MemoryIdentityStore, @@ -61,7 +75,19 @@ "MemoryCredentialStore", "MemoryIdentityStore", "MemoryInspectCache", + "MemoryPlatformDiscoveryCache", "OperationKey", + "PlatformAuthenticationHeaders", + "PlatformCommandError", + "PlatformContextProvider", + "PlatformDiscoveryCache", + "PlatformDiscoveryCacheEntry", + "PlatformIdempotencyKeyFactory", + "PlatformIdentityProvider", + "PlatformIdentityProviderOptions", + "PlatformPendingSign", + "PlatformPendingSignResolver", + "PlatformSignPendingError", "RandomIdempotencyKeyProvider", "RevokeOptions", "ServiceIdentity", diff --git a/src/agent_enrollment_protocol/agent/platform_provider.py b/src/agent_enrollment_protocol/agent/platform_provider.py new file mode 100644 index 0000000..095fca1 --- /dev/null +++ b/src/agent_enrollment_protocol/agent/platform_provider.py @@ -0,0 +1,665 @@ +from __future__ import annotations + +import asyncio +import json +import math +import secrets +from collections.abc import Awaitable, Callable, Mapping +from copy import deepcopy +from dataclasses import dataclass, field +from datetime import UTC, datetime, timedelta +from types import MappingProxyType +from typing import Protocol, TypeVar +from urllib.parse import quote, urlencode, urljoin, urlsplit, urlunsplit + +from agent_enrollment_protocol.core import ( + AEP_MEDIA_TYPE, + AEP_PLATFORM_WELL_KNOWN_PATH, + AEP_PROBLEM_MEDIA_TYPE, + AepValidationError, + ClientAssertionClaims, + HttpRequest, + HttpResponse, + ManagedAgentStatus, + PlatformAgentIdentity, + PlatformAgentIdentityListResponse, + PlatformDiscoveryDocument, + PlatformProvisionRequest, + PlatformSignCompleted, + PlatformSignPending, + PlatformSignRequest, + ProblemDetails, + SigningAlgorithm, + did_web_document_url, + media_type_essence, + parse_json_model, + parse_platform_sign_response, + same_origin, +) + +from .transport import AsyncHttpTransport, HttpxTransport +from .types import AssertionSigner, IdentityRequest, ServiceIdentity + +Clock = Callable[[], datetime] +PlatformAuthenticationHeaders = Callable[[], Awaitable[Mapping[str, str]]] +PlatformIdempotencyKeyFactory = Callable[[], Awaitable[str]] +PlatformContextProvider = Callable[ + [ServiceIdentity, ClientAssertionClaims], Awaitable[Mapping[str, object] | None] +] +ResponseT = TypeVar("ResponseT") + + +@dataclass(frozen=True, slots=True) +class PlatformDiscoveryCacheEntry: + cached_at: datetime + document: PlatformDiscoveryDocument + final_url: str + cache_control: str = "" + etag: str = "" + last_modified: str = "" + + +class PlatformDiscoveryCache(Protocol): + async def delete(self, key: str) -> None: ... + + async def find(self, key: str) -> PlatformDiscoveryCacheEntry | None: ... + + async def save(self, key: str, entry: PlatformDiscoveryCacheEntry) -> None: ... + + +class MemoryPlatformDiscoveryCache: + def __init__(self) -> None: + self._lock = asyncio.Lock() + self._records: dict[str, PlatformDiscoveryCacheEntry] = {} + + async def delete(self, key: str) -> None: + async with self._lock: + self._records.pop(key, None) + + async def find(self, key: str) -> PlatformDiscoveryCacheEntry | None: + async with self._lock: + entry = self._records.get(key) + return _copy_cache_entry(entry) if entry is not None else None + + async def save(self, key: str, entry: PlatformDiscoveryCacheEntry) -> None: + async with self._lock: + self._records[key] = _copy_cache_entry(entry) + + +@dataclass(frozen=True, slots=True) +class PlatformPendingSign: + identity: ServiceIdentity + platform_context: Mapping[str, object] + retry_after_seconds: int + + def __post_init__(self) -> None: + if self.retry_after_seconds < 1 or self.retry_after_seconds > 300: + raise ValueError("retry_after_seconds must be between 1 and 300") + object.__setattr__(self, "platform_context", _copy_context(self.platform_context)) + + +PlatformPendingSignResolver = Callable[ + [PlatformPendingSign], Awaitable[Mapping[str, object] | None] +] + + +class PlatformSignPendingError(Exception): + def __init__(self, pending: PlatformPendingSign) -> None: + super().__init__("AEP Platform signing is pending") + self.pending = pending + + +class PlatformCommandError(Exception): + def __init__(self, status: int, problem: ProblemDetails | None = None) -> None: + message = ( + f"AEP Platform command failed: {problem.title}" + if problem is not None + else f"AEP Platform command failed with HTTP {status}" + ) + super().__init__(message) + self.status = status + self.problem = problem + + +@dataclass(frozen=True, slots=True) +class PlatformIdentityProviderOptions: + platform_url: str + allow_insecure_loopback: bool = False + authentication_headers: PlatformAuthenticationHeaders | None = None + authorization: str | None = field(default=None, repr=False) + clock: Clock = lambda: datetime.now(UTC) + discovery_cache: PlatformDiscoveryCache | None = None + idempotency_key: PlatformIdempotencyKeyFactory | None = None + maximum_response_bytes: int = 1 << 20 + pending_sign_resolver: PlatformPendingSignResolver | None = None + platform_context: PlatformContextProvider | None = None + request_timeout: float = 30.0 + transport: AsyncHttpTransport | None = None + + +class PlatformIdentityProvider: + def __init__(self, options: PlatformIdentityProviderOptions) -> None: + if options.maximum_response_bytes < 1: + raise ValueError("AEP Platform maximum response bytes must be positive") + if options.request_timeout <= 0 or not math.isfinite(options.request_timeout): + raise ValueError("AEP Platform request timeout must be positive and finite") + self._allow_insecure_loopback = options.allow_insecure_loopback + self._authentication_headers = options.authentication_headers + self._authorization = options.authorization + self._clock = _validated_clock(options.clock) + self._discovery_cache = options.discovery_cache or MemoryPlatformDiscoveryCache() + self._idempotency_key = options.idempotency_key or _random_idempotency_key + self._maximum_response_bytes = options.maximum_response_bytes + self._pending_sign_resolver = options.pending_sign_resolver + self._platform_context = options.platform_context + self._platform_url = _platform_url(options.platform_url, options.allow_insecure_loopback) + self._request_timeout = options.request_timeout + self._transport = options.transport or HttpxTransport( + maximum_response_bytes=options.maximum_response_bytes + ) + self._owns_transport = options.transport is None + self._discovery_lock = asyncio.Lock() + self._identity_lock = asyncio.Lock() + self._clock() + + async def __aenter__(self) -> PlatformIdentityProvider: + return self + + async def __aexit__(self, *args: object) -> None: + await self.aclose() + + async def aclose(self) -> None: + if self._owns_transport: + await self._transport.aclose() + + async def find_identity_by_service_did(self, service_did: str) -> ServiceIdentity | None: + PlatformProvisionRequest(service_did=service_did) + discovery = await self._discover() + endpoint = _endpoint(self._platform_url, discovery.document.endpoints.list) + query = urlencode( + { + "descending": "true", + "limit": "100", + "service_did": service_did, + } + ) + _, _, listed = await self._command( + "GET", + f"{endpoint}?{query}", + None, + None, + lambda body: parse_json_model( + body, PlatformAgentIdentityListResponse, "Platform identity list" + ), + ) + for candidate in listed.data: + _validate_platform_identity(candidate, self._allow_insecure_loopback) + if ( + candidate.service_did == service_did + and candidate.status is ManagedAgentStatus.ACTIVE + ): + return self._service_identity(candidate) + return None + + async def get_or_create_identity(self, request: IdentityRequest) -> ServiceIdentity: + if request.service_did != request.inspect.service.did: + raise ValueError("AEP identity request does not match the inspected Service") + if "did:web" not in request.inspect.identity.methods: + raise ValueError("AEP Service does not support Platform-hosted did:web identities") + async with self._identity_lock: + existing = await self.find_identity_by_service_did(request.service_did) + if existing is not None: + return existing + discovery = await self._discover() + endpoint = _endpoint(self._platform_url, discovery.document.endpoints.provision) + key = await self._new_idempotency_key() + provision = PlatformProvisionRequest(service_did=request.service_did) + _, _, created = await self._command( + "POST", + endpoint, + key, + provision.to_wire(), + lambda body: parse_json_model( + body, PlatformAgentIdentity, "Platform provision response" + ), + ) + _validate_platform_identity(created, self._allow_insecure_loopback) + if ( + created.service_did != request.service_did + or created.status is not ManagedAgentStatus.ACTIVE + ): + raise ValueError( + "AEP Platform provisioned an identity outside the requested Service scope" + ) + identity = self._service_identity(created) + return identity + + async def signer_for(self, identity: ServiceIdentity) -> AssertionSigner: + self._validate_owned_identity(identity) + + async def signer( + claims: ClientAssertionClaims, algorithms: tuple[SigningAlgorithm, ...] + ) -> str: + if ( + claims.iss != identity.agent_did + or claims.sub != identity.agent_did + or claims.aud != identity.service_did + ): + raise ValueError("AEP Platform signer received claims for another identity") + if not set(identity.signing_algorithms).intersection(algorithms): + raise ValueError("AEP Platform and Service have no compatible signing algorithm") + context: Mapping[str, object] | None = None + if self._platform_context is not None: + context = await self._platform_context(identity, claims) + previous_key: str | None = None + while True: + key = await self._new_idempotency_key() + if key == previous_key: + raise ValueError( + "AEP Platform pending Sign stages require distinct idempotency keys" + ) + response = await self._sign(identity, claims, context, key) + if isinstance(response, PlatformSignCompleted): + return response.client_assertion + pending = PlatformPendingSign( + identity=identity, + platform_context=response.platform_context or {}, + retry_after_seconds=int(response.retry_after_seconds), + ) + if self._pending_sign_resolver is None: + raise PlatformSignPendingError(pending) + previous_key = key + context = await self._pending_sign_resolver(pending) + + return signer + + async def _sign( + self, + identity: ServiceIdentity, + claims: ClientAssertionClaims, + platform_context: Mapping[str, object] | None, + idempotency_key: str, + ) -> PlatformSignCompleted | PlatformSignPending: + discovery = await self._discover() + agent_identity_id = identity.metadata.get("agent_identity_id", "") + endpoint = _endpoint( + self._platform_url, + discovery.document.endpoints.sign, + agent_identity_id=agent_identity_id, + ) + request_data: dict[str, object] = { + "jti": claims.jti, + "lifetime_seconds": str(claims.exp - claims.iat), + "op": claims.op, + "service_did": claims.aud, + } + if platform_context is not None: + request_data["platform_context"] = deepcopy(dict(platform_context)) + if claims.resource is not None: + request_data["resource"] = claims.resource + request = PlatformSignRequest.model_validate(request_data) + status, headers, response = await self._command( + "POST", + endpoint, + idempotency_key, + request.to_wire(), + parse_platform_sign_response, + ) + if _header(headers, "retry-after") is not None: + raise ValueError("AEP Platform Sign response included Retry-After") + if isinstance(response, PlatformSignPending): + if status != 202: + raise ValueError("AEP Platform returned an invalid pending Sign status") + return response + if status != 200 or not _valid_completed_sign(response, identity, claims): + raise ValueError("AEP Platform returned an invalid completed Sign response") + return response + + async def _discover(self) -> PlatformDiscoveryCacheEntry: + async with self._discovery_lock: + discovery_url = urljoin(self._platform_url, AEP_PLATFORM_WELL_KNOWN_PATH) + cached = await self._discovery_cache.find(discovery_url) + now = self._clock() + if cached is not None: + try: + cached = _validate_cache_entry( + cached, + discovery_url, + self._allow_insecure_loopback, + ) + except ValueError: + await self._discovery_cache.delete(discovery_url) + cached = None + if cached is not None and _cache_fresh(cached, now): + return cached + current = cached.final_url if cached is not None else discovery_url + headers = {"Accept": AEP_MEDIA_TYPE} + if cached is not None: + if cached.etag: + headers["If-None-Match"] = cached.etag + if cached.last_modified: + headers["If-Modified-Since"] = cached.last_modified + redirects = 0 + while True: + response = await self._send(HttpRequest(method="GET", url=current, headers=headers)) + if response.status in {301, 302, 303, 307, 308}: + location = _header(response.headers, "location") + if location is None: + raise ValueError("AEP Platform discovery redirect omitted Location") + if redirects >= 5: + raise ValueError("AEP Platform discovery exceeded five redirects") + target = urljoin(current, location) + if not _valid_url(target, self._allow_insecure_loopback) or not same_origin( + current, target + ): + raise ValueError("AEP Platform discovery redirect changed origin or scheme") + current = target + redirects += 1 + continue + entry = self._discovery_response(response, current, now, cached) + if _cache_directive(entry.cache_control, "no-store") is not None: + await self._discovery_cache.delete(discovery_url) + else: + await self._discovery_cache.save(discovery_url, entry) + return entry + + def _discovery_response( + self, + response: HttpResponse, + final_url: str, + now: datetime, + cached: PlatformDiscoveryCacheEntry | None, + ) -> PlatformDiscoveryCacheEntry: + if response.status == 304: + if cached is None: + raise ValueError("AEP Platform discovery returned 304 without a cached document") + return PlatformDiscoveryCacheEntry( + cached_at=now, + document=cached.document, + final_url=final_url, + cache_control=_header(response.headers, "cache-control") or cached.cache_control, + etag=_header(response.headers, "etag") or cached.etag, + last_modified=_header(response.headers, "last-modified") or cached.last_modified, + ) + if response.status < 200 or response.status >= 300: + raise PlatformCommandError(response.status) + if media_type_essence(_header(response.headers, "content-type") or "") != AEP_MEDIA_TYPE: + raise ValueError("AEP Platform discovery response media type is invalid") + _bounded(response.body, self._maximum_response_bytes) + document = parse_json_model( + response.body, PlatformDiscoveryDocument, "Platform discovery document" + ) + if "did:web" not in document.identity.did_methods: + raise ValueError("AEP Platform does not advertise did:web") + return PlatformDiscoveryCacheEntry( + cached_at=now, + document=document, + final_url=final_url, + cache_control=_header(response.headers, "cache-control") or "", + etag=_header(response.headers, "etag") or "", + last_modified=_header(response.headers, "last-modified") or "", + ) + + async def _command( + self, + method: str, + url: str, + idempotency_key: str | None, + body: Mapping[str, object] | None, + parser: Callable[[bytes], ResponseT], + ) -> tuple[int, Mapping[str, str], ResponseT]: + headers = await self._headers() + headers["Accept"] = AEP_MEDIA_TYPE + data = None + if body is not None: + data = json.dumps(body, separators=(",", ":")).encode() + headers["Content-Type"] = AEP_MEDIA_TYPE + if idempotency_key is not None: + headers["Idempotency-Key"] = idempotency_key + response = await self._send(HttpRequest(method=method, url=url, headers=headers, body=data)) + _bounded(response.body, self._maximum_response_bytes) + if response.status < 200 or response.status >= 300: + problem = None + if ( + media_type_essence(_header(response.headers, "content-type") or "") + == AEP_PROBLEM_MEDIA_TYPE + ): + try: + candidate = parse_json_model(response.body, ProblemDetails, "Problem Details") + if candidate.status == response.status: + problem = candidate + except AepValidationError: + pass + raise PlatformCommandError(response.status, problem) + if media_type_essence(_header(response.headers, "content-type") or "") != AEP_MEDIA_TYPE: + raise ValueError("AEP Platform response media type is invalid") + return response.status, response.headers, parser(response.body) + + async def _headers(self) -> dict[str, str]: + headers: dict[str, str] = {} + if self._authorization: + headers["Authorization"] = self._authorization + if self._authentication_headers is not None: + supplied = await self._authentication_headers() + for name, value in supplied.items(): + if name.lower() in {"accept", "content-type", "idempotency-key"}: + continue + _set_header(headers, name, value) + return headers + + async def _send(self, request: HttpRequest) -> HttpResponse: + return await asyncio.wait_for(self._transport.send(request), timeout=self._request_timeout) + + async def _new_idempotency_key(self) -> str: + key = await self._idempotency_key() + if not key.strip(): + raise ValueError("AEP Platform idempotency key generation failed") + return key + + def _service_identity(self, value: PlatformAgentIdentity) -> ServiceIdentity: + return ServiceIdentity( + agent_did=value.agent_did, + identity_method="did:web", + service_did=value.service_did, + signing_algorithms=value.signing_algorithms, + metadata={ + "agent_identity_id": value.agent_identity_id, + "created_at": value.created_at, + "did_document_url": value.did_document_url, + "key_id": value.key_id, + "platform_url": self._platform_url, + "status": value.status.value, + "updated_at": value.updated_at, + }, + ) + + def _validate_owned_identity(self, identity: ServiceIdentity) -> None: + if ( + identity.metadata.get("platform_url") != self._platform_url + or not identity.metadata.get("agent_identity_id") + or identity.metadata.get("status") != ManagedAgentStatus.ACTIVE.value + or identity.identity_method != "did:web" + or not identity.agent_did.startswith("did:web:") + or identity.metadata.get("key_id") != identity.agent_did + or not identity.signing_algorithms + ): + raise ValueError("AEP identity is not an active identity from this Platform") + expected = did_web_document_url( + identity.agent_did, + allow_insecure_loopback=self._allow_insecure_loopback, + ) + if identity.metadata.get("did_document_url") != expected: + raise ValueError("AEP Platform DID document URL does not match the Agent DID") + + +async def _random_idempotency_key() -> str: + return secrets.token_hex(16) + + +def _platform_url(value: str, allow_insecure_loopback: bool) -> str: + raw = value.strip() + if not raw: + raise ValueError("invalid AEP Platform URL") + if "://" not in raw: + raw = f"https://{raw}" + if not _valid_url(raw, allow_insecure_loopback): + raise ValueError("AEP Platform URLs require HTTPS") + parsed = urlsplit(raw) + return urlunsplit((parsed.scheme, parsed.netloc, "/", "", "")) + + +def _valid_url(value: str, allow_insecure_loopback: bool) -> bool: + parsed = urlsplit(value) + if parsed.username or parsed.password or parsed.fragment or not parsed.hostname: + return False + if parsed.scheme == "https": + return True + return ( + allow_insecure_loopback + and parsed.scheme == "http" + and parsed.hostname in {"localhost", "127.0.0.1", "::1"} + ) + + +def _endpoint(platform_url: str, path: str, *, agent_identity_id: str | None = None) -> str: + if agent_identity_id is not None: + path = path.replace("{agent_identity_id}", quote(agent_identity_id, safe="")) + parsed = urlsplit(path) + if ( + not path.startswith("/") + or path.startswith("//") + or parsed.scheme + or parsed.netloc + or parsed.query + or parsed.fragment + or "{" in path + ): + raise ValueError("AEP Platform advertised an invalid endpoint") + return urljoin(platform_url, path) + + +def _validate_platform_identity( + value: PlatformAgentIdentity, allow_insecure_loopback: bool +) -> None: + if ( + not value.agent_did.startswith("did:web:") + or value.key_id != value.agent_did + or not value.signing_algorithms + ): + raise ValueError("AEP Platform returned an invalid identity") + expected = did_web_document_url( + value.agent_did, allow_insecure_loopback=allow_insecure_loopback + ) + if value.did_document_url != expected: + raise ValueError("AEP Platform DID document URL does not match the Agent DID") + + +def _valid_completed_sign( + response: PlatformSignCompleted, + identity: ServiceIdentity, + claims: ClientAssertionClaims, +) -> bool: + issued_at = datetime.fromisoformat(response.issued_at.replace("Z", "+00:00")) + expires_at = datetime.fromisoformat(response.expires_at.replace("Z", "+00:00")) + return ( + response.agent_did == identity.agent_did + and response.service_did == identity.service_did + and response.jti == claims.jti + and int((expires_at - issued_at).total_seconds()) == claims.exp - claims.iat + ) + + +def _copy_cache_entry(entry: PlatformDiscoveryCacheEntry) -> PlatformDiscoveryCacheEntry: + return PlatformDiscoveryCacheEntry( + cached_at=entry.cached_at, + document=entry.document.model_copy(deep=True), + final_url=entry.final_url, + cache_control=entry.cache_control, + etag=entry.etag, + last_modified=entry.last_modified, + ) + + +def _copy_context(value: Mapping[str, object]) -> Mapping[str, object]: + return MappingProxyType(deepcopy(dict(value))) + + +def _validate_cache_entry( + entry: PlatformDiscoveryCacheEntry, + discovery_url: str, + allow_insecure_loopback: bool, +) -> PlatformDiscoveryCacheEntry: + if entry.cached_at.utcoffset() is None: + raise ValueError("cached AEP Platform discovery timestamp has no UTC offset") + if not _valid_url(entry.final_url, allow_insecure_loopback) or not same_origin( + entry.final_url, discovery_url + ): + raise ValueError("cached AEP Platform discovery URL is invalid") + document = parse_json_model( + json.dumps(entry.document.to_wire()), + PlatformDiscoveryDocument, + "Platform discovery document", + ) + if "did:web" not in document.identity.did_methods: + raise ValueError("AEP Platform does not advertise did:web") + return PlatformDiscoveryCacheEntry( + cached_at=entry.cached_at, + document=document, + final_url=entry.final_url, + cache_control=entry.cache_control, + etag=entry.etag, + last_modified=entry.last_modified, + ) + + +def _cache_fresh(entry: PlatformDiscoveryCacheEntry, now: datetime) -> bool: + if any( + _cache_directive(entry.cache_control, directive) is not None + for directive in ("no-cache", "no-store") + ): + return False + maximum_age = _cache_directive(entry.cache_control, "max-age") + if maximum_age is None: + seconds = 300 + else: + try: + seconds = int(maximum_age.strip('"')) + except ValueError: + return False + if seconds < 0: + return False + return now < entry.cached_at + timedelta(seconds=seconds) + + +def _cache_directive(value: str, name: str) -> str | None: + for item in value.split(","): + key, separator, content = item.strip().partition("=") + if key.lower() == name.lower(): + return content if separator else "" + return None + + +def _header(headers: Mapping[str, str], name: str) -> str | None: + return next((value for key, value in headers.items() if key.lower() == name.lower()), None) + + +def _set_header(headers: dict[str, str], name: str, value: str) -> None: + existing = next((key for key in headers if key.lower() == name.lower()), None) + if existing is not None: + del headers[existing] + headers[name] = value + + +def _bounded(body: bytes, maximum: int) -> None: + if len(body) > maximum: + raise ValueError("AEP Platform response exceeds the configured limit") + + +def _validated_clock(clock: Clock) -> Clock: + def current() -> datetime: + value = clock() + if value.utcoffset() is None: + raise ValueError("AEP Platform clock must return an offset-aware datetime") + return value + + return current diff --git a/src/agent_enrollment_protocol/core/__init__.py b/src/agent_enrollment_protocol/core/__init__.py index 984547a..7dfc12f 100644 --- a/src/agent_enrollment_protocol/core/__init__.py +++ b/src/agent_enrollment_protocol/core/__init__.py @@ -29,6 +29,7 @@ AEP_GRANT_TYPE_OAUTH_BEARER, AEP_IDENTITY_METHOD_DID_WEB, AEP_MEDIA_TYPE, + AEP_PLATFORM_WELL_KNOWN_PATH, AEP_PROBLEM_MEDIA_TYPE, AEP_SIGNING_ALGORITHMS, AEP_VERSION, @@ -149,6 +150,7 @@ "AEP_GRANT_TYPE_OAUTH_BEARER", "AEP_IDENTITY_METHOD_DID_WEB", "AEP_MEDIA_TYPE", + "AEP_PLATFORM_WELL_KNOWN_PATH", "AEP_PROBLEM_MEDIA_TYPE", "AEP_SIGNING_ALGORITHMS", "AEP_VERSION", diff --git a/src/agent_enrollment_protocol/core/constants.py b/src/agent_enrollment_protocol/core/constants.py index f81b08b..f0a83fe 100644 --- a/src/agent_enrollment_protocol/core/constants.py +++ b/src/agent_enrollment_protocol/core/constants.py @@ -3,6 +3,7 @@ AEP_VERSION: Final = "1.0" AEP_MEDIA_TYPE: Final = "application/aep+json" AEP_PROBLEM_MEDIA_TYPE: Final = "application/problem+json" +AEP_PLATFORM_WELL_KNOWN_PATH: Final = "/.well-known/aep-platform" AEP_AUTH_SCHEME: Final = "AEP" AEP_AUTHORIZATION_HEADER: Final = "AEP-Authorization" AEP_WELL_KNOWN_PATH: Final = "/.well-known/aep" diff --git a/src/agent_enrollment_protocol/platform/document.py b/src/agent_enrollment_protocol/platform/document.py index 62912a0..e2832ef 100644 --- a/src/agent_enrollment_protocol/platform/document.py +++ b/src/agent_enrollment_protocol/platform/document.py @@ -4,6 +4,7 @@ from urllib.parse import quote, urlsplit from agent_enrollment_protocol.core import ( + AEP_PLATFORM_WELL_KNOWN_PATH, AEP_VERSION, PlatformDiscoveryDocument, PlatformEndpoints, @@ -18,7 +19,7 @@ DID_MEDIA_TYPE = "application/did+json" HOSTED_IDENTITY_DRAFT = "draft-kavian-aep-platform-hosted-identity-01" -WELL_KNOWN_PATH = "/.well-known/aep-platform" +WELL_KNOWN_PATH = AEP_PLATFORM_WELL_KNOWN_PATH _DID_CONTEXT = "https://www.w3.org/ns/did/v1" _DID_PLACEHOLDER = "{agent_did_id}" diff --git a/tests/test_agent_platform_provider.py b/tests/test_agent_platform_provider.py new file mode 100644 index 0000000..903686e --- /dev/null +++ b/tests/test_agent_platform_provider.py @@ -0,0 +1,768 @@ +from __future__ import annotations + +import asyncio +import json +from collections.abc import Callable, Mapping +from dataclasses import replace +from datetime import UTC, datetime, timedelta +from typing import cast +from urllib.parse import parse_qs, urlsplit + +import pytest + +from agent_enrollment_protocol.agent import ( + MemoryPlatformDiscoveryCache, + PlatformAuthenticationHeaders, + PlatformCommandError, + PlatformContextProvider, + PlatformDiscoveryCache, + PlatformDiscoveryCacheEntry, + PlatformIdempotencyKeyFactory, + PlatformIdentityProvider, + PlatformIdentityProviderOptions, + PlatformPendingSign, + PlatformPendingSignResolver, + PlatformSignPendingError, + ServiceIdentity, +) +from agent_enrollment_protocol.agent.platform_provider import ( + _cache_directive, + _cache_fresh, + _endpoint, + _platform_url, + _valid_url, + _validate_cache_entry, +) +from agent_enrollment_protocol.agent.types import IdentityRequest +from agent_enrollment_protocol.core import ( + AEP_MEDIA_TYPE, + AEP_PROBLEM_MEDIA_TYPE, + AssertionOperation, + ClientAssertionClaims, + HttpRequest, + HttpResponse, + InspectDocument, + PlatformDiscoveryDocument, + SigningAlgorithm, +) + +from .test_core_models import inspect_document + +NOW = datetime(2026, 9, 3, 12, tzinfo=UTC) +SERVICE_DID = "did:web:api.example.com" +AGENT_DID = "did:web:platform.example:agents:one" + + +class QueueTransport: + def __init__(self, *responses: HttpResponse) -> None: + self.closed = False + self.requests: list[HttpRequest] = [] + self.responses = list(responses) + + async def send(self, request: HttpRequest) -> HttpResponse: + self.requests.append(request) + if not self.responses: + raise AssertionError("unexpected request") + return self.responses.pop(0) + + async def aclose(self) -> None: + self.closed = True + + +class BlockingTransport: + async def send(self, request: HttpRequest) -> HttpResponse: + del request + await asyncio.sleep(1) + raise AssertionError("unreachable") + + async def aclose(self) -> None: + return None + + +def json_response( + body: object, + *, + status: int = 200, + headers: dict[str, str] | None = None, +) -> HttpResponse: + return HttpResponse( + status=status, + headers={"Content-Type": AEP_MEDIA_TYPE, **(headers or {})}, + body=json.dumps(body, separators=(",", ":")).encode(), + ) + + +def discovery(**changes: object) -> dict[str, object]: + value: dict[str, object] = { + "aep_version": "1.0", + "endpoints": { + "hosted_verification": "/v1/aep/verifications", + "lifecycle": "/v1/aep/agent-identities/{agent_identity_id}", + "list": "/v1/aep/agent-identities", + "provision": "/v1/aep/agent-identities", + "sign": "/v1/aep/agent-identities/{agent_identity_id}/sign", + }, + "http": {"endpoint_base": "/v1/aep"}, + "identity": { + "did_methods": ["did:web"], + "did_url_template": "https://platform.example/agents/{agent_did_id}/did.json", + }, + "platform": { + "did": "did:web:platform.example", + "hosted_verification": True, + "name": "Example Platform", + }, + "signing": {"algorithms": ["ES256"], "default_lifetime_seconds": "300"}, + } + value.update(changes) + return value + + +def identity(**changes: object) -> dict[str, object]: + value: dict[str, object] = { + "agent_did": AGENT_DID, + "agent_identity_id": "pai_one", + "created_at": "2026-09-03T12:00:00Z", + "did_document_url": "https://platform.example/agents/one/did.json", + "key_id": AGENT_DID, + "service_did": SERVICE_DID, + "signing_algorithms": ["ES256"], + "status": "active", + "updated_at": "2026-09-03T12:00:00Z", + } + value.update(changes) + return value + + +def listed(*values: dict[str, object]) -> dict[str, object]: + return {"count": str(len(values)), "data": list(values), "total": str(len(values))} + + +def request( + *, + inspect: InspectDocument | None = None, + service_did: str | None = None, + service_url: str = "https://api.example.com/", +) -> IdentityRequest: + document = inspect or inspect_document() + return IdentityRequest( + inspect=document, + service_did=service_did or document.service.did, + service_url=service_url, + ) + + +def service_identity( + *, + agent_did: str = AGENT_DID, + identity_method: str = "did:web", + metadata: Mapping[str, str] | None = None, + service_did: str = SERVICE_DID, + signing_algorithms: tuple[SigningAlgorithm, ...] = (SigningAlgorithm.ES256,), +) -> ServiceIdentity: + return ServiceIdentity( + agent_did=agent_did, + identity_method=identity_method, + service_did=service_did, + signing_algorithms=signing_algorithms, + metadata={ + "agent_identity_id": "pai_one", + "created_at": "2026-09-03T12:00:00Z", + "did_document_url": "https://platform.example/agents/one/did.json", + "key_id": AGENT_DID, + "platform_url": "https://platform.example/", + "status": "active", + "updated_at": "2026-09-03T12:00:00Z", + } + if metadata is None + else metadata, + ) + + +def claims(**changes: object) -> ClientAssertionClaims: + values: dict[str, object] = { + "aud": SERVICE_DID, + "exp": int((NOW + timedelta(minutes=2)).timestamp()), + "iat": int(NOW.timestamp()), + "iss": AGENT_DID, + "jti": "assertion-one", + "op": AssertionOperation.ENROLL, + "sub": AGENT_DID, + } + values.update(changes) + return ClientAssertionClaims.model_validate(values) + + +def completed(**changes: object) -> dict[str, object]: + value: dict[str, object] = { + "agent_did": AGENT_DID, + "client_assertion": "header.payload.signature", + "expires_at": "2026-09-03T12:02:00Z", + "issued_at": "2026-09-03T12:00:00Z", + "jti": "assertion-one", + "service_did": SERVICE_DID, + "status": "completed", + } + value.update(changes) + return value + + +def provider( + transport: QueueTransport, + *, + allow_insecure_loopback: bool = False, + authentication_headers: PlatformAuthenticationHeaders | None = None, + authorization: str | None = None, + clock: Callable[[], datetime] = lambda: NOW, + discovery_cache: PlatformDiscoveryCache | None = None, + idempotency_key: PlatformIdempotencyKeyFactory | None = None, + maximum_response_bytes: int = 1 << 20, + pending_sign_resolver: PlatformPendingSignResolver | None = None, + platform_context: PlatformContextProvider | None = None, + platform_url: str = "https://platform.example", + request_timeout: float = 30.0, +) -> PlatformIdentityProvider: + return PlatformIdentityProvider( + PlatformIdentityProviderOptions( + allow_insecure_loopback=allow_insecure_loopback, + authentication_headers=authentication_headers, + authorization=authorization, + clock=clock, + discovery_cache=discovery_cache, + idempotency_key=idempotency_key, + maximum_response_bytes=maximum_response_bytes, + pending_sign_resolver=pending_sign_resolver, + platform_context=platform_context, + platform_url=platform_url, + request_timeout=request_timeout, + transport=transport, + ) + ) + + +@pytest.mark.asyncio +async def test_recovers_existing_identity_with_authenticated_list() -> None: + async def headers() -> dict[str, str]: + return { + "authorization": "Bearer dynamic", + "Content-Type": "ignored", + "Idempotency-Key": "ignored", + "X-Platform-Tenant": "tenant", + } + + transport = QueueTransport( + json_response(discovery(), headers={"Cache-Control": "max-age=300"}), + json_response(listed(identity())), + ) + instance = provider( + transport, + authentication_headers=headers, + authorization="Bearer static", + ) + recovered = await instance.get_or_create_identity(request()) + assert recovered == service_identity() + assert len(transport.requests) == 2 + listed_request = transport.requests[1] + assert listed_request.method == "GET" + assert parse_qs(urlsplit(listed_request.url).query) == { + "descending": ["true"], + "limit": ["100"], + "service_did": [SERVICE_DID], + } + assert listed_request.headers["authorization"] == "Bearer dynamic" + assert listed_request.headers["X-Platform-Tenant"] == "tenant" + assert listed_request.headers["Accept"] == AEP_MEDIA_TYPE + assert "Content-Type" not in listed_request.headers + assert "Idempotency-Key" not in listed_request.headers + await instance.aclose() + assert not transport.closed + + +@pytest.mark.asyncio +async def test_provisions_when_recovery_is_empty_and_serializes_concurrent_calls() -> None: + keys = iter(("provision-one", "provision-two")) + + async def key() -> str: + return next(keys) + + transport = QueueTransport( + json_response(discovery(), headers={"Cache-Control": "max-age=300"}), + json_response(listed()), + json_response(identity()), + json_response(listed(identity())), + ) + instance = provider(transport, idempotency_key=key) + first, second = await asyncio.gather( + instance.get_or_create_identity(request()), + instance.get_or_create_identity(request()), + ) + assert first == second == service_identity() + provision = transport.requests[2] + assert provision.method == "POST" + assert provision.headers["Idempotency-Key"] == "provision-one" + assert json.loads(provision.body or b"") == {"service_did": SERVICE_DID} + assert [item.method for item in transport.requests] == ["GET", "GET", "POST", "GET"] + + +@pytest.mark.asyncio +async def test_identity_request_and_platform_response_boundaries() -> None: + instance = provider(QueueTransport()) + with pytest.raises(ValueError, match="does not match"): + await instance.get_or_create_identity(request(service_did="did:web:other.example")) + + document_data = inspect_document().to_wire() + cast(dict[str, object], document_data["identity"])["methods"] = [] + cast(dict[str, object], document_data["commands"])["supported"] = ["inspect"] + document = InspectDocument.model_validate_json(json.dumps(document_data)) + with pytest.raises(ValueError, match="does not support"): + await instance.get_or_create_identity(request(inspect=document)) + + for changed, message in ( + ({"agent_did": "did:key:one", "key_id": "did:key:one"}, "invalid identity"), + ({"key_id": "did:web:other.example"}, "invalid identity"), + ({"did_document_url": "https://platform.example/wrong"}, "does not match"), + ): + broken = provider( + QueueTransport(json_response(discovery()), json_response(listed(identity(**changed)))) + ) + with pytest.raises(ValueError, match=message): + await broken.find_identity_by_service_did(SERVICE_DID) + + inactive = provider( + QueueTransport( + json_response(discovery()), + json_response(listed(identity(status="suspended"))), + ) + ) + assert await inactive.find_identity_by_service_did(SERVICE_DID) is None + + for changed in ( + {"service_did": "did:web:other.example"}, + {"status": "suspended"}, + ): + broken = provider( + QueueTransport( + json_response(discovery()), + json_response(listed()), + json_response(identity(**changed)), + ) + ) + with pytest.raises(ValueError, match="outside the requested Service scope"): + await broken.get_or_create_identity(request()) + + +@pytest.mark.asyncio +async def test_delegates_completed_signing_with_context() -> None: + context: dict[str, object] = {"authorization_handle": {"value": "opaque"}} + + async def context_provider( + selected: ServiceIdentity, assertion: ClientAssertionClaims + ) -> dict[str, object]: + assert selected.agent_did == assertion.iss + return context + + transport = QueueTransport(json_response(discovery()), json_response(completed())) + instance = provider(transport, platform_context=context_provider) + signer = await instance.signer_for(service_identity()) + assert await signer(claims(), (SigningAlgorithm.ES256,)) == "header.payload.signature" + outbound = transport.requests[1] + assert outbound.url.endswith("/v1/aep/agent-identities/pai_one/sign") + assert outbound.headers["Idempotency-Key"] + body = json.loads(outbound.body or b"") + assert body == { + "jti": "assertion-one", + "lifetime_seconds": "120", + "op": "enroll", + "platform_context": context, + "service_did": SERVICE_DID, + } + context["authorization_handle"] = "changed" + assert body["platform_context"] != context + + +@pytest.mark.asyncio +async def test_pending_signing_exposes_or_resolves_opaque_context() -> None: + pending_response = { + "platform_context": {"authorization_handle": {"value": "opaque"}}, + "retry_after_seconds": "5", + "status": "pending", + } + without_resolver = provider( + QueueTransport(json_response(discovery()), json_response(pending_response, status=202)) + ) + signer = await without_resolver.signer_for(service_identity()) + with pytest.raises(PlatformSignPendingError) as raised: + await signer(claims(), (SigningAlgorithm.ES256,)) + assert raised.value.pending.retry_after_seconds == 5 + assert raised.value.pending.platform_context == pending_response["platform_context"] + + resolved: list[PlatformPendingSign] = [] + + async def resolver(pending: PlatformPendingSign) -> dict[str, object]: + resolved.append(pending) + return {"authorization_handle": "approved"} + + keys = iter(("initial", "final")) + + async def key() -> str: + return next(keys) + + transport = QueueTransport( + json_response(discovery()), + json_response(pending_response, status=202), + json_response(completed(platform_context={"authorization_handle": "approved"})), + ) + instance = provider( + transport, + idempotency_key=key, + pending_sign_resolver=resolver, + ) + signer = await instance.signer_for(service_identity()) + assert await signer(claims(), (SigningAlgorithm.ES256,)) == "header.payload.signature" + assert len(resolved) == 1 + assert transport.requests[1].headers["Idempotency-Key"] == "initial" + assert transport.requests[2].headers["Idempotency-Key"] == "final" + assert json.loads(transport.requests[2].body or b"")["platform_context"] == { + "authorization_handle": "approved" + } + + +@pytest.mark.asyncio +async def test_signing_rejects_identity_claim_algorithm_and_stage_mismatches() -> None: + instance = provider(QueueTransport()) + invalid_identities = ( + service_identity(metadata={}), + service_identity(identity_method="did:key"), + service_identity(agent_did="did:key:one"), + service_identity(metadata={**service_identity().metadata, "key_id": "wrong"}), + service_identity(signing_algorithms=()), + service_identity( + metadata={**service_identity().metadata, "did_document_url": "https://wrong.example"} + ), + ) + for value in invalid_identities: + with pytest.raises(ValueError): + await instance.signer_for(value) + + signer = await instance.signer_for(service_identity()) + with pytest.raises(ValueError, match="another identity"): + await signer(claims(aud="did:web:other.example"), (SigningAlgorithm.ES256,)) + with pytest.raises(ValueError, match="no compatible"): + await signer(claims(), (SigningAlgorithm.EDDSA,)) + + async def duplicate() -> str: + return "same" + + duplicate_provider = provider( + QueueTransport( + json_response(discovery()), + json_response({"retry_after_seconds": "1", "status": "pending"}, status=202), + ), + idempotency_key=duplicate, + pending_sign_resolver=_empty_context, + ) + duplicate_signer = await duplicate_provider.signer_for(service_identity()) + with pytest.raises(ValueError, match="distinct idempotency"): + await duplicate_signer(claims(), (SigningAlgorithm.ES256,)) + + +async def _empty_context(pending: PlatformPendingSign) -> None: + del pending + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "sign_response,status,headers,message", + [ + (completed(), 202, {}, "completed Sign"), + (completed(jti="wrong"), 200, {}, "completed Sign"), + ({"retry_after_seconds": "1", "status": "pending"}, 200, {}, "pending Sign"), + (completed(), 200, {"Retry-After": "1"}, "included Retry-After"), + ], +) +async def test_sign_response_contract_boundaries( + sign_response: dict[str, object], + status: int, + headers: dict[str, str], + message: str, +) -> None: + instance = provider( + QueueTransport( + json_response(discovery()), + json_response(sign_response, status=status, headers=headers), + ) + ) + signer = await instance.signer_for(service_identity()) + with pytest.raises(ValueError, match=message): + await signer(claims(), (SigningAlgorithm.ES256,)) + + +@pytest.mark.asyncio +async def test_authenticate_signing_includes_resource() -> None: + transport = QueueTransport(json_response(discovery()), json_response(completed())) + instance = provider(transport) + signer = await instance.signer_for(service_identity()) + authentication = claims( + op=AssertionOperation.AUTHENTICATE, + resource="https://api.example.com/orders", + ) + await signer(authentication, (SigningAlgorithm.ES256,)) + assert json.loads(transport.requests[1].body or b"")["resource"] == ( + "https://api.example.com/orders" + ) + + +@pytest.mark.asyncio +async def test_discovery_cache_freshness_revalidation_and_no_store() -> None: + cache = MemoryPlatformDiscoveryCache() + transport = QueueTransport( + json_response(discovery(), headers={"Cache-Control": "max-age=0", "ETag": '"one"'}), + json_response(listed()), + HttpResponse(status=304, headers={"Cache-Control": "max-age=300"}), + json_response(listed()), + ) + instance = provider(transport, discovery_cache=cache) + assert await instance.find_identity_by_service_did(SERVICE_DID) is None + assert await instance.find_identity_by_service_did(SERVICE_DID) is None + assert transport.requests[2].headers["If-None-Match"] == '"one"' + + modified_transport = QueueTransport( + json_response( + discovery(), + headers={"Cache-Control": "max-age=0", "Last-Modified": "latest"}, + ), + json_response(listed()), + HttpResponse(status=304), + json_response(listed()), + ) + modified = provider(modified_transport) + await modified.find_identity_by_service_did(SERVICE_DID) + await modified.find_identity_by_service_did(SERVICE_DID) + assert modified_transport.requests[2].headers["If-Modified-Since"] == "latest" + assert "If-None-Match" not in modified_transport.requests[2].headers + + no_store_transport = QueueTransport( + json_response(discovery(), headers={"Cache-Control": "no-store"}), + json_response(listed()), + json_response(discovery(), headers={"Cache-Control": "no-store"}), + json_response(listed()), + ) + no_store = provider(no_store_transport) + await no_store.find_identity_by_service_did(SERVICE_DID) + await no_store.find_identity_by_service_did(SERVICE_DID) + assert ( + sum(item.url.endswith("/.well-known/aep-platform") for item in no_store_transport.requests) + == 2 + ) + + +@pytest.mark.asyncio +async def test_discovery_redirect_and_failure_boundaries() -> None: + redirected = provider( + QueueTransport( + HttpResponse(status=307, headers={"Location": "/metadata/platform"}), + json_response(discovery()), + json_response(listed()), + ) + ) + await redirected.find_identity_by_service_did(SERVICE_DID) + + failures = ( + (QueueTransport(HttpResponse(status=307)), "omitted Location"), + ( + QueueTransport( + HttpResponse(status=307, headers={"Location": "https://other.example/platform"}) + ), + "changed origin", + ), + (QueueTransport(json_response({}, status=500)), "HTTP 500"), + ( + QueueTransport(json_response(discovery(), headers={"Content-Type": "text/plain"})), + "media type", + ), + (QueueTransport(HttpResponse(status=304)), "without a cached"), + ( + QueueTransport( + json_response( + discovery( + identity={ + "did_methods": ["did:key"], + "did_url_template": ( + "https://platform.example/agents/{agent_did_id}/did.json" + ), + } + ) + ) + ), + "does not advertise did:web", + ), + ) + for transport, message in failures: + with pytest.raises((PlatformCommandError, ValueError), match=message): + await provider(transport).find_identity_by_service_did(SERVICE_DID) + + redirects = [ + HttpResponse(status=307, headers={"Location": f"/redirect-{value}"}) for value in range(6) + ] + with pytest.raises(ValueError, match="five redirects"): + await provider(QueueTransport(*redirects)).find_identity_by_service_did(SERVICE_DID) + + +@pytest.mark.asyncio +async def test_invalid_cached_discovery_is_evicted() -> None: + cache = MemoryPlatformDiscoveryCache() + document = PlatformDiscoveryDocument.model_validate_json(json.dumps(discovery())) + unsupported = PlatformDiscoveryDocument.model_validate_json( + json.dumps( + discovery( + identity={ + "did_methods": ["did:key"], + "did_url_template": ("https://platform.example/agents/{agent_did_id}/did.json"), + } + ) + ) + ) + key = "https://platform.example/.well-known/aep-platform" + for entry in ( + PlatformDiscoveryCacheEntry(datetime(2026, 9, 3), document, key), + PlatformDiscoveryCacheEntry(NOW, document, "https://other.example/platform"), + PlatformDiscoveryCacheEntry(NOW, unsupported, key), + ): + await cache.save(key, entry) + transport = QueueTransport(json_response(discovery()), json_response(listed())) + await provider(transport, discovery_cache=cache).find_identity_by_service_did(SERVICE_DID) + assert transport.requests[0].url == key + await cache.delete(key) + assert await cache.find(key) is None + + +@pytest.mark.asyncio +async def test_platform_command_errors_media_bounds_and_timeout() -> None: + problem = { + "code": "not_recognized", + "status": 401, + "title": "Not recognized", + "type": "urn:aep:error:not_recognized", + } + for body, content_type, expected_problem in ( + (problem, AEP_PROBLEM_MEDIA_TYPE, True), + ({**problem, "status": 403}, AEP_PROBLEM_MEDIA_TYPE, False), + ({"invalid": True}, AEP_PROBLEM_MEDIA_TYPE, False), + (problem, "application/json", False), + ): + instance = provider( + QueueTransport( + json_response(discovery()), + json_response(body, status=401, headers={"Content-Type": content_type}), + ) + ) + with pytest.raises(PlatformCommandError) as raised: + await instance.find_identity_by_service_did(SERVICE_DID) + assert (raised.value.problem is not None) is expected_problem + assert raised.value.status == 401 + + invalid_media = provider( + QueueTransport( + json_response(discovery()), + json_response(listed(), headers={"Content-Type": "application/json"}), + ) + ) + with pytest.raises(ValueError, match="response media type"): + await invalid_media.find_identity_by_service_did(SERVICE_DID) + + oversized = provider(QueueTransport(json_response(discovery())), maximum_response_bytes=10) + with pytest.raises(ValueError, match="configured limit"): + await oversized.find_identity_by_service_did(SERVICE_DID) + + timed_out = PlatformIdentityProvider( + PlatformIdentityProviderOptions( + platform_url="https://platform.example", + request_timeout=0.001, + transport=BlockingTransport(), + ) + ) + with pytest.raises(TimeoutError): + await timed_out.find_identity_by_service_did(SERVICE_DID) + + +def test_configuration_and_url_boundaries() -> None: + for changes, message in ( + ({"maximum_response_bytes": 0}, "response bytes"), + ({"request_timeout": 0}, "timeout"), + ({"request_timeout": float("inf")}, "timeout"), + ({"clock": lambda: datetime(2026, 9, 3)}, "offset-aware"), + ): + with pytest.raises(ValueError, match=message): + provider(QueueTransport(), **changes) + for value in ("", "http://platform.example", "https://user:secret@platform.example"): + with pytest.raises(ValueError): + provider(QueueTransport(), platform_url=value) + loopback = provider( + QueueTransport(), + allow_insecure_loopback=True, + platform_url="http://localhost:8080/path", + ) + assert loopback._platform_url == "http://localhost:8080/" + assert _platform_url("platform.example/path", False) == "https://platform.example/" + assert _valid_url("https://platform.example", False) + assert not _valid_url("https:///missing", False) + + for path in ( + "relative", + "//other.example/path", + "/path?query=true", + "/path#fragment", + "/path/{missing}", + ): + with pytest.raises(ValueError, match="invalid endpoint"): + _endpoint("https://platform.example/", path) + assert _endpoint( + "https://platform.example/", + "/identities/{agent_identity_id}", + agent_identity_id="a/b", + ).endswith("/identities/a%2Fb") + + +def test_cache_and_pending_value_boundaries() -> None: + document = PlatformDiscoveryDocument.model_validate_json(json.dumps(discovery())) + entry = PlatformDiscoveryCacheEntry(NOW, document, "https://platform.example/platform") + assert _cache_fresh(entry, NOW + timedelta(seconds=299)) + assert not _cache_fresh(entry, NOW + timedelta(seconds=300)) + for control in ("no-cache", "no-store", "max-age=invalid", "max-age=-1"): + assert not _cache_fresh(replace(entry, cache_control=control), NOW) + assert _cache_directive('public, max-age="60"', "max-age") == '"60"' + assert _cache_directive("public", "missing") is None + with pytest.raises(ValueError, match="timestamp"): + _validate_cache_entry( + replace(entry, cached_at=datetime(2026, 9, 3)), + "https://platform.example/.well-known/aep-platform", + False, + ) + with pytest.raises(ValueError, match="URL"): + _validate_cache_entry( + replace(entry, final_url="https://other.example/platform"), + "https://platform.example/.well-known/aep-platform", + False, + ) + for seconds in (0, 301): + with pytest.raises(ValueError, match="between 1 and 300"): + PlatformPendingSign(service_identity(), {}, seconds) + + +@pytest.mark.asyncio +async def test_default_owned_transport_closes_and_empty_key_fails() -> None: + instance = PlatformIdentityProvider( + PlatformIdentityProviderOptions(platform_url="https://platform.example") + ) + async with instance as entered: + assert entered is instance + + async def empty_key() -> str: + return " " + + invalid = provider( + QueueTransport(json_response(discovery()), json_response(listed())), + idempotency_key=empty_key, + ) + with pytest.raises(ValueError, match="key generation"): + await invalid.get_or_create_identity(request()) diff --git a/tests/test_package.py b/tests/test_package.py index e2249bc..6cb647d 100644 --- a/tests/test_package.py +++ b/tests/test_package.py @@ -2,8 +2,18 @@ from agent_enrollment_protocol import __version__, adapters, agent, core, platform, service from agent_enrollment_protocol.adapters import AepAsgiApplication, AepAuthenticationMiddleware -from agent_enrollment_protocol.agent import Agent, AgentOptions, HttpxTransport, ServiceIdentity -from agent_enrollment_protocol.core import ClaimSupportEvaluation, evaluate_claim_support +from agent_enrollment_protocol.agent import ( + Agent, + AgentOptions, + HttpxTransport, + PlatformIdentityProvider, + ServiceIdentity, +) +from agent_enrollment_protocol.core import ( + AEP_PLATFORM_WELL_KNOWN_PATH, + ClaimSupportEvaluation, + evaluate_claim_support, +) from agent_enrollment_protocol.service import ( MemoryServiceCredentialStore, Service, @@ -24,6 +34,10 @@ def test_public_package_modules() -> None: assert Agent.__module__ == "agent_enrollment_protocol.agent.client" assert AgentOptions.__module__ == "agent_enrollment_protocol.agent.client" assert HttpxTransport.__module__ == "agent_enrollment_protocol.agent.transport" + assert PlatformIdentityProvider.__module__ == ( + "agent_enrollment_protocol.agent.platform_provider" + ) + assert AEP_PLATFORM_WELL_KNOWN_PATH == "/.well-known/aep-platform" assert ServiceIdentity.__module__ == "agent_enrollment_protocol.agent.types" assert Service.__module__ == "agent_enrollment_protocol.service.service" assert ServiceOptions.__module__ == "agent_enrollment_protocol.service.types"