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.