fix(bedrock): list the account's invocable models behind bedrock/* when check_provider_endpoint is on (#43981)

* fix(bedrock): list the account's invocable models behind bedrock/* when check_provider_endpoint is on

* fix(bedrock): sign the listing query RFC 3986 style and raise BedrockError on a failed listing

* test(utils): keep test_utils.py at main's formatting, only the Bedrock discovery test is new

* fix(bedrock): chain BedrockModelInfo's constructor so the passthrough config keeps its event parser

* fix(bedrock): list discovered ids without the provider prefix so partial wildcards filter

get_valid_models returns bare vendor ids for every other provider and for the static bedrock catalog, and the proxy's wildcard expansion adds the provider prefix itself. The lister prefixed its ids, so a partial wildcard like bedrock/anthropic.* never matched the proxy's filter and listed uncallable bedrock/anthropic.bedrock/<id> entries

* test(bedrock): integration cells for wildcard discovery through the proxy and the SDK

* fix(proxy): keep a partial wildcard a filter when its deployment repeats the prefix

The wildcard expansion guessed filter-or-alias by whether any provider id carried the prefix. A deployment such as model_name bedrock/anthropic.* over model bedrock/anthropic.* can only ever route names that share the prefix, so when the account has no on-demand anthropic.* id (every Anthropic model behind an inference profile) the guess fell into the alias branch and listed bedrock/anthropic.<every id>, none of them callable. A deployment whose model repeats the suffix now always filters, and lists nothing when nothing matches

* fix(proxy): prefix every discovered id under a custom wildcard prefix that starts a vendor id

* test(bedrock): Final and read-only annotations in the wildcard discovery tests

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-09 14:14:17 -07:00 • committed by GitHub
parent 8080601a95
commit 43f07d0bd8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 2178 additions and 29 deletions

View file

@ -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]:
# """

View file

@ -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()})

View file

@ -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

View file

@ -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<key>[^/]+)/(?P<scope>[^,]+), "
r"SignedHeaders=(?P<signed>[^,]+), Signature=(?P<signature>[0-9a-f]{64})$"
)
_BEARER: Final = re.compile(r"^Bearer (?P<token>.+)$")
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.<region>.<suffix>`` 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}"

View file

@ -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
)

View file

@ -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

View file

@ -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"

View file

@ -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()

View file

@ -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"<html>upstream maintenance</html>", 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))

View file

@ -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

View file

@ -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())

View file

@ -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",
]

View file

@ -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.