From 43f07d0bd86f16539909c2cc0eb318fa38837c4f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 9 Oct 2026 14:14:17 -0700 Subject: [PATCH] fix(bedrock): list the account's invocable models behind bedrock/* when check_provider_endpoint is on (#43981) * fix(bedrock): list the account's invocable models behind bedrock/* when check_provider_endpoint is on * fix(bedrock): sign the listing query RFC 3986 style and raise BedrockError on a failed listing * test(utils): keep test_utils.py at main's formatting, only the Bedrock discovery test is new * fix(bedrock): chain BedrockModelInfo's constructor so the passthrough config keeps its event parser * fix(bedrock): list discovered ids without the provider prefix so partial wildcards filter get_valid_models returns bare vendor ids for every other provider and for the static bedrock catalog, and the proxy's wildcard expansion adds the provider prefix itself. The lister prefixed its ids, so a partial wildcard like bedrock/anthropic.* never matched the proxy's filter and listed uncallable bedrock/anthropic.bedrock/ entries * test(bedrock): integration cells for wildcard discovery through the proxy and the SDK * fix(proxy): keep a partial wildcard a filter when its deployment repeats the prefix The wildcard expansion guessed filter-or-alias by whether any provider id carried the prefix. A deployment such as model_name bedrock/anthropic.* over model bedrock/anthropic.* can only ever route names that share the prefix, so when the account has no on-demand anthropic.* id (every Anthropic model behind an inference profile) the guess fell into the alias branch and listed bedrock/anthropic., none of them callable. A deployment whose model repeats the suffix now always filters, and lists nothing when nothing matches * fix(proxy): prefix every discovered id under a custom wildcard prefix that starts a vendor id * test(bedrock): Final and read-only annotations in the wildcard discovery tests --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/llms/bedrock/common_utils.py | 15 +- litellm/llms/bedrock/model_listing.py | 119 +++ litellm/proxy/auth/model_checks.py | 9 +- .../integration/_support/aws_control_plane.py | 281 ++++++ .../integration/_support/bedrock_discovery.py | 294 +++++++ tests/integration/_support/client.py | 27 + tests/integration/_support/sigv4.py | 35 +- .../test_bedrock_model_discovery_chaos.py | 107 +++ .../test_bedrock_model_discovery_wire.py | 815 ++++++++++++++++++ .../sdk/test_bedrock_model_discovery_sdk.py | 130 +++ .../bedrock/test_bedrock_model_listing.py | 166 ++++ tests/unit/proxy/auth/test_model_checks.py | 161 +++- tests/unit/test_utils.py | 48 ++ 13 files changed, 2178 insertions(+), 29 deletions(-) create mode 100644 litellm/llms/bedrock/model_listing.py create mode 100644 tests/integration/_support/aws_control_plane.py create mode 100644 tests/integration/_support/bedrock_discovery.py create mode 100644 tests/integration/providers/test_bedrock_model_discovery_chaos.py create mode 100644 tests/integration/providers/test_bedrock_model_discovery_wire.py create mode 100644 tests/integration/sdk/test_bedrock_model_discovery_sdk.py create mode 100644 tests/unit/llms/bedrock/test_bedrock_model_listing.py diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index e046b0f3616..6b28ecd02df 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -38,6 +38,7 @@ from litellm.secret_managers.main import get_secret, get_secret_str from litellm.types.llms.bedrock import AWS_AUTH_PARAM_KEYS, AwsAuthParams if TYPE_CHECKING: + from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.types.llms.openai import AllMessageValues @@ -1363,6 +1364,10 @@ class BedrockModelInfo(BaseLLMModelInfo): global_config = AmazonBedrockGlobalConfig() all_global_regions = global_config.get_all_regions() + def __init__(self, client: HTTPHandler | None = None) -> None: + super().__init__() + self._client: Final = client + @staticmethod def get_api_base(api_base: str | None = None) -> str | None: """ @@ -1390,7 +1395,15 @@ class BedrockModelInfo(BaseLLMModelInfo): return headers def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: - return [] + return self.discover_models({"api_key": api_key}) + + def discover_models( + self, litellm_params: Mapping[str, object] | None = None + ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override + from litellm.llms.bedrock.model_listing import BedrockModelLister + + client: Final = self._client if self._client is not None else litellm.module_level_client + return sorted(BedrockModelLister(deployment=litellm_params or {}, client=client).invocable_model_ids()) # def get_provider_info(self, model: str) -> Optional[ProviderSpecificModelInfo]: # """ diff --git a/litellm/llms/bedrock/model_listing.py b/litellm/llms/bedrock/model_listing.py new file mode 100644 index 00000000000..99ff48f9400 --- /dev/null +++ b/litellm/llms/bedrock/model_listing.py @@ -0,0 +1,119 @@ +from collections.abc import Iterator, Mapping +from types import MappingProxyType +from typing import Final, TypeAlias +from urllib.parse import quote, urlencode + +import httpx +from pydantic import BaseModel, ConfigDict, Field + +from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, bedrock_bearer_token +from litellm.llms.bedrock.common_utils import BedrockError +from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.types.llms.bedrock import AwsAuthParams + +LIST_MODELS_TIMEOUT: Final = 10.0 +INFERENCE_PROFILES_PAGE_SIZE: Final = 1000 +INFERENCE_PROFILES_PAGE_CAP: Final = 20 + +QueryPairs: TypeAlias = tuple[tuple[str, str], ...] + + +class FoundationModelSummary(BaseModel): + model_config = ConfigDict(extra="ignore") + + id: str = Field(alias="modelId") + + +class ListFoundationModelsResponse(BaseModel): + model_config = ConfigDict(extra="ignore") + + summaries: tuple[FoundationModelSummary, ...] = Field(default=(), alias="modelSummaries") + + +class InferenceProfileSummary(BaseModel): + model_config = ConfigDict(extra="ignore") + + id: str = Field(alias="inferenceProfileId") + status: str + + +class ListInferenceProfilesResponse(BaseModel): + model_config = ConfigDict(extra="ignore") + + summaries: tuple[InferenceProfileSummary, ...] = Field(default=(), alias="inferenceProfileSummaries") + next_token: str | None = Field(default=None, alias="nextToken") + + +class BedrockModelLister(BaseAWSLLM): + def __init__(self, deployment: Mapping[str, object], client: HTTPHandler) -> None: + super().__init__() + self._aws_params: Final = dict(deployment) + self._aws_region_name: Final = self._get_aws_region_name(optional_params=self._aws_params) + api_key: Final = deployment.get("api_key") + self._bearer_token: Final = bedrock_bearer_token(api_key if isinstance(api_key, str) else None) + self._client: Final = client + + def invocable_model_ids(self) -> frozenset[str]: + return frozenset(self._active_inference_profile_ids()) | frozenset(self._on_demand_foundation_model_ids()) + + def _on_demand_foundation_model_ids(self) -> Iterator[str]: + response: Final = ListFoundationModelsResponse.model_validate( + self._get_json("/foundation-models", (("byInferenceType", "ON_DEMAND"),)) + ) + return (summary.id for summary in response.summaries) + + def _active_inference_profile_ids(self) -> tuple[str, ...]: + collected: tuple[str, ...] = () # rebind-ok: accumulates one page of ids per iteration + next_token: str | None = None # rebind-ok: advances to each page's nextToken + for _ in range(INFERENCE_PROFILES_PAGE_CAP): + page = self._inference_profiles_page(next_token) + collected += tuple(summary.id for summary in page.summaries if summary.status == "ACTIVE") + if page.next_token is None: + return collected + next_token = page.next_token + raise BedrockError( + status_code=500, + message=( + f"Bedrock inference profile listing in {self._aws_region_name} did not end within " + f"{INFERENCE_PROFILES_PAGE_CAP} pages." + ), + ) + + def _inference_profiles_page(self, next_token: str | None) -> ListInferenceProfilesResponse: + continuation: Final[QueryPairs] = (("nextToken", next_token),) if next_token is not None else () + return ListInferenceProfilesResponse.model_validate( + self._get_json( + "/inference-profiles", + (("maxResults", str(INFERENCE_PROFILES_PAGE_SIZE)), ("typeEquals", "SYSTEM_DEFINED"), *continuation), + ) + ) + + def _get_json(self, path: str, query: QueryPairs) -> object: + host: Final = f"bedrock.{self._aws_region_name}.{get_aws_dns_suffix(self._aws_region_name)}" + url: Final = f"https://{host}{path}?{urlencode(query, quote_via=quote)}" + response: Final = self._client.get(url=url, headers=self._auth_headers(url), timeout=LIST_MODELS_TIMEOUT) + try: + response.raise_for_status() + except httpx.HTTPStatusError: + raise BedrockError( + status_code=response.status_code, + message=( + f"Failed to list Bedrock models in {self._aws_region_name}. " + f"Status code: {response.status_code}, Response: {response.text}" + ), + ) + return response.json() + + def _auth_headers(self, url: str) -> Mapping[str, str]: + if self._bearer_token is not None: + return MappingProxyType({"Authorization": f"Bearer {self._bearer_token}"}) + from botocore.auth import SigV4Auth + from botocore.awsrequest import AWSRequest + + credentials: Final = self.resolve_credentials( + AwsAuthParams.model_validate(self._aws_params), self._aws_region_name + ) + request: Final = AWSRequest(method="GET", url=url, data="") + SigV4Auth(credentials, "bedrock", self._aws_region_name).add_auth(request) + return MappingProxyType({name: str(value) for name, value in request.prepare().headers.items()}) diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index 133e95eadba..3101f19ced2 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -303,7 +303,12 @@ def get_known_models_from_wildcard(wildcard_model: str, litellm_params: LiteLLM_ ## CHECK IF PARTIAL FILTER e.g. `gemini-*` model_prefix: Final = wildcard_suffix.replace("*", "") - is_partial_filter: Final = any(wc_model.startswith(model_prefix) for wc_model in wildcard_models) + deployment_repeats_prefix: Final = litellm_params is not None and litellm_params.model.endswith( + f"/{wildcard_suffix}" + ) + is_partial_filter: Final = deployment_repeats_prefix or any( + wc_model.startswith(model_prefix) for wc_model in wildcard_models + ) if is_partial_filter: filtered_wildcard_models = [wc_model for wc_model in wildcard_models if wc_model.startswith(model_prefix)] wildcard_models = filtered_wildcard_models @@ -314,7 +319,7 @@ def get_known_models_from_wildcard(wildcard_model: str, litellm_params: LiteLLM_ known_providers: Final = {provider.value for provider in LlmProviders} suffix_appended_wildcard_models: Final = [] for model in wildcard_models: - if not model.startswith(wildcard_provider_prefix): + if not model.startswith(f"{wildcard_provider_prefix}/"): # `get_provider_models` returns provider-prefixed ids (e.g. "ollama/gemma3:1b"). # When the wildcard uses a custom prefix (e.g. "ollama_server1/*" to distinguish # multiple instances), replace that existing provider prefix instead of stacking diff --git a/tests/integration/_support/aws_control_plane.py b/tests/integration/_support/aws_control_plane.py new file mode 100644 index 00000000000..8cef78f8bee --- /dev/null +++ b/tests/integration/_support/aws_control_plane.py @@ -0,0 +1,281 @@ +from __future__ import annotations + +import json +import re +import threading +from collections.abc import Callable, Generator, Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, TypeAlias +from urllib.parse import parse_qsl, urlsplit + +from integration._support import tls +from integration._support.sigv4 import canonical_query, signature +from integration._support.wire import Reply +from pydantic import JsonValue + +_AUTHORIZATION: Final = re.compile( + r"^AWS4-HMAC-SHA256 Credential=(?P[^/]+)/(?P[^,]+), " + r"SignedHeaders=(?P[^,]+), Signature=(?P[0-9a-f]{64})$" +) +_BEARER: Final = re.compile(r"^Bearer (?P.+)$") +UNKNOWN_CREDENTIAL: Final = "none" +NO_PROXY_HOSTS: Final = "127.0.0.1,localhost" +INFERENCE_PROFILES: Final = "/inference-profiles" +FOUNDATION_MODELS: Final = "/foundation-models" + + +@dataclass(frozen=True, slots=True) +class ControlPlaneRequest: + host: str + method: str + target: str + headers: Mapping[str, str] + body: bytes + + @property + def path(self) -> str: + return urlsplit(self.target).path + + @property + def query(self) -> Mapping[str, str]: + return MappingProxyType(dict(parse_qsl(urlsplit(self.target).query, keep_blank_values=True))) + + @property + def credential(self) -> str: + """The access key id of a SigV4 request, the token of a bearer one, UNKNOWN_CREDENTIAL otherwise.""" + authorization: Final = self.headers.get("authorization", "") + signed: Final = _AUTHORIZATION.match(authorization) + if signed is not None: + return signed.group("key") + bearer: Final = _BEARER.match(authorization) + return bearer.group("token") if bearer is not None else UNKNOWN_CREDENTIAL + + +def verified_scope(request: ControlPlaneRequest, secret: str) -> str: + """Recompute the SigV4 signature with `secret`; returns the credential scope the request was signed under.""" + found: Final = _AUTHORIZATION.match(request.headers.get("authorization", "")) + assert found is not None, dict(request.headers) + parts: Final = urlsplit(request.target) + _, expected = signature( + request.method, + parts.path, + request.headers, + found.group("signed"), + request.body, + secret, + found.group("scope"), + canonical_query(parts.query), + ) + assert found.group("signature") == expected, request + return found.group("scope") + + +Responder: TypeAlias = Callable[[ControlPlaneRequest], Reply] + + +def json_reply(status: int, payload: Mapping[str, JsonValue]) -> Reply: + return Reply(status=status, body=json.dumps(payload).encode()) + + +def _unregistered(request: ControlPlaneRequest) -> Reply: + return json_reply(403, {"message": f"no scripted catalog for credential {request.credential}"}) + + +@dataclass(frozen=True, slots=True) +class Catalog: + """What the scripted control plane advertises for one credential; `page_size` splits the profiles.""" + + active_profiles: tuple[str, ...] = () + inactive_profiles: tuple[str, ...] = () + on_demand_models: tuple[str, ...] = () + page_size: int | None = None + + def vendor_ids(self) -> frozenset[str]: + return frozenset((*self.active_profiles, *self.on_demand_models)) + + def invocable_ids(self) -> frozenset[str]: + return frozenset(f"bedrock/{model_id}" for model_id in self.vendor_ids()) + + def pages(self) -> tuple[tuple[Mapping[str, JsonValue], ...], ...]: + summaries: Final[tuple[Mapping[str, JsonValue], ...]] = ( + *({"inferenceProfileId": name, "status": "ACTIVE"} for name in self.active_profiles), + *({"inferenceProfileId": name, "status": "INACTIVE"} for name in self.inactive_profiles), + ) + size: Final = len(summaries) if self.page_size is None else self.page_size + return tuple(summaries[start : start + size] for start in range(0, len(summaries), size)) or ((),) + + def respond(self, request: ControlPlaneRequest) -> Reply: + if request.method != "GET": + return json_reply(405, {"message": f"{request.method} is not a listing"}) + if request.path == FOUNDATION_MODELS: + if request.query != {"byInferenceType": "ON_DEMAND"}: + return json_reply(400, {"message": f"unexpected foundation-models query {dict(request.query)}"}) + return json_reply(200, {"modelSummaries": [{"modelId": name} for name in self.on_demand_models]}) + if request.path != INFERENCE_PROFILES: + return json_reply(404, {"message": f"no scripted listing at {request.path}"}) + pages: Final = self.pages() + index: Final = int(request.query.get("nextToken", "page-0").removeprefix("page-")) + expected_query: Final = {"maxResults": "1000", "typeEquals": "SYSTEM_DEFINED"} | ( + {"nextToken": f"page-{index}"} if index else {} + ) + if dict(request.query) != expected_query or index >= len(pages): + return json_reply(400, {"message": f"unexpected inference-profiles query {dict(request.query)}"}) + continuation: Final = {"nextToken": f"page-{index + 1}"} if index + 1 < len(pages) else {} + return json_reply(200, {"inferenceProfileSummaries": list(pages[index]), **continuation}) + + +@dataclass(frozen=True, slots=True) +class ControlPlane: + """An owned HTTP CONNECT proxy terminating TLS for the hosted Bedrock control-plane names. + + The proxy under test gets it as ``HTTPS_PROXY`` and trusts its certificate through ``SSL_VERIFY``, so the + lister's own ``https://bedrock..`` URLs reach a scripted catalog without a DNS override. + A CONNECT to any other host is refused with 403 and recorded in `refused`.""" + + url: str + certificate: Path + received: SimpleQueue[ControlPlaneRequest] + refused: SimpleQueue[str] + responders: dict[str, Responder] # mutable-ok: answering() adds a responder per test and removes it + + def environment(self) -> Mapping[str, str]: + return MappingProxyType( + { + "HTTPS_PROXY": self.url, + "https_proxy": self.url, + "NO_PROXY": NO_PROXY_HOSTS, + "no_proxy": NO_PROXY_HOSTS, + "SSL_VERIFY": str(self.certificate), + } + ) + + def drain(self) -> tuple[ControlPlaneRequest, ...]: + return tuple(self.received.get_nowait() for _ in range(self.received.qsize())) + + def refusals(self) -> tuple[str, ...]: + return tuple(self.refused.get_nowait() for _ in range(self.refused.qsize())) + + @contextmanager + def answering(self, credential: str, respond: Responder) -> Iterator[None]: + self.responders[credential] = respond + try: + yield + finally: + del self.responders[credential] + + +def _header_map(handler: BaseHTTPRequestHandler) -> Mapping[str, str]: + return MappingProxyType({name.lower(): value for name, value in handler.headers.items()}) + + +def _answer(request: ControlPlaneRequest, responders: Mapping[str, Responder], errors: SimpleQueue[Exception]) -> Reply: + try: + return responders.get(request.credential, _unregistered)(request) + except Exception as error: + errors.put(error) + return Reply(status=500) + + +def _tunneled_handler( + host: str, + received: SimpleQueue[ControlPlaneRequest], + responders: Mapping[str, Responder], + errors: SimpleQueue[Exception], +) -> type[BaseHTTPRequestHandler]: + class Tunneled(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + timeout = 60 + + def respond(self) -> None: + request: Final = ControlPlaneRequest( + host, + self.command, + self.path, + _header_map(self), + self.rfile.read(int(self.headers.get("content-length", "0"))), + ) + received.put(request) + reply: Final = _answer(request, responders, errors) + if reply.drop_connection: + self.close_connection = True + return + try: + self.send_response(reply.status) + self.send_header("content-type", reply.content_type) + for name, value in reply.headers.items(): + self.send_header(name, value) + self.send_header("content-length", str(len(reply.body))) + self.end_headers() + self.wfile.write(reply.body) + self.wfile.flush() + except (BrokenPipeError, ConnectionResetError, TimeoutError): + self.close_connection = True + + do_GET = respond + do_POST = respond + + def log_message(self, format: str, *args: object) -> None: + pass + + return Tunneled + + +@contextmanager +def control_plane(directory: Path, hosts: tuple[str, ...]) -> Generator[ControlPlane, None, None]: + cert_file, key_file = tls.write_self_signed_cert(directory, names=hosts) + context: Final = tls.server_context(cert_file, key_file) + trusted: Final = frozenset(hosts) + received: Final[SimpleQueue[ControlPlaneRequest]] = SimpleQueue() + refused: Final[SimpleQueue[str]] = SimpleQueue() + errors: Final[SimpleQueue[Exception]] = SimpleQueue() + responders: Final[dict[str, Responder]] = {} # mutable-ok: cells register and remove their catalogs + + class Proxy(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + timeout = 60 + + def do_CONNECT(self) -> None: + host: Final = self.path.rsplit(":", 1)[0] + self.close_connection = True + if host not in trusted: + refused.put(self.path) + self.send_response(403) + self.send_header("content-length", "0") + self.end_headers() + return + self.send_response(200, "Connection established") + self.end_headers() + self.wfile.flush() + try: + secured: Final = context.wrap_socket(self.connection, server_side=True) + except OSError as error: + errors.put(error) + return + with secured: + _tunneled_handler(host, received, responders, errors)(secured, self.client_address, self.server) + + def log_message(self, format: str, *args: object) -> None: + pass + + class Server(ThreadingHTTPServer): + daemon_threads = True + block_on_close = False + request_queue_size = 128 + + with Server(("127.0.0.1", 0), Proxy) as server: + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}) + thread.start() + try: + yield ControlPlane(f"http://127.0.0.1:{server.server_port}", cert_file, received, refused, responders) + finally: + server.shutdown() + thread.join(timeout=6) + assert not thread.is_alive(), "Owned control plane survived cleanup" + server.server_close() + failure: Final = None if errors.empty() else errors.get_nowait() + assert failure is None, f"Owned control plane failed: {failure!r}" diff --git a/tests/integration/_support/bedrock_discovery.py b/tests/integration/_support/bedrock_discovery.py new file mode 100644 index 00000000000..5b3a091dc17 --- /dev/null +++ b/tests/integration/_support/bedrock_discovery.py @@ -0,0 +1,294 @@ +"""Owned rig and assertions for Bedrock wildcard model discovery through a scripted AWS control plane.""" + +from __future__ import annotations + +import os +import re +import uuid +from collections import Counter +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from pathlib import Path +from typing import Final +from urllib.parse import urlsplit + +import psutil +import yaml +from integration._support.aws_control_plane import ( + FOUNDATION_MODELS, + INFERENCE_PROFILES, + Catalog, + ControlPlane, + ControlPlaneRequest, + verified_scope, +) +from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, held, object_value, string_value +from integration._support.process import OwnedProxy, owned_proxy_process +from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix +from pydantic import JsonValue + +HOSTED_REGIONS: Final = ("us-east-1", "us-west-2", "eu-west-1", "us-gov-west-1", "cn-north-1") +UNHOSTED_REGION: Final = "eu-central-1" +CONTROL_MODEL: Final = "integration-discovery-control" +CONVERSE_MODEL: Final = "us.openai.gpt-5.6-sol" +WORKERS: Final = 2 +RELOAD_SECONDS: Final = 5 +LISTING_TIMEOUT_SECONDS: Final = 10.0 +PROFILE_PAGE_CAP: Final = 20 +LISTING_PATHS: Final = ("/v1/models", "/models") +INFO_PATHS: Final = ("/model/info", "/v1/model/info") +PROFILES_QUERY: Final[Mapping[str, str]] = {"maxResults": "1000", "typeEquals": "SYSTEM_DEFINED"} +FOUNDATION_QUERY: Final[Mapping[str, str]] = {"byInferenceType": "ON_DEMAND"} +CLOSE: Final[Mapping[str, str]] = {"connection": "close"} +LISTING_FAILURE_LOG: Final = "Error getting valid models" +STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +STARTUP_COMPLETE: Final = "Application startup complete." + + +def control_plane_host(region: str) -> str: + return f"bedrock.{region}.{get_aws_dns_suffix(region)}" + + +def credential() -> tuple[str, str]: + return f"AKIA{uuid.uuid4().hex[:16].upper()}", f"secret-{uuid.uuid4().hex}" + + +def stem() -> str: + return f"disc{uuid.uuid4().hex[:10]}" + + +def catalog_for(marker: str, *, page_size: int | None = None) -> Catalog: + return Catalog( + active_profiles=(f"us.{marker}.sonnet-v1:0", f"eu.{marker}.haiku-v1:0"), + inactive_profiles=(f"us.{marker}.retired-v1:0",), + on_demand_models=(f"{marker}.nova-micro-v1:0",), + page_size=page_size, + ) + + +def discovery_config(parent: Gateway, directory: Path, *, check_provider_endpoint: bool) -> Path: + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + configuration["model_list"] = [ + { + "model_name": CONTROL_MODEL, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "integration-provider-key", + "api_base": f"{parent.upstream_url}/v1", + }, + } + ] + configuration["litellm_settings"] = { + **configuration.get("litellm_settings", {}), + "check_provider_endpoint": check_provider_endpoint, + } + path: Final = directory / f"discovery-{'on' if check_provider_endpoint else 'off'}.yaml" + path.write_text(yaml.safe_dump(configuration)) + return path + + +def discovery_environment(plane: ControlPlane, directory: Path) -> Mapping[str, str]: + empty: Final = directory / "empty-aws-config" + empty.write_text("") + return { + **plane.environment(), + "AWS_CONFIG_FILE": str(empty), + "AWS_SHARED_CREDENTIALS_FILE": str(empty), + "AWS_EC2_METADATA_DISABLED": "true", + "LITELLM_LOCAL_MODEL_COST_MAP": "True", + "PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": str(RELOAD_SECONDS), + } + + +def aws_environment_names() -> tuple[str, ...]: + return tuple(name for name in os.environ if name.startswith("AWS_")) + + +@contextmanager +def discovery_proxy( + parent: Gateway, directory: Path, plane: ControlPlane, *, check_provider_endpoint: bool, workers: int +) -> Iterator[OwnedProxy]: + with owned_proxy_process( + parent, + directory, + discovery_environment(plane, directory), + config=discovery_config(parent, directory, check_provider_endpoint=check_provider_endpoint), + remove_environment=aws_environment_names(), + workers=workers, + ) as owned: + wait_for_every_worker(owned.log, workers) + yield owned + + +def deployment_ids(gateway: Gateway) -> frozenset[str]: + entries: Final = gateway.get("/model/info")["data"] + assert isinstance(entries, list) + return frozenset(string_value(object_value(object_value(entry)["model_info"])["id"]) for entry in entries) + + +def remove_deployment(gateway: Gateway, identity: str) -> None: + gateway.post("/model/delete", {"id": identity}) + eventually(lambda: deployment_ids(gateway), lambda ids: identity not in ids, seconds=RELOAD_SECONDS * 4) + + +def deployment( + scenario: Scenario, + *, + model_name: str = "bedrock/*", + model: str = "bedrock/*", + model_info: Mapping[str, JsonValue] | None = None, + litellm_params: Mapping[str, JsonValue] | None = None, +) -> str: + """A wildcard deployment posted straight to /model/new, since Scenario.model fixes a non-wildcard name.""" + created: Final = scenario.gateway.post( + "/model/new", + { + "model_name": model_name, + "litellm_params": {"model": model, **(litellm_params or {})}, + "model_info": dict(model_info) if model_info is not None else {}, + }, + ) + identity: Final = string_value(object_value(created["model_info"])["id"]) + scenario.cleanups.callback(remove_deployment, scenario.gateway, identity) + settled_on_every_worker(scenario.gateway, identity) + return identity + + +def settled_on_every_worker(gateway: Gateway, identity: str) -> None: + """Each worker picks a new deployment up on its own reload tick, so a listing straight after /model/new can land on + a worker without it; the by-id lookup never lists, so holding on it leaves every worker's discovery cache cold.""" + + def present() -> bool: + lookup: Final = gateway.request("GET", "/model/info", params={"litellm_model_id": identity}, headers=CLOSE) + return lookup.status_code == 200 + + held(present, lambda found: found, holding=RELOAD_SECONDS * 2, seconds=RELOAD_SECONDS * 6) + + +def sigv4_deployment( + scenario: Scenario, + key: str, + secret: str, + region: str, + *, + model_name: str = "bedrock/*", + model: str = "bedrock/*", + model_info: Mapping[str, JsonValue] | None = None, + litellm_params: Mapping[str, JsonValue] | None = None, +) -> str: + return deployment( + scenario, + model_name=model_name, + model=model, + model_info=model_info, + litellm_params={ + "aws_access_key_id": key, + "aws_secret_access_key": secret, + "aws_region_name": region, + **(litellm_params or {}), + }, + ) + + +def listed_ids( + gateway: Gateway, + path: str = "/v1/models", + *, + key: str | None = None, + params: Mapping[str, str] | None = None, +) -> frozenset[str]: + response: Final = gateway.request("GET", path, key=key, params=params, headers=CLOSE) + assert response.status_code == 200, f"GET {path}: {response.status_code} {response.text}" + data: Final = JSON_OBJECT.validate_json(response.content)["data"] + assert isinstance(data, list), response.text + return frozenset(string_value(object_value(entry)["id"]) for entry in data) + + +def from_stem(marker: str, ids: frozenset[str]) -> frozenset[str]: + return frozenset(name for name in ids if marker in name) + + +def discovered( + gateway: Gateway, + catalog: Catalog, + marker: str, + path: str = "/v1/models", + *, + key: str | None = None, + params: Mapping[str, str] | None = None, +) -> frozenset[str]: + """Poll until the catalog's invocable ids are listed; returns every listed id carrying the stem.""" + listed: Final = eventually( + lambda: listed_ids(gateway, path, key=key, params=params), + lambda ids: catalog.invocable_ids() <= ids, + seconds=RELOAD_SECONDS * 4, + ) + return from_stem(marker, listed) + + +def mine(plane: ControlPlane, credential_id: str) -> tuple[ControlPlaneRequest, ...]: + return tuple(request for request in plane.drain() if request.credential == credential_id) + + +def listings(requests: tuple[ControlPlaneRequest, ...]) -> Mapping[str, int]: + return Counter(request.path for request in requests) + + +def assert_listing_shape(requests: tuple[ControlPlaneRequest, ...], *, pages: int = 1) -> int: + """Every listing is one foundation-models GET plus `pages` inference-profiles GETs; returns the listing count.""" + assert requests, "no listing reached the control plane" + counts: Final = listings(requests) + assert set(counts) == {FOUNDATION_MODELS, INFERENCE_PROFILES}, counts + assert 1 <= counts[FOUNDATION_MODELS] <= WORKERS, counts + assert counts[INFERENCE_PROFILES] == counts[FOUNDATION_MODELS] * pages, counts + for request in requests: + assert request.method == "GET" and request.body == b"", request + if request.path == FOUNDATION_MODELS: + assert request.query == FOUNDATION_QUERY, request + else: + assert {name: value for name, value in request.query.items() if name != "nextToken"} == PROFILES_QUERY + return counts[FOUNDATION_MODELS] + + +def assert_sigv4(requests: tuple[ControlPlaneRequest, ...], *, key: str, secret: str, region: str) -> None: + host: Final = control_plane_host(region) + for request in requests: + assert request.host == host, request + assert request.headers["host"] == host, request + scope: Final = verified_scope(request, secret) + assert scope == f"{request.headers['x-amz-date'][:8]}/{region}/bedrock/aws4_request", request + assert request.headers["authorization"].startswith(f"AWS4-HMAC-SHA256 Credential={key}/"), request + assert "x-amz-security-token" not in request.headers, request + + +def worker_pids(log: Path) -> tuple[int, ...]: + return tuple(int(pid) for pid in STARTED_WORKER.findall(log.read_text())) + + +def live_worker_pids(log: Path) -> tuple[int, ...]: + return tuple(pid for pid in worker_pids(log) if psutil.pid_exists(pid)) + + +def wait_for_every_worker(log: Path, workers: int) -> None: + """Readiness answers from the first worker up; uvicorn replaces a child that dies at boot, so the rest can lag.""" + + def every_worker_serving() -> bool: + return len(live_worker_pids(log)) >= workers and log.read_text().count(STARTUP_COMPLETE) >= workers + + eventually(every_worker_serving, bool, seconds=150) + + +def wait_for_replacement_worker(log: Path, original: tuple[int, ...]) -> None: + def replacement_is_serving(pids: tuple[int, ...]) -> bool: + return len(pids) > len(original) and log.read_text().count(STARTUP_COMPLETE) > len(original) + + eventually(lambda: worker_pids(log), replacement_is_serving, seconds=150) + + +def open_connections_to(pid: int, url: str) -> int: + port: Final = urlsplit(url).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) diff --git a/tests/integration/_support/client.py b/tests/integration/_support/client.py index a57792255ab..1c58284e42f 100644 --- a/tests/integration/_support/client.py +++ b/tests/integration/_support/client.py @@ -28,6 +28,11 @@ def string_value(value: JsonValue) -> str: return value +def list_value(value: JsonValue) -> Sequence[JsonValue]: + assert isinstance(value, list), f"Expected a list, received {type(value).__name__}" + return value + + def delete_key_if_present(candidate: Gateway, key: str) -> None: digest: Final = sha256(key.encode()).hexdigest() if read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,)): @@ -52,6 +57,28 @@ def eventually( time.sleep(0.1) +def held(read: Callable[[], T], satisfied: Callable[[T], bool], *, holding: float, seconds: float = 10) -> T: + """`eventually`, and then `satisfied` must hold on every read for `holding` more seconds; a miss restarts the hold.""" + deadline: Final = time.monotonic() + seconds + while True: + converged: Final = eventually(read, satisfied, seconds=max(deadline - time.monotonic(), 0.0)) + misses: Final = _misses_within(read, satisfied, holding) + if not misses: + return converged + assert time.monotonic() < deadline, f"State did not hold: {misses[0]!r}" + time.sleep(0.1) + + +def _misses_within(read: Callable[[], T], satisfied: Callable[[T], bool], holding: float) -> tuple[T, ...]: + until: Final = time.monotonic() + holding + while time.monotonic() < until: + time.sleep(0.1) + observed: Final = read() + if not satisfied(observed): + return (observed,) + return () + + @dataclass(frozen=True, slots=True) class Gateway: client: httpx.Client diff --git a/tests/integration/_support/sigv4.py b/tests/integration/_support/sigv4.py index e02283a719d..a6abf58efef 100644 --- a/tests/integration/_support/sigv4.py +++ b/tests/integration/_support/sigv4.py @@ -2,19 +2,42 @@ import hashlib import hmac from collections.abc import Mapping from typing import Final +from urllib.parse import parse_qsl + +_UNRESERVED: Final = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_.~" -def encoded_path(value: str) -> str: - safe: Final = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_.~/" +def _encoded(value: str, safe: bytes) -> str: return "".join(chr(byte) if byte in safe else f"%{byte:02X}" for byte in value.encode("utf-8")) +def encoded_path(value: str) -> str: + return _encoded(value, _UNRESERVED + b"/") + + +def canonical_query(query: str) -> str: + """The SigV4 canonical query string: pairs sorted by name, each name and value URI-encoded.""" + return "&".join( + _encoded(name, _UNRESERVED) + "=" + _encoded(value, _UNRESERVED) + for name, value in sorted(parse_qsl(query, keep_blank_values=True)) + ) + + def signature( - method: str, path: str, headers: Mapping[str, str], signed: str, body: bytes, secret: str, scope: str, + method: str, + path: str, + headers: Mapping[str, str], + signed: str, + body: bytes, + secret: str, + scope: str, + query: str = "", ) -> tuple[str, str]: - """AWS SigV4 equations, independent of botocore and LiteLLM's signer.""" - canonical_headers: Final = "".join(name + ":" + " ".join(headers[name].split()) + "\n" for name in signed.split(";")) - canonical: Final = "\n".join((method, path, "", canonical_headers, signed, hashlib.sha256(body).hexdigest())) + """AWS SigV4 equations, independent of botocore and LiteLLM's signer; `query` is already canonical.""" + canonical_headers: Final = "".join( + name + ":" + " ".join(headers[name].split()) + "\n" for name in signed.split(";") + ) + canonical: Final = "\n".join((method, path, query, canonical_headers, signed, hashlib.sha256(body).hexdigest())) canonical_hash: Final = hashlib.sha256(canonical.encode()).hexdigest() date, region, service, terminator = scope.split("/") assert terminator == "aws4_request" diff --git a/tests/integration/providers/test_bedrock_model_discovery_chaos.py b/tests/integration/providers/test_bedrock_model_discovery_chaos.py new file mode 100644 index 00000000000..f1358837b7e --- /dev/null +++ b/tests/integration/providers/test_bedrock_model_discovery_chaos.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import signal +import threading +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import psutil +import pytest +from integration._support.aws_control_plane import ControlPlane, ControlPlaneRequest, control_plane +from integration._support.bedrock_discovery import ( + CLOSE, + CONTROL_MODEL, + HOSTED_REGIONS, + LISTING_TIMEOUT_SECONDS, + WORKERS, + catalog_for, + control_plane_host, + credential, + discovered, + discovery_proxy, + from_stem, + listed_ids, + live_worker_pids, + mine, + open_connections_to, + sigv4_deployment, + stem, + wait_for_replacement_worker, +) +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.process import graceful_stop_seconds +from integration._support.wire import Reply + +pytestmark: Final = pytest.mark.timeout(2 * graceful_stop_seconds() + 120) + + +@dataclass(frozen=True, slots=True) +class Rig: + gateway: Gateway + log: Path + plane: ControlPlane + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]: + hosts: Final = tuple(control_plane_host(region) for region in HOSTED_REGIONS) + directory: Final = tmp_path_factory.mktemp("bedrock_discovery_chaos") + with control_plane(directory, hosts) as plane, gateway_from_environment() as parent: + with discovery_proxy(parent, directory, plane, check_provider_endpoint=True, workers=WORKERS) as owned: + yield Rig(owned.gateway, owned.log, plane) + + +def _victim(pids: tuple[int, ...], plane_url: str) -> int | None: + busy: Final = tuple(pid for pid in pids if psutil.pid_exists(pid) and open_connections_to(pid, plane_url) > 0) + return busy[0] if busy else None + + +def test_worker_killed_mid_listing_leaves_the_sibling_serving_and_a_replacement_listing_again(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + release: Final = threading.Event() + attempts: Final[list[str]] = [] # mutable-ok: appended by the control plane's handler threads + + def hang(request: ControlPlaneRequest) -> Reply: + attempts.append(request.path) + if not release.is_set(): + release.wait(timeout=LISTING_TIMEOUT_SECONDS * 3) + return catalog.respond(request) + + original: Final = live_worker_pids(rig.log) + assert len(original) == WORKERS, original + try: + with rig.gateway.scenario() as scenario, rig.plane.answering(key, hang): + sigv4_deployment(scenario, key, secret, "us-east-1") + with ThreadPoolExecutor(max_workers=1) as pool: + in_flight: Final = pool.submit(lambda: rig.gateway.request("GET", "/v1/models", headers=CLOSE)) + victim: Final = eventually( + lambda: _victim(original, rig.plane.url), + lambda pid: pid is not None, + seconds=LISTING_TIMEOUT_SECONDS, + ) + assert victim is not None + psutil.Process(victim).send_signal(signal.SIGKILL) + sibling: Final = next(pid for pid in original if pid != victim) + for _ in range(5): + assert rig.gateway.chat(CONTROL_MODEL)["choices"], "the sibling worker stopped answering" + assert psutil.pid_exists(sibling) + severed: Final = in_flight.exception(timeout=LISTING_TIMEOUT_SECONDS * 2) + assert severed is not None or in_flight.result().status_code >= 500 + wait_for_replacement_worker(rig.log, original) + release.set() + assert discovered(rig.gateway, catalog, marker) == catalog.invocable_ids() + for _ in range(2 * WORKERS): + assert from_stem(marker, listed_ids(rig.gateway)) == catalog.invocable_ids() + assert attempts, "the victim never reached the control plane" + assert mine(rig.plane, key) != () + survivors: Final = frozenset(live_worker_pids(rig.log)) + replacements: Final = survivors - frozenset(original) + assert replacements, (survivors, original, victim) + assert survivors == (frozenset(original) - {victim}) | replacements, (survivors, original, victim) + finally: + release.set() diff --git a/tests/integration/providers/test_bedrock_model_discovery_wire.py b/tests/integration/providers/test_bedrock_model_discovery_wire.py new file mode 100644 index 00000000000..3bf4583e736 --- /dev/null +++ b/tests/integration/providers/test_bedrock_model_discovery_wire.py @@ -0,0 +1,815 @@ +from __future__ import annotations + +import threading +import time +import uuid +from collections import Counter +from collections.abc import Callable, Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import litellm +import pytest +from integration._support.aws_control_plane import ( + FOUNDATION_MODELS, + INFERENCE_PROFILES, + UNKNOWN_CREDENTIAL, + Catalog, + ControlPlane, + ControlPlaneRequest, + control_plane, + json_reply, +) +from integration._support.bedrock_discovery import ( + CLOSE, + CONTROL_MODEL, + CONVERSE_MODEL, + HOSTED_REGIONS, + INFO_PATHS, + LISTING_FAILURE_LOG, + LISTING_PATHS, + LISTING_TIMEOUT_SECONDS, + PROFILE_PAGE_CAP, + RELOAD_SECONDS, + UNHOSTED_REGION, + WORKERS, + assert_listing_shape, + assert_sigv4, + catalog_for, + control_plane_host, + credential, + deployment, + discovered, + discovery_proxy, + from_stem, + listed_ids, + listings, + mine, + sigv4_deployment, + stem, +) +from integration._support.bedrock_runtime_peer import MARKER, marker_of, respond, target_of +from integration._support.client import ( + JSON_OBJECT, + Gateway, + eventually, + gateway_from_environment, + list_value, + object_value, + string_value, +) +from integration._support.database import read_rows +from integration._support.process import graceful_stop_seconds +from integration._support.wire import Reply, Wire, wire_server +from pydantic import JsonValue + +pytestmark: Final = pytest.mark.timeout(2 * graceful_stop_seconds() + 120) + +STREAM_TERMINALS: Final[Mapping[str, str]] = { + "/v1/chat/completions": "data: [DONE]", + "/v1/messages": "event: message_stop", + "/v1/responses": "response.completed", +} + + +@dataclass(frozen=True, slots=True) +class Rig: + gateway: Gateway + log: Path + plane: ControlPlane + runtime: Wire + + +@pytest.fixture(scope="module") +def plane(tmp_path_factory: pytest.TempPathFactory) -> Iterator[ControlPlane]: + hosts: Final = tuple(control_plane_host(region) for region in HOSTED_REGIONS) + with control_plane(tmp_path_factory.mktemp("bedrock_control_plane"), hosts) as value: + yield value + + +@pytest.fixture(scope="module") +def runtime() -> Iterator[Wire]: + with wire_server(respond) as wire: + yield wire + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory, plane: ControlPlane, runtime: Wire) -> Iterator[Rig]: + with gateway_from_environment() as parent: + directory: Final = tmp_path_factory.mktemp("bedrock_discovery_on") + with discovery_proxy(parent, directory, plane, check_provider_endpoint=True, workers=WORKERS) as owned: + yield Rig(owned.gateway, owned.log, plane, runtime) + + +@pytest.fixture(scope="module") +def flag_off_rig(tmp_path_factory: pytest.TempPathFactory, plane: ControlPlane, runtime: Wire) -> Iterator[Rig]: + with gateway_from_environment() as parent: + directory: Final = tmp_path_factory.mktemp("bedrock_discovery_off") + with discovery_proxy(parent, directory, plane, check_provider_endpoint=False, workers=1) as owned: + yield Rig(owned.gateway, owned.log, plane, runtime) + + +def test_admin_models_list_the_accounts_invocable_models_signed_for_the_deployment_region(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1") + for path in LISTING_PATHS: + assert discovered(rig.gateway, catalog, marker, path) == catalog.invocable_ids() + assert CONTROL_MODEL in listed_ids(rig.gateway, path) + requests: Final = mine(rig.plane, key) + assert_listing_shape(requests) + assert_sigv4(requests, key=key, secret=secret, region="us-east-1") + + +def _info_rows(gateway: Gateway, path: str, marker: str, identity: str) -> Mapping[str, str]: + entries: Final = gateway.get(path)["data"] + assert isinstance(entries, list) + rows: Final = tuple(object_value(entry) for entry in entries) + return { + string_value(row["model_name"]): string_value(object_value(row["litellm_params"])["model"]) + for row in rows + if marker in string_value(row["model_name"]) and string_value(object_value(row["model_info"])["id"]) == identity + } + + +def test_model_info_expands_the_wildcard_into_one_row_per_invocable_model(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + identity: Final = sigv4_deployment(scenario, key, secret, "us-east-1") + for path in INFO_PATHS: + expanded: Final = _info_rows(rig.gateway, path, marker, identity) + assert expanded == {name: name for name in catalog.invocable_ids()}, expanded + assert_listing_shape(mine(rig.plane, key)) + + +def test_model_group_info_carries_every_invocable_model(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1") + + def groups() -> frozenset[str]: + entries: Final = rig.gateway.get("/model_group/info")["data"] + assert isinstance(entries, list) + return frozenset(string_value(object_value(entry)["model_group"]) for entry in entries) + + listed: Final = eventually(groups, lambda names: catalog.invocable_ids() <= names, seconds=RELOAD_SECONDS * 4) + assert from_stem(marker, listed) == catalog.invocable_ids() + assert_listing_shape(mine(rig.plane, key)) + + +def test_key_scoped_to_the_wildcard_lists_the_discovered_models_and_nothing_else(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1") + scoped: Final = scenario.key(models=["bedrock/*"]) + assert discovered(rig.gateway, catalog, marker, key=scoped) == catalog.invocable_ids() + assert CONTROL_MODEL not in listed_ids(rig.gateway, key=scoped) + assert_listing_shape(mine(rig.plane, key)) + + +def test_team_scoped_to_the_wildcard_lists_the_discovered_models(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1") + team: Final = scenario.team(models=["bedrock/*"]) + member: Final = scenario.key(team_id=team) + assert discovered(rig.gateway, catalog, marker, key=member) == catalog.invocable_ids() + assert_listing_shape(mine(rig.plane, key)) + + +def test_access_group_holding_the_wildcard_lists_the_discovered_models(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + group: Final = f"bedrock-group-{marker}" + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1", model_info={"access_groups": [group]}) + grouped: Final = scenario.key(models=[group]) + assert discovered(rig.gateway, catalog, marker, key=grouped) == catalog.invocable_ids() + assert_listing_shape(mine(rig.plane, key)) + + +def test_bearer_token_deployment_lists_with_the_token_and_no_signature(rig: Rig) -> None: + token: Final = f"bearer-{uuid.uuid4().hex}" + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(token, catalog.respond): + deployment(scenario, litellm_params={"api_key": token, "aws_region_name": "us-east-1"}) + assert discovered(rig.gateway, catalog, marker) == catalog.invocable_ids() + requests: Final = mine(rig.plane, token) + assert_listing_shape(requests) + for request in requests: + assert request.headers["authorization"] == f"Bearer {token}", request + assert request.host == control_plane_host("us-east-1"), request + assert "x-amz-date" not in request.headers, request + + +@pytest.mark.parametrize("region", ("eu-west-1", "us-gov-west-1", "cn-north-1")) +def test_listing_reaches_the_control_plane_of_the_deployment_region_and_partition(rig: Rig, region: str) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, region) + assert discovered(rig.gateway, catalog, marker) == catalog.invocable_ids() + requests: Final = mine(rig.plane, key) + assert_listing_shape(requests) + assert_sigv4(requests, key=key, secret=secret, region=region) + + +def _frame_id(frame: Mapping[str, JsonValue]) -> str | None: + if frame.get("type") == "message_start": + return string_value(object_value(frame["message"])["id"]) + response: Final = frame.get("response") + if isinstance(response, dict) and "id" in response: + return string_value(response["id"]) + identity: Final = frame.get("id") + return identity if isinstance(identity, str) else None + + +def _response_id(text: str, *, stream: bool) -> str: + if not stream: + return string_value(JSON_OBJECT.validate_json(text)["id"]) + frames: Final = tuple( + JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + ids: Final = tuple(identity for identity in map(_frame_id, frames) if identity is not None) + assert ids, text + return ids[-1] + + +def _call(gateway: Gateway, path: str, body: Mapping[str, JsonValue], *, key: str, tag: str) -> tuple[str, str]: + """POST one call tagged for its spend row and return (response id, response text); a stream is read to its terminal frame.""" + response: Final = gateway.request("POST", path, body, key=key, headers={**CLOSE, "x-litellm-tags": tag}) + assert response.status_code == 200, f"POST {path}: {response.status_code} {response.text}" + stream: Final = body.get("stream") is True + if stream: + assert STREAM_TERMINALS[path] in response.text, response.text + return _response_id(response.text, stream=stream), response.text + + +def _calls_on(model: str, markers: tuple[str, ...]) -> tuple[tuple[str, Mapping[str, JsonValue]], ...]: + def chat(marker: str, stream: bool) -> Mapping[str, JsonValue]: + return {"model": model, "messages": [{"role": "user", "content": f"marker-{marker}"}], "stream": stream} + + def messages(marker: str, stream: bool) -> Mapping[str, JsonValue]: + return {**chat(marker, stream), "max_tokens": 64} + + def responses(marker: str, stream: bool) -> Mapping[str, JsonValue]: + return {"model": model, "input": f"marker-{marker}", "stream": stream} + + return ( + ("/v1/chat/completions", chat(markers[0], False)), + ("/v1/chat/completions", chat(markers[1], True)), + ("/v1/messages", messages(markers[2], False)), + ("/v1/messages", messages(markers[3], True)), + ("/v1/responses", responses(markers[4], False)), + ("/v1/responses", responses(markers[5], True)), + ) + + +def _chat_rows_landed(ids: tuple[str, ...]) -> None: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s)', + (list(ids),), # pyright: ignore[reportArgumentType] # psycopg adapts the list to a text array + ), + lambda found: len(found) >= len(ids), + seconds=60, + ) + assert sorted(string_value(row["request_id"]) for row in rows) == sorted(ids), rows + + +def _spend_row_id(rows: Sequence[Mapping[str, JsonValue]], tag: str) -> str: + tagged: Final = tuple(row for row in rows if tag in list_value(row["request_tags"])) + assert len(tagged) == 1, (tag, rows) + return string_value(tagged[0]["request_id"]) + + +def _landed_once(served: tuple[tuple[str, str, bool], ...], tags: tuple[str, ...]) -> None: + """One spend row per tagged call, carrying the id the client saw. + + FIXME: a streamed /v1/responses row carries the managed `resp_` id LiteLLM built before the proxy encrypted the + advertised one, so that row is matched by tag alone until the spend log reads the id the client received. + """ + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, request_tags FROM "LiteLLM_SpendLogs" WHERE request_tags ?| %s', + (list(tags),), # pyright: ignore[reportArgumentType] # psycopg adapts the list to a text array + ), + lambda found: len(found) >= len(tags), + seconds=60, + ) + for (identity, path, stream), tag in zip(served, tags, strict=True): + row_id: Final = _spend_row_id(rows, tag) + if path == "/v1/responses" and stream: + assert row_id.startswith("resp_"), (row_id, identity) + continue + assert row_id == identity, (path, stream, row_id, identity) + + +def test_discovered_model_is_callable_through_the_wildcard_on_every_endpoint(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = Catalog(active_profiles=(f"us.{marker}.sonnet-v1:0",), on_demand_models=(CONVERSE_MODEL,)) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment( + scenario, key, secret, "us-east-1", litellm_params={"aws_bedrock_runtime_endpoint": rig.runtime.url} + ) + scoped: Final = scenario.key(models=["bedrock/*"]) + assert discovered(rig.gateway, catalog, marker, key=scoped) == from_stem(marker, catalog.invocable_ids()) + assert f"bedrock/{CONVERSE_MODEL}" in listed_ids(rig.gateway, key=scoped) + rig.runtime.drain() + markers: Final = tuple(uuid.uuid4().hex for _ in range(6)) + calls: Final = _calls_on(f"bedrock/{CONVERSE_MODEL}", markers) + served: Final = tuple( + _call(rig.gateway, path, body, key=scoped, tag=f"discovery-{call_marker}") + for (path, body), call_marker in zip(calls, markers, strict=True) + ) + for (identity, text), call_marker in zip(served, markers, strict=True): + assert identity, text + assert set(MARKER.findall(text)) == {call_marker}, text + received: Final = rig.runtime.drain() + assert sorted(marker_of(request) for request in received) == sorted(markers), received + for request in received: + assert target_of(request).startswith("/openai/v1/"), request.target + assert request.headers["authorization"].startswith(f"AWS4-HMAC-SHA256 Credential={key}/"), request + _landed_once( + tuple( + (identity, path, body.get("stream") is True) + for (identity, _), (path, body) in zip(served, calls, strict=True) + ), + tuple(f"discovery-{call_marker}" for call_marker in markers), + ) + assert_listing_shape(mine(rig.plane, key)) + + +def test_flag_off_proxy_lists_the_static_catalog_and_never_calls_the_control_plane(flag_off_rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + static: Final = frozenset( + name if name.startswith("bedrock/") else f"bedrock/{name}" for name in litellm.models_by_provider["bedrock"] + ) + with flag_off_rig.gateway.scenario() as scenario, flag_off_rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1") + listed: Final = eventually( + lambda: listed_ids(flag_off_rig.gateway), lambda ids: static <= ids, seconds=RELOAD_SECONDS * 4 + ) + assert from_stem(marker, listed) == frozenset() + for _ in range(3): + assert from_stem(marker, listed_ids(flag_off_rig.gateway)) == frozenset() + assert mine(flag_off_rig.plane, key) == () + + +def _sad_listing(rig: Rig, answer: Callable[[ControlPlaneRequest], Reply]) -> tuple[ControlPlaneRequest, ...]: + """Deploy against a control plane answering `answer`; the proxy keeps listing without the account's ids.""" + key, secret = credential() + marker: Final = stem() + with rig.gateway.scenario() as scenario, rig.plane.answering(key, answer): + sigv4_deployment(scenario, key, secret, "us-east-1") + for path in LISTING_PATHS: + for _ in range(3): + listed: Final = listed_ids(rig.gateway, path) + assert from_stem(marker, listed) == frozenset(), listed + assert CONTROL_MODEL in listed + assert rig.gateway.chat(CONTROL_MODEL)["choices"], "the control deployment stopped answering" + return mine(rig.plane, key) + + +@pytest.mark.parametrize( + "answer", + ( + pytest.param(lambda _: json_reply(403, {"message": "AccessDeniedException"}), id="403"), + pytest.param(lambda _: json_reply(404, {"message": "ResourceNotFoundException"}), id="404"), + pytest.param(lambda _: Reply(body=b"upstream maintenance", content_type="text/html"), id="html"), + pytest.param(lambda _: json_reply(200, {"modelSummaries": "nope", "inferenceProfileSummaries": 7}), id="shape"), + pytest.param(lambda _: json_reply(200, {"modelSummaries": [], "inferenceProfileSummaries": []}), id="empty"), + ), +) +def test_control_plane_errors_leave_the_listing_without_bedrock_ids_and_the_proxy_serving( + rig: Rig, answer: Callable[[ControlPlaneRequest], Reply] +) -> None: + requests: Final = _sad_listing(rig, answer) + assert requests, "the proxy never asked the control plane" + assert set(listings(requests)) <= {FOUNDATION_MODELS, INFERENCE_PROFILES}, listings(requests) + + +def test_deployment_without_credentials_never_reaches_the_control_plane(rig: Rig) -> None: + marker: Final = stem() + rig.plane.drain() + with rig.gateway.scenario() as scenario: + deployment(scenario, litellm_params={"aws_region_name": "us-east-1"}) + for _ in range(6): + listed: Final = listed_ids(rig.gateway) + assert from_stem(marker, listed) == frozenset() and CONTROL_MODEL in listed + assert tuple(request for request in rig.plane.drain() if request.credential == UNKNOWN_CREDENTIAL) == () + + +def test_refused_egress_leaves_the_listing_without_bedrock_ids(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + rig.plane.refusals() + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog_for(marker).respond): + sigv4_deployment(scenario, key, secret, UNHOSTED_REGION) + for _ in range(6): + assert from_stem(marker, listed_ids(rig.gateway)) == frozenset() + assert mine(rig.plane, key) == () + assert f"{control_plane_host(UNHOSTED_REGION)}:443" in rig.plane.refusals() + + +def test_hung_control_plane_is_bounded_by_the_listing_timeout_while_the_proxy_keeps_serving( + rig: Rig, record_property: Callable[[str, object], None] +) -> None: + key, secret = credential() + marker: Final = stem() + release: Final = threading.Event() + catalog: Final = catalog_for(marker) + + def hang(request: ControlPlaneRequest) -> Reply: + assert release.wait(timeout=120), "the hung listing was never released" + return catalog.respond(request) + + stalled: Final = threading.Event() + latencies: Final[list[float]] = [] # mutable-ok: appended by the probe thread + + def probe() -> None: + while not stalled.is_set(): + started: Final = time.monotonic() + health: Final = rig.gateway.request("GET", "/health/liveliness", headers=CLOSE) + latencies.append(time.monotonic() - started) + assert health.status_code == 200, health.text + + failures_before: Final = rig.log.read_text().count(LISTING_FAILURE_LOG) + try: + with rig.gateway.scenario() as scenario, rig.plane.answering(key, hang): + sigv4_deployment(scenario, key, secret, "us-east-1") + prober: Final = threading.Thread(target=probe) + prober.start() + started: Final = time.monotonic() + listed: Final = listed_ids(rig.gateway) + elapsed: Final = time.monotonic() - started + assert rig.gateway.chat(CONTROL_MODEL)["choices"] + stalled.set() + prober.join(timeout=30) + assert not prober.is_alive() + assert from_stem(marker, listed) == frozenset() and CONTROL_MODEL in listed + assert LISTING_TIMEOUT_SECONDS <= elapsed < LISTING_TIMEOUT_SECONDS * 2, elapsed + record_property("listing_seconds", round(elapsed, 2)) + record_property("max_liveliness_seconds", round(max(latencies), 2)) + eventually( + lambda: rig.log.read_text().count(LISTING_FAILURE_LOG), + lambda count: count > failures_before, + seconds=10, + ) + assert mine(rig.plane, key) != () + release.set() + finally: + release.set() + + +def test_endless_pagination_stops_at_the_page_cap(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + + def endless(request: ControlPlaneRequest) -> Reply: + page: Final = int(request.query.get("nextToken", "page-0").removeprefix("page-")) + return json_reply(200, {"inferenceProfileSummaries": [], "nextToken": f"page-{page + 1}"}) + + with rig.gateway.scenario() as scenario, rig.plane.answering(key, endless): + sigv4_deployment(scenario, key, secret, "us-east-1") + for _ in range(3): + assert from_stem(marker, listed_ids(rig.gateway)) == frozenset() + requests: Final = mine(rig.plane, key) + counts: Final = listings(requests) + assert set(counts) == {INFERENCE_PROFILES}, counts + assert counts[INFERENCE_PROFILES] % PROFILE_PAGE_CAP == 0 and counts[INFERENCE_PROFILES] > 0, counts + tokens: Final = tuple(request.query.get("nextToken", "page-0") for request in requests) + expected: Final = tuple(f"page-{index}" for index in range(PROFILE_PAGE_CAP)) + assert tokens == expected * (counts[INFERENCE_PROFILES] // PROFILE_PAGE_CAP), tokens + + +@pytest.mark.parametrize("region", (pytest.param("", id="empty"), pytest.param("x" * 5120, id="5kb"))) +def test_unusable_region_never_reaches_a_control_plane(rig: Rig, region: str) -> None: + key, secret = credential() + marker: Final = stem() + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog_for(marker).respond): + sigv4_deployment(scenario, key, secret, region) + for _ in range(6): + listed: Final = listed_ids(rig.gateway) + assert from_stem(marker, listed) == frozenset() and CONTROL_MODEL in listed + assert rig.gateway.chat(CONTROL_MODEL)["choices"] + assert mine(rig.plane, key) == () + + +@pytest.mark.parametrize("field", ("aws_region_name", "api_key")) +def test_non_string_credential_fields_are_refused_at_model_creation(rig: Rig, field: str) -> None: + key, secret = credential() + response: Final = rig.gateway.request( + "POST", + "/model/new", + { + "model_name": "bedrock/*", + "litellm_params": { + "model": "bedrock/*", + "aws_access_key_id": key, + "aws_secret_access_key": secret, + "aws_region_name": "us-east-1", + field: 12345, + }, + "model_info": {}, + }, + ) + assert response.status_code in (400, 422), response.text + assert field in response.text, response.text + assert CONTROL_MODEL in listed_ids(rig.gateway) + + +def test_empty_api_key_falls_back_to_sigv4(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1", litellm_params={"api_key": ""}) + assert discovered(rig.gateway, catalog, marker) == catalog.invocable_ids() + requests: Final = mine(rig.plane, key) + assert_listing_shape(requests) + assert_sigv4(requests, key=key, secret=secret, region="us-east-1") + + +def test_repeated_listings_are_served_from_the_cache_once_per_worker(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1") + assert discovered(rig.gateway, catalog, marker) == catalog.invocable_ids() + for path in LISTING_PATHS * 3: + assert from_stem(marker, listed_ids(rig.gateway, path)) == catalog.invocable_ids() + assert assert_listing_shape(mine(rig.plane, key)) <= WORKERS + + +def test_two_regions_list_the_union_from_each_regions_control_plane(rig: Rig) -> None: + first_key, first_secret = credential() + second_key, second_secret = credential() + marker: Final = stem() + first: Final = Catalog(active_profiles=(f"us.{marker}.sonnet-v1:0",), on_demand_models=(f"{marker}.us-only",)) + second: Final = Catalog(active_profiles=(f"eu.{marker}.haiku-v1:0",), on_demand_models=(f"{marker}.eu-only",)) + with ( + rig.gateway.scenario() as scenario, + rig.plane.answering(first_key, first.respond), + rig.plane.answering(second_key, second.respond), + ): + sigv4_deployment(scenario, first_key, first_secret, "us-east-1") + sigv4_deployment(scenario, second_key, second_secret, "eu-west-1") + union: Final = Catalog( + active_profiles=first.active_profiles + second.active_profiles, + on_demand_models=first.on_demand_models + second.on_demand_models, + ) + assert discovered(rig.gateway, union, marker) == union.invocable_ids() + assert_sigv4(mine(rig.plane, first_key), key=first_key, secret=first_secret, region="us-east-1") + assert_sigv4(mine(rig.plane, second_key), key=second_key, secret=second_secret, region="eu-west-1") + + +def test_partial_wildcard_lists_only_the_matching_models_under_their_real_ids(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = Catalog( + active_profiles=(f"us.anthropic.{marker}-v1:0",), + on_demand_models=(f"anthropic.{marker}-v1:0", f"amazon.{marker}-v1:0"), + ) + expected: Final = frozenset({f"bedrock/anthropic.{marker}-v1:0"}) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment( + scenario, key, secret, "us-east-1", model_name="bedrock/anthropic.*", model="bedrock/anthropic.*" + ) + listed: Final = eventually( + lambda: from_stem(marker, listed_ids(rig.gateway)), + lambda ids: ids != frozenset(), + seconds=RELOAD_SECONDS * 4, + ) + assert listed == expected, listed + + +def test_partial_wildcard_matching_no_invocable_id_lists_nothing_instead_of_alias_names(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = Catalog( + active_profiles=(f"us.anthropic.{marker}-v1:0",), + on_demand_models=(f"amazon.{marker}-v1:0",), + ) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1", model_name="bedrock/*", model="bedrock/*") + sigv4_deployment( + scenario, key, secret, "us-east-1", model_name="bedrock/anthropic.*", model="bedrock/anthropic.*" + ) + listed: Final = eventually( + lambda: from_stem(marker, listed_ids(rig.gateway)), + lambda ids: ids != frozenset(), + seconds=RELOAD_SECONDS * 4, + ) + assert listed == catalog.invocable_ids(), listed + + +@pytest.mark.parametrize( + "prefix_of", + ( + pytest.param(lambda marker: f"team-{marker}", id="distinct"), + pytest.param(lambda marker: marker, id="starts-a-discovered-id"), + ), +) +def test_custom_prefix_wildcard_lists_the_discovered_models_under_that_prefix( + rig: Rig, prefix_of: Callable[[str], str] +) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + prefix: Final = prefix_of(marker) + expected: Final = frozenset(name.replace("bedrock/", f"{prefix}/", 1) for name in catalog.invocable_ids()) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1", model_name=f"{prefix}/*") + listed: Final = eventually( + lambda: from_stem(marker, listed_ids(rig.gateway)), lambda ids: expected <= ids, seconds=RELOAD_SECONDS * 4 + ) + assert listed == expected, listed + assert_listing_shape(mine(rig.plane, key)) + + +def test_wildcard_route_is_listed_alongside_the_discovered_models_when_asked(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1") + params: Final = {"return_wildcard_routes": "true"} + assert discovered(rig.gateway, catalog, marker, params=params) == catalog.invocable_ids() + assert "bedrock/*" in listed_ids(rig.gateway, params=params) + assert "bedrock/*" not in listed_ids(rig.gateway) + + +def test_inactive_profiles_are_excluded_and_a_model_in_both_listings_appears_once(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + shared: Final = f"{marker}.shared-v1:0" + catalog: Final = Catalog( + active_profiles=(shared, f"us.{marker}.sonnet-v1:0"), + inactive_profiles=(f"us.{marker}.retired-v1:0", f"eu.{marker}.retired-v1:0"), + on_demand_models=(shared,), + ) + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + sigv4_deployment(scenario, key, secret, "us-east-1") + assert discovered(rig.gateway, catalog, marker) == catalog.invocable_ids() + response: Final = rig.gateway.request("GET", "/v1/models", headers=CLOSE) + data: Final = JSON_OBJECT.validate_json(response.content)["data"] + assert isinstance(data, list) + names: Final = [string_value(object_value(entry)["id"]) for entry in data] + assert names.count(f"bedrock/{shared}") == 1, names + assert not any("retired" in name for name in names), names + + +def test_concurrent_cold_listings_all_answer_with_the_discovered_models( + rig: Rig, record_property: Callable[[str, object], None] +) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker, page_size=1) + + def slow(request: ControlPlaneRequest) -> Reply: + time.sleep(0.3) + return catalog.respond(request) + + with rig.gateway.scenario() as scenario, rig.plane.answering(key, slow): + sigv4_deployment(scenario, key, secret, "us-east-1") + with ThreadPoolExecutor(max_workers=8) as pool: + answers: Final = tuple(pool.map(lambda _: listed_ids(rig.gateway), range(8))) + for ids in answers: + assert from_stem(marker, ids) == catalog.invocable_ids(), ids + requests: Final = mine(rig.plane, key) + counts: Final = listings(requests) + assert counts[INFERENCE_PROFILES] == counts[FOUNDATION_MODELS] * len(catalog.pages()), counts + assert 1 <= counts[FOUNDATION_MODELS] <= len(answers), counts + record_property("listings_during_burst", counts[FOUNDATION_MODELS]) + + +def test_region_update_while_listings_flow_moves_the_listing_to_the_new_regions_control_plane(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + stop: Final = threading.Event() + statuses: Final[list[int]] = [] # mutable-ok: appended by the polling thread + hosts: Final[set[str]] = set() # mutable-ok: collects every control plane host the listings reached + + def poll() -> None: + while not stop.is_set(): + statuses.append(rig.gateway.request("GET", "/v1/models", headers=CLOSE).status_code) + + def reached() -> frozenset[str]: + hosts.update(request.host for request in mine(rig.plane, key)) + return frozenset(hosts) + + with rig.gateway.scenario() as scenario, rig.plane.answering(key, catalog.respond): + identity: Final = sigv4_deployment(scenario, key, secret, "us-east-1") + assert discovered(rig.gateway, catalog, marker) == catalog.invocable_ids() + poller: Final = threading.Thread(target=poll) + poller.start() + try: + updated: Final = rig.gateway.request( + "PATCH", f"/model/{identity}/update", {"litellm_params": {"aws_region_name": "eu-west-1"}} + ) + assert updated.status_code == 200, updated.text + eventually(reached, lambda seen: control_plane_host("eu-west-1") in seen, seconds=RELOAD_SECONDS * 4) + finally: + stop.set() + poller.join(timeout=30) + assert not poller.is_alive() + assert statuses and set(statuses) == {200}, Counter(statuses) + assert from_stem(marker, listed_ids(rig.gateway)) == catalog.invocable_ids() + + +def _burst(gateway: Gateway, count: int) -> tuple[tuple[str, int, str], ...]: + """`count` concurrent calls: model listings and control chats, a third of the chats streamed.""" + + def one(index: int) -> tuple[str, int, str]: + if index % 2 == 0: + listing: Final = gateway.request("GET", "/v1/models", headers=CLOSE) + return "models", listing.status_code, listing.text + body: Final[Mapping[str, JsonValue]] = { + "model": CONTROL_MODEL, + "messages": [{"role": "user", "content": f"burst {index}"}], + "stream": index % 3 == 0, + } + chat: Final = gateway.request("POST", "/v1/chat/completions", body, headers=CLOSE) + return "chat", chat.status_code, chat.text + + with ThreadPoolExecutor(max_workers=count) as pool: + return tuple(pool.map(one, range(count))) + + +def _chat_ids(served: tuple[tuple[str, int, str], ...]) -> tuple[str, ...]: + return tuple(_response_id(text, stream=text.startswith("data: ")) for kind, _, text in served if kind == "chat") + + +def _listed_stems(marker: str, served: tuple[tuple[str, int, str], ...]) -> tuple[frozenset[str], ...]: + def ids(text: str) -> frozenset[str]: + data: Final = JSON_OBJECT.validate_json(text)["data"] + assert isinstance(data, list), text + return from_stem(marker, frozenset(string_value(object_value(entry)["id"]) for entry in data)) + + return tuple(ids(text) for kind, _, text in served if kind == "models") + + +def test_control_plane_outage_mid_burst_recovers_without_a_proxy_restart(rig: Rig) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + up: Final = threading.Event() + + def flaky(request: ControlPlaneRequest) -> Reply: + return catalog.respond(request) if up.is_set() else Reply(drop_connection=True) + + with rig.gateway.scenario() as scenario, rig.plane.answering(key, flaky): + sigv4_deployment(scenario, key, secret, "us-east-1") + served: Final = _burst(rig.gateway, 20) + assert {status for _, status, _ in served} == {200}, served + assert set(_listed_stems(marker, served)) == {frozenset()}, served + assert mine(rig.plane, key) != (), "the outage was never attempted" + up.set() + assert discovered(rig.gateway, catalog, marker) == catalog.invocable_ids() + _chat_rows_landed(_chat_ids(served)) + + +def test_slow_control_plane_under_a_burst_answers_everything_without_a_deadlock( + rig: Rig, record_property: Callable[[str, object], None] +) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + + def slow(request: ControlPlaneRequest) -> Reply: + time.sleep(1.5) + return catalog.respond(request) + + with rig.gateway.scenario() as scenario, rig.plane.answering(key, slow): + sigv4_deployment(scenario, key, secret, "us-east-1") + started: Final = time.monotonic() + served: Final = _burst(rig.gateway, 20) + record_property("burst_seconds", round(time.monotonic() - started, 2)) + assert {status for _, status, _ in served} == {200}, served + assert set(_listed_stems(marker, served)) <= {frozenset(), catalog.invocable_ids()}, served + assert discovered(rig.gateway, catalog, marker) == catalog.invocable_ids() + record_property("listings_during_burst", assert_listing_shape(mine(rig.plane, key))) + _chat_rows_landed(_chat_ids(served)) diff --git a/tests/integration/sdk/test_bedrock_model_discovery_sdk.py b/tests/integration/sdk/test_bedrock_model_discovery_sdk.py new file mode 100644 index 00000000000..95c3fc022cc --- /dev/null +++ b/tests/integration/sdk/test_bedrock_model_discovery_sdk.py @@ -0,0 +1,130 @@ +from __future__ import annotations + +import json +import os +import subprocess +import sys +from collections.abc import Iterator, Mapping, Sequence +from pathlib import Path +from typing import Final + +import pytest +from integration._support.aws_control_plane import FOUNDATION_MODELS, INFERENCE_PROFILES, ControlPlane, control_plane +from integration._support.bedrock_discovery import ( + assert_listing_shape, + assert_sigv4, + catalog_for, + control_plane_host, + credential, + listings, + mine, + stem, +) +from pydantic import JsonValue, TypeAdapter + +REGION: Final = "us-east-1" +IDS: Final = TypeAdapter(list[str]) +SCRIPT: Final = """ +import json, sys +import litellm +from litellm.types.router import LiteLLM_Params +arguments = json.loads(sys.argv[1]) +params = LiteLLM_Params(**arguments["litellm_params"]) if arguments["litellm_params"] is not None else None +models = litellm.get_valid_models( + check_provider_endpoint=True, custom_llm_provider=arguments["provider"], litellm_params=params +) +print(json.dumps(sorted(models))) +""" + + +@pytest.fixture(scope="module") +def plane(tmp_path_factory: pytest.TempPathFactory) -> Iterator[ControlPlane]: + with control_plane(tmp_path_factory.mktemp("bedrock_control_plane_sdk"), (control_plane_host(REGION),)) as value: + yield value + + +def _sdk_listing( + plane: ControlPlane, + cwd: Path, + *, + provider: str | None, + litellm_params: Mapping[str, JsonValue] | None, + process_env: Mapping[str, str] | None = None, +) -> Sequence[str]: + environment: Final = { + **{name: value for name, value in os.environ.items() if not name.startswith("AWS_")}, + **plane.environment(), + "AWS_CONFIG_FILE": str(cwd / "empty-aws-config"), + "AWS_SHARED_CREDENTIALS_FILE": str(cwd / "empty-aws-config"), + "AWS_EC2_METADATA_DISABLED": "true", + "LITELLM_LOCAL_MODEL_COST_MAP": "True", + **(process_env or {}), + } + (cwd / "empty-aws-config").write_text("") + arguments: Final = json.dumps({"provider": provider, "litellm_params": litellm_params}) + completed: Final = subprocess.run( + [sys.executable, "-P", "-c", SCRIPT, arguments], cwd=cwd, env=environment, capture_output=True, text=True + ) + assert completed.returncode == 0, completed.stderr + return IDS.validate_json(completed.stdout.strip().splitlines()[-1]) + + +def test_sdk_get_valid_models_lists_the_account_through_the_deployment_credentials( + plane: ControlPlane, tmp_path: Path +) -> None: + key, secret = credential() + marker: Final = stem() + catalog: Final = catalog_for(marker) + with plane.answering(key, catalog.respond): + listed: Final = _sdk_listing( + plane, + tmp_path, + provider="bedrock", + litellm_params={ + "model": "bedrock/*", + "aws_access_key_id": key, + "aws_secret_access_key": secret, + "aws_region_name": REGION, + }, + ) + assert frozenset(name for name in listed if marker in name) == catalog.vendor_ids(), listed + requests: Final = mine(plane, key) + assert assert_listing_shape(requests) == 1 + assert_sigv4(requests, key=key, secret=secret, region=REGION) + + +def test_sdk_get_valid_models_with_only_aws_environment_never_infers_bedrock( + plane: ControlPlane, tmp_path: Path +) -> None: + key, secret = credential() + marker: Final = stem() + with plane.answering(key, catalog_for(marker).respond): + listed: Final = _sdk_listing( + plane, + tmp_path, + provider=None, + litellm_params=None, + process_env={"AWS_ACCESS_KEY_ID": key, "AWS_SECRET_ACCESS_KEY": secret, "AWS_REGION_NAME": REGION}, + ) + assert not any(name.startswith("bedrock/") for name in listed), listed + assert mine(plane, key) == () + + +def test_sdk_get_valid_models_infers_bedrock_from_its_api_key_and_lists_with_the_bearer_token( + plane: ControlPlane, tmp_path: Path +) -> None: + token: Final = f"bearer-{stem()}" + marker: Final = stem() + catalog: Final = catalog_for(marker) + with plane.answering(token, catalog.respond): + listed: Final = _sdk_listing( + plane, + tmp_path, + provider=None, + litellm_params=None, + process_env={"BEDROCK_API_KEY": token, "AWS_BEARER_TOKEN_BEDROCK": token, "AWS_REGION_NAME": REGION}, + ) + assert frozenset(name for name in listed if marker in name) == catalog.vendor_ids(), listed + requests: Final = mine(plane, token) + assert set(listings(requests)) == {FOUNDATION_MODELS, INFERENCE_PROFILES}, listings(requests) + assert all(request.headers["authorization"] == f"Bearer {token}" for request in requests), requests diff --git a/tests/unit/llms/bedrock/test_bedrock_model_listing.py b/tests/unit/llms/bedrock/test_bedrock_model_listing.py new file mode 100644 index 00000000000..ca57b5bad82 --- /dev/null +++ b/tests/unit/llms/bedrock/test_bedrock_model_listing.py @@ -0,0 +1,166 @@ +import hashlib +import hmac +import re +from collections.abc import Callable, Mapping +from functools import reduce +from typing import Final +from urllib.parse import parse_qsl, quote + +import httpx +import pytest + +from litellm.llms.bedrock.common_utils import BedrockError, BedrockModelInfo +from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.types.router import LiteLLM_Params + +ACCESS_KEY: Final = "AKIAIOSFODNN7EXAMPLE" +SECRET_KEY: Final = "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY" +REGION: Final = "eu-central-1" +BEARER_TOKEN: Final = "bedrock-api-key-for-tests" +SECOND_PAGE_TOKEN: Final = "page2+token/with=reserved chars" + +ON_DEMAND_MODELS: Final = ( + {"modelId": "anthropic.claude-3-haiku-20240307-v1:0", "inferenceTypesSupported": ["ON_DEMAND"]}, + {"modelId": "amazon.nova-micro-v1:0", "inferenceTypesSupported": ["ON_DEMAND"]}, +) +PROFILE_ONLY_MODEL: Final = { + "modelId": "anthropic.claude-opus-4-5-20251101-v1:0", + "inferenceTypesSupported": ["INFERENCE_PROFILE"], +} +PROFILE_PAGES: Final = { + None: { + "inferenceProfileSummaries": [ + {"inferenceProfileId": "global.anthropic.claude-opus-4-5-20251101-v1:0", "status": "ACTIVE"}, + {"inferenceProfileId": "eu.anthropic.claude-3-haiku-20240307-v1:0", "status": "INACTIVE"}, + ], + "nextToken": SECOND_PAGE_TOKEN, + }, + SECOND_PAGE_TOKEN: { + "inferenceProfileSummaries": [{"inferenceProfileId": "eu.amazon.nova-micro-v1:0", "status": "ACTIVE"}], + }, +} +EXPECTED_MODELS: Final = [ + "amazon.nova-micro-v1:0", + "anthropic.claude-3-haiku-20240307-v1:0", + "eu.amazon.nova-micro-v1:0", + "global.anthropic.claude-opus-4-5-20251101-v1:0", +] +SIGV4_AUTHORIZATION: Final = re.compile( + rf"^AWS4-HMAC-SHA256 Credential={ACCESS_KEY}/\d{{8}}/{REGION}/bedrock/aws4_request, " + r"SignedHeaders=([^,]+), Signature=([0-9a-f]{64})$" +) + + +@pytest.fixture(autouse=True) +def no_ambient_bearer_token(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + + +def _sigv4_signature_as_aws_computes_it(request: httpx.Request, signed_headers: tuple[str, ...]) -> str: + amz_date: Final = request.headers["x-amz-date"] + canonical_query: Final = "&".join( + f"{quote(name, safe='-_.~')}={quote(value, safe='-_.~')}" + for name, value in sorted(parse_qsl(request.url.query.decode(), keep_blank_values=True)) + ) + canonical_headers: Final = "".join(f"{name}:{request.headers[name].strip()}\n" for name in signed_headers) + canonical_request: Final = "\n".join( + ( + request.method, + request.url.path, + canonical_query, + canonical_headers, + ";".join(signed_headers), + hashlib.sha256(request.content).hexdigest(), + ) + ) + scope: Final = f"{amz_date[:8]}/{REGION}/bedrock/aws4_request" + string_to_sign: Final = "\n".join( + ("AWS4-HMAC-SHA256", amz_date, scope, hashlib.sha256(canonical_request.encode()).hexdigest()) + ) + signing_key: Final = reduce( + lambda key, part: hmac.new(key, part.encode(), hashlib.sha256).digest(), + (amz_date[:8], REGION, "bedrock", "aws4_request"), + f"AWS4{SECRET_KEY}".encode(), + ) + return hmac.new(signing_key, string_to_sign.encode(), hashlib.sha256).hexdigest() + + +def _authorized(request: httpx.Request, bearer_token: str | None) -> bool: + authorization: Final = request.headers.get("authorization", "") + if bearer_token is not None: + return authorization == f"Bearer {bearer_token}" + signed: Final = SIGV4_AUTHORIZATION.match(authorization) + if signed is None or "x-amz-date" not in request.headers: + return False + signed_headers: Final = tuple(signed.group(1).split(";")) + if "host" not in signed_headers or "x-amz-date" not in signed_headers: + return False + return signed.group(2) == _sigv4_signature_as_aws_computes_it(request, signed_headers) + + +def _control_plane(bearer_token: str | None = None) -> Callable[[httpx.Request], httpx.Response]: + def handler(request: httpx.Request) -> httpx.Response: + if request.url.host != f"bedrock.{REGION}.amazonaws.com": + return httpx.Response(404, json={"message": f"no Bedrock at {request.url.host}"}) + if not _authorized(request, bearer_token): + return httpx.Response(403, json={"message": "not signed for this account, region and service"}) + if request.url.path == "/foundation-models": + on_demand_only: Final = request.url.params.get("byInferenceType") == "ON_DEMAND" + summaries: Final = ON_DEMAND_MODELS if on_demand_only else (*ON_DEMAND_MODELS, PROFILE_ONLY_MODEL) + return httpx.Response(200, json={"modelSummaries": list(summaries)}) + if request.url.path == "/inference-profiles": + if request.url.params.get("typeEquals") != "SYSTEM_DEFINED": + return httpx.Response(400, json={"message": "application profiles are not listable here"}) + return httpx.Response(200, json=PROFILE_PAGES[request.url.params.get("nextToken")]) + return httpx.Response(404, json={"message": f"unknown path {request.url.path}"}) + + return handler + + +def _bedrock(handler: Callable[[httpx.Request], httpx.Response]) -> BedrockModelInfo: + return BedrockModelInfo(client=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler)))) + + +def _deployment(region: str = REGION, api_key: str | None = None) -> Mapping[str, object]: + return LiteLLM_Params( + model="bedrock/*", + aws_access_key_id=ACCESS_KEY, + aws_secret_access_key=SECRET_KEY, + aws_region_name=region, + api_key=api_key, + ).model_dump(exclude_none=True) + + +def test_lists_active_profiles_and_on_demand_models_signed_for_the_deployment() -> None: + assert _bedrock(_control_plane()).discover_models(_deployment()) == EXPECTED_MODELS + + +def test_lists_in_the_deployment_region_not_the_ambient_one(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AWS_REGION_NAME", REGION) + + with pytest.raises(BedrockError, match=r"us-west-2.*404"): + _bedrock(_control_plane()).discover_models(_deployment(region="us-west-2")) + + +def test_bearer_token_api_key_replaces_sigv4() -> None: + bedrock: Final = _bedrock(_control_plane(bearer_token=BEARER_TOKEN)) + + assert bedrock.discover_models(_deployment(api_key=BEARER_TOKEN)) == EXPECTED_MODELS + with pytest.raises(BedrockError, match="403"): + bedrock.discover_models(_deployment()) + + +def test_get_models_without_a_deployment_uses_ambient_credentials(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("AWS_ACCESS_KEY_ID", ACCESS_KEY) + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", SECRET_KEY) + monkeypatch.setenv("AWS_REGION_NAME", REGION) + + assert _bedrock(_control_plane()).get_models() == EXPECTED_MODELS + + +def test_listing_failure_names_the_region_and_response() -> None: + def throttled(request: httpx.Request) -> httpx.Response: + return httpx.Response(429, json={"message": "Too many requests"}) + + with pytest.raises(BedrockError, match=rf"{REGION}.*429.*Too many requests"): + _bedrock(throttled).discover_models(_deployment()) diff --git a/tests/unit/proxy/auth/test_model_checks.py b/tests/unit/proxy/auth/test_model_checks.py index 88c2d9c684a..9f9a6687210 100644 --- a/tests/unit/proxy/auth/test_model_checks.py +++ b/tests/unit/proxy/auth/test_model_checks.py @@ -1,6 +1,9 @@ +from typing import Final from unittest.mock import patch +import httpx import pytest +import respx def test_get_team_models_for_all_models_and_team_only_models(): @@ -11,9 +14,7 @@ def test_get_team_models_for_all_models_and_team_only_models(): model_access_groups = {} include_model_access_groups = False - result = get_team_models( - team_models, proxy_model_list, model_access_groups, include_model_access_groups - ) + result: Final = get_team_models(team_models, proxy_model_list, model_access_groups, include_model_access_groups) combined_models = team_models + proxy_model_list assert set(result) == set(combined_models) @@ -246,9 +247,7 @@ def test_get_key_models_does_not_mutate_input(): ), ], ) -def test_get_complete_model_list_order( - key_models, team_models, proxy_model_list, model_list, expected -): +def test_get_complete_model_list_order(key_models, team_models, proxy_model_list, model_list, expected): """ Test that get_complete_model_list preserves order """ @@ -401,9 +400,7 @@ def test_wildcard_credential_hydration_preserves_deployment_params( captured_params["api_key"] = litellm_params.api_key captured_params["api_version"] = litellm_params.api_version captured_params["credential_name"] = litellm_params.litellm_credential_name - captured_params["has_unexpected_field"] = hasattr( - litellm_params, "unexpected_field" - ) + captured_params["has_unexpected_field"] = hasattr(litellm_params, "unexpected_field") return ["gpt-4o"] monkeypatch.setattr(model_checks, "get_provider_models", fake_get_provider_models) @@ -448,9 +445,7 @@ def test_wildcard_custom_prefix_does_not_stack_provider_prefix(monkeypatch): result = get_known_models_from_wildcard( wildcard_model="ollama_server1/*", - litellm_params=LiteLLM_Params( - model="ollama_chat/*", custom_llm_provider="ollama_chat" - ), + litellm_params=LiteLLM_Params(model="ollama_chat/*", custom_llm_provider="ollama_chat"), ) assert result == ["ollama_server1/gemma3:1b", "ollama_server1/llama3:8b"] @@ -477,9 +472,7 @@ def test_wildcard_custom_prefix_keeps_org_segment_for_non_provider_first_segment result = get_known_models_from_wildcard( wildcard_model="my_hf/*", - litellm_params=LiteLLM_Params( - model="huggingface/*", custom_llm_provider="huggingface" - ), + litellm_params=LiteLLM_Params(model="huggingface/*", custom_llm_provider="huggingface"), ) assert result == ["my_hf/meta-llama/Llama-3-8B"] @@ -927,9 +920,7 @@ def test_add_known_models_refreshes_models_by_provider_for_wildcard_expansion(): assert fake_model not in litellm.models_by_provider["vertex_ai"] try: litellm.add_known_models( - model_cost_map={ - fake_model: {"litellm_provider": "vertex_ai-language-models", "mode": "chat"} - } + model_cost_map={fake_model: {"litellm_provider": "vertex_ai-language-models", "mode": "chat"}} ) assert fake_model in litellm.models_by_provider["vertex_ai"] assert litellm.models_by_provider is captured_reference @@ -1052,6 +1043,136 @@ def test_transcribe_is_a_known_provider_for_wildcard_expansion(): assert "transcribe" in litellm.models_by_provider assert "transcribe/StartTranscriptionJob" in litellm.models_by_provider["transcribe"] assert get_provider_models("transcribe") == ["transcribe/StartTranscriptionJob"] - assert get_known_models_from_wildcard("transcribe/*") == [ - "transcribe/StartTranscriptionJob" + assert get_known_models_from_wildcard("transcribe/*") == ["transcribe/StartTranscriptionJob"] + + +@respx.mock +def test_partial_bedrock_wildcard_filters_the_discovered_ids(monkeypatch: pytest.MonkeyPatch) -> None: + import litellm + from litellm.proxy.auth.model_checks import get_known_models_from_wildcard + from litellm.types.router import LiteLLM_Params + + monkeypatch.setattr(litellm, "check_provider_endpoint", True) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + region: Final = "ap-south-1" + respx.get( + f"https://bedrock.{region}.amazonaws.com/foundation-models", params={"byInferenceType": "ON_DEMAND"} + ).mock( + return_value=httpx.Response( + 200, + json={ + "modelSummaries": [ + {"modelId": "anthropic.claude-haiku-4-5-20251001-v1:0"}, + {"modelId": "amazon.nova-micro-v1:0"}, + ] + }, + ) + ) + respx.get( + f"https://bedrock.{region}.amazonaws.com/inference-profiles", params={"typeEquals": "SYSTEM_DEFINED"} + ).mock( + return_value=httpx.Response( + 200, + json={ + "inferenceProfileSummaries": [ + {"inferenceProfileId": "us.anthropic.claude-haiku-4-5-20251001-v1:0", "status": "ACTIVE"} + ] + }, + ) + ) + deployment: Final = LiteLLM_Params( + model="bedrock/anthropic.*", + aws_access_key_id="AKIAPARTIALWILDCARD", + aws_secret_access_key="partial-wildcard-secret", + aws_region_name=region, + ) + + assert get_known_models_from_wildcard("bedrock/anthropic.*", deployment) == [ + "bedrock/anthropic.claude-haiku-4-5-20251001-v1:0" + ] + + +@respx.mock +def test_partial_bedrock_wildcard_lists_nothing_when_no_discovered_id_carries_its_prefix( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm + from litellm.proxy.auth.model_checks import get_known_models_from_wildcard + from litellm.types.router import LiteLLM_Params + + monkeypatch.setattr(litellm, "check_provider_endpoint", True) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + region: Final = "ap-southeast-2" + respx.get( + f"https://bedrock.{region}.amazonaws.com/foundation-models", params={"byInferenceType": "ON_DEMAND"} + ).mock(return_value=httpx.Response(200, json={"modelSummaries": [{"modelId": "amazon.nova-micro-v1:0"}]})) + respx.get( + f"https://bedrock.{region}.amazonaws.com/inference-profiles", params={"typeEquals": "SYSTEM_DEFINED"} + ).mock( + return_value=httpx.Response( + 200, + json={ + "inferenceProfileSummaries": [ + {"inferenceProfileId": "us.anthropic.claude-haiku-4-5-20251001-v1:0", "status": "ACTIVE"} + ] + }, + ) + ) + deployment: Final = LiteLLM_Params( + model="bedrock/anthropic.*", + aws_access_key_id="AKIAPROFILEONLYACCOUNT", + aws_secret_access_key="profile-only-secret", + aws_region_name=region, + ) + + assert get_known_models_from_wildcard("bedrock/anthropic.*", deployment) == [] + + +@respx.mock +def test_custom_prefix_that_starts_a_discovered_id_still_prefixes_every_listed_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import litellm + from litellm.proxy.auth.model_checks import get_known_models_from_wildcard + from litellm.types.router import LiteLLM_Params + + monkeypatch.setattr(litellm, "check_provider_endpoint", True) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + region: Final = "ca-central-1" + respx.get( + f"https://bedrock.{region}.amazonaws.com/foundation-models", params={"byInferenceType": "ON_DEMAND"} + ).mock( + return_value=httpx.Response( + 200, + json={ + "modelSummaries": [ + {"modelId": "anthropic.claude-haiku-4-5-20251001-v1:0"}, + {"modelId": "amazon.nova-micro-v1:0"}, + ] + }, + ) + ) + respx.get( + f"https://bedrock.{region}.amazonaws.com/inference-profiles", params={"typeEquals": "SYSTEM_DEFINED"} + ).mock( + return_value=httpx.Response( + 200, + json={ + "inferenceProfileSummaries": [ + {"inferenceProfileId": "us.anthropic.claude-haiku-4-5-20251001-v1:0", "status": "ACTIVE"} + ] + }, + ) + ) + deployment: Final = LiteLLM_Params( + model="bedrock/*", + aws_access_key_id="AKIACUSTOMPREFIXACCOUNT", + aws_secret_access_key="custom-prefix-secret", + aws_region_name=region, + ) + + assert get_known_models_from_wildcard("anthropic/*", deployment) == [ + "anthropic/amazon.nova-micro-v1:0", + "anthropic/anthropic.claude-haiku-4-5-20251001-v1:0", + "anthropic/us.anthropic.claude-haiku-4-5-20251001-v1:0", ] diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 63e2280dc86..bf3eb7d9885 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -2553,6 +2553,54 @@ class TestGetValidModelsWithCLI: assert headers.get("Authorization") == "Bearer sk-test-cli-key-123" +@respx.mock +def test_get_valid_models_bedrock_lists_what_the_deployment_credentials_can_invoke( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + region: Final = "eu-west-3" + deployment: Final = LiteLLM_Params( + model="bedrock/*", + aws_access_key_id="AKIAIOSFODNN7EXAMPLE", + aws_secret_access_key="wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + aws_region_name=region, + ) + expected_scope: Final = f"/{region}/bedrock/aws4_request" + + def signed_by_the_deployment(request: httpx.Request, body: Mapping[str, object]) -> httpx.Response: + authorization: Final = request.headers.get("authorization", "") + if not authorization.startswith("AWS4-HMAC-SHA256 Credential=AKIAIOSFODNN7EXAMPLE/"): + return httpx.Response(403, json={"message": "not signed with the deployment's key"}) + if expected_scope not in authorization: + return httpx.Response(403, json={"message": "not signed for the deployment's region"}) + return httpx.Response(200, json=body) + + respx.get( + f"https://bedrock.{region}.amazonaws.com/foundation-models", params={"byInferenceType": "ON_DEMAND"} + ).mock( + side_effect=lambda request: signed_by_the_deployment( + request, {"modelSummaries": [{"modelId": "amazon.nova-micro-v1:0"}]} + ) + ) + respx.get( + f"https://bedrock.{region}.amazonaws.com/inference-profiles", + params={"typeEquals": "SYSTEM_DEFINED"}, + ).mock( + side_effect=lambda request: signed_by_the_deployment( + request, + { + "inferenceProfileSummaries": [ + {"inferenceProfileId": "eu.anthropic.claude-sonnet-4-5-20250929-v1:0", "status": "ACTIVE"} + ] + }, + ) + ) + + assert litellm.get_valid_models( + check_provider_endpoint=True, custom_llm_provider="bedrock", litellm_params=deployment + ) == ["amazon.nova-micro-v1:0", "eu.anthropic.claude-sonnet-4-5-20250929-v1:0"] + + class TestIsCachedMessage: """Test is_cached_message function for context caching detection.