feat(e2e): key the provider cache per test and mount Bedrock behind it

The exact-request cache reused 5% of routed traffic (build 218: 19 hits,
350 misses) because every test salts its prompt with a fresh unique_marker(),
so the same test could never match itself across builds. It also routed only
openai and anthropic, while the week's flakiness was Bedrock.

Key is now HMAC(test id + method + URL + headers + body, with every
unique_marker() token replaced by a placeholder, + FIFO slot index). The slot
index is what keeps two marker-only-different calls in one test on two
recordings and therefore two provider response ids, so spend rows still
reconcile one per invocation. A call outside any test is not cacheable.

Bedrock gets a region-qualified mount and SigV4 re-signing, since the edge
rewrites the Host the proxy signed. Signature headers are excluded from the
key for signing mounts only, because x-amz-date would otherwise make every
Bedrock request a permanent miss; every other mount still keys on its
credentials whole. Only Anthropic-on-Bedrock chat deployments route:
embeddings, image generation, rerank and realtime keep their direct path, and
so do deployments carrying their own aws_role_name or static keys, whose whole
point is to prove the product's assume-role chain rather than the runner's.
The two eventstream actions bypass the cache and go live, still signed.

Counters are now attributed per mount as well as in total, so a build can
report a per-provider hit rate instead of one number.
This commit is contained in:
Yuneng Jiang 2026-09-16 02:15:46 -07:00
parent a8979fe054
commit 2d40254b57
8 changed files with 768 additions and 118 deletions

View file

@ -7,7 +7,7 @@ import subprocess
import threading
import time
import uuid
from collections.abc import Generator
from collections.abc import Generator, Mapping
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from dataclasses import dataclass, replace
@ -18,21 +18,58 @@ from urllib.parse import urlsplit
import pytest
from e2e_http import NetworkError, PreparedForward, RawResponse, StreamChunk, StreamHead, forward, prepare_forward
from models import LiteLLMParamsBody
from provider_cache import CacheEdge, CacheHit, CaptureLease, exact_key, successful_response
from models import LiteLLMParamsBody, ModelMode
from botocore.credentials import Credentials
from provider_cache import (
SIGNATURE_HEADERS,
CacheEdge,
CacheHit,
CaptureLease,
ResponseStore,
cacheable_endpoint,
request_identity,
slotted_key,
successful_response,
)
from provider_cache_redis import PUBLISH, RedisCommands, RedisResponseStore, configured_cache, redis_store
from provider_cache_routing import LIVE_PROVIDER_REQUIRED, route_cache_model
from provider_edge import configured_cache_backend, start_provider_edge
from fixture_mode import SESSION_TEST_KEY
from provider_edge import EDGE_MOUNTS, configured_cache_backend, resolve_mount, start_provider_edge
from provider_edge_bedrock import bedrock_signer
from redis.exceptions import ConnectionError as RedisConnectionError
SECRET: Final = b"synthetic-cache-hmac-key-for-tests"
BODY: Final = b'{"model":"test","messages":[{"role":"user","content":"hello"}]}'
SUCCESS: Final = b'{"id":"provider-fixed-id","choices":[{"message":{"content":"hello"},"finish_reason":"stop"}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}'
HEADERS: Final = {"content-type": "application/json", "authorization": "Bearer synthetic-account-one"}
TEST_KEY: Final = "tests/e2e/synthetic_suite.py::TestCase::test_case"
OTHER_TEST_KEY: Final = "tests/e2e/synthetic_suite.py::TestCase::test_other_case"
def marked(marker: str) -> bytes:
"""One request body shaped like the suite's own: a fixed prompt salted with a
12-lowercase-hex ``unique_marker()`` token, fresh on every run."""
return b'{"model":"test","messages":[{"role":"user","content":"hello %s"}]}' % marker.encode()
MARKED: Final = marked("0a1b2c3d4e5f")
BEDROCK_MOUNT: Final = "bedrock/us-east-1"
BEDROCK_MODEL: Final = "us.anthropic.claude-haiku-4-5-20251001-v1%3A0"
BEDROCK_BODY: Final = b'{"messages":[{"role":"user","content":[{"text":"hello 0a1b2c3d4e5f"}]}]}'
CONVERSE_SUCCESS: Final = (
b'{"output":{"message":{"role":"assistant","content":[{"text":"hi"}]}},'
b'"stopReason":"end_turn","usage":{"inputTokens":1,"outputTokens":1,"totalTokens":2}}'
)
INVOKE_SUCCESS: Final = (
b'{"id":"msg_synthetic","type":"message","role":"assistant",'
b'"content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}'
)
STATIC_CREDENTIALS: Final = Credentials("AKIAIOSFODNN7EXAMPLE", "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY")
class Provider(ThreadingHTTPServer):
hits: tuple[tuple[str, bytes], ...] = ()
authorizations: tuple[str, ...] = ()
response: bytes = SUCCESS
status: int = 200
delay: float = 0
@ -49,6 +86,7 @@ class Handler(BaseHTTPRequestHandler):
assert isinstance(server, Provider)
body: Final = self.rfile.read(int(self.headers.get("content-length", "0")))
server.hits += ((self.path, body),)
server.authorizations += (self.headers.get("authorization", ""),)
time.sleep(server.delay)
self.send_response(server.status)
if server.stream:
@ -122,6 +160,29 @@ def store(redis_url: str) -> RedisResponseStore:
return redis_store(redis_url, "test-" + uuid.uuid4().hex)
def cache_edge(store: ResponseStore, test_key: str = TEST_KEY) -> CacheEdge:
"""A cache edge standing in for one pytest process. A fresh instance over the
same store is the next build running the same test: the recordings survive,
the per-test FIFO slot counters start over."""
return CacheEdge(store, SECRET, test_key=lambda: test_key)
def slot_key(
url: str, slot: int = 0, body: bytes | None = BODY,
headers: dict[str, str] = HEADERS, test_key: str = TEST_KEY,
) -> str:
prepared: Final = prepare_forward("POST", url, headers, body)
assert isinstance(prepared, PreparedForward)
return slotted_key(SECRET, request_identity(SECRET, test_key, "POST", url, prepared.headers, body), slot)
def bedrock_cache_edge(store: ResponseStore, test_key: str = TEST_KEY) -> CacheEdge:
return CacheEdge(
store, SECRET, test_key=lambda: test_key,
signers={BEDROCK_MOUNT: bedrock_signer("us-east-1", lambda: STATIC_CREDENTIALS)},
)
@contextmanager
def edge(cache: CacheEdge, provider: Provider) -> Generator[str, None, None]:
upstream: Final = f"http://127.0.0.1:{provider.server_port}"
@ -132,34 +193,54 @@ def edge(cache: CacheEdge, provider: Provider) -> Generator[str, None, None]:
running.shutdown()
@contextmanager
def bedrock_edge(cache: CacheEdge, provider: Provider, action: str = "converse") -> Generator[str, None, None]:
upstream: Final = f"http://127.0.0.1:{provider.server_port}"
running: Final = start_provider_edge(cache, mounts={BEDROCK_MOUNT: upstream})
try:
yield f"{running.edge.api_base(BEDROCK_MOUNT)}/model/{BEDROCK_MODEL}/{action}"
finally:
running.shutdown()
def call(url: str, body: bytes = BODY, headers: dict[str, str] = HEADERS) -> RawResponse:
result: Final = forward("POST", url, headers=headers, body=body, timeout=5)
assert isinstance(result, RawResponse), result
return result
def test_success_is_reusable_across_fresh_edges(store: RedisResponseStore, provider: Provider) -> None:
with edge(CacheEdge(store, SECRET), provider) as url:
def test_repeated_call_takes_its_own_slot_and_both_replay_next_run(
store: RedisResponseStore, provider: Provider,
) -> None:
with edge(cache_edge(store), provider) as url:
assert call(url).body == SUCCESS
assert call(url).body == SUCCESS
with edge(CacheEdge(store, SECRET), provider) as other:
assert len(provider.hits) == 2
with edge(cache_edge(store), provider) as other:
assert call(other).body == SUCCESS
assert len(provider.hits) == 1
assert call(other).body == SUCCESS
assert len(provider.hits) == 2
@pytest.mark.parametrize("body", [BODY + b" ", BODY.replace(b"hello", b"Hello"), BODY.replace(b"test", b"test2")])
def test_any_body_change_calls_live(store: RedisResponseStore, provider: Provider, body: bytes) -> None:
with edge(CacheEdge(store, SECRET), provider) as url:
with edge(cache_edge(store), provider) as url:
call(url)
assert len(provider.hits) == 1
with edge(cache_edge(store), provider) as url:
call(url, body)
assert len(provider.hits) == 2
with edge(cache_edge(store), provider) as url:
call(url, body)
assert len(provider.hits) == 2
@pytest.mark.parametrize("name,value", [("authorization", "Bearer another-account"), ("x-request-id", "one"), ("anthropic-version", "new")])
def test_changed_header_cannot_reuse(store: RedisResponseStore, provider: Provider, name: str, value: str) -> None:
with edge(CacheEdge(store, SECRET), provider) as url:
with edge(cache_edge(store), provider) as url:
call(url)
assert len(provider.hits) == 1
with edge(cache_edge(store), provider) as url:
call(url, headers=HEADERS | {name: value})
call(url + "?x=1")
assert len(provider.hits) == 3
@ -169,36 +250,59 @@ def test_changed_header_cannot_reuse(store: RedisResponseStore, provider: Provid
def test_failed_provider_responses_never_enter_cache(store: RedisResponseStore, provider: Provider, status: int, response: bytes) -> None:
provider.status = status
provider.response = response
with edge(CacheEdge(store, SECRET), provider) as url:
with edge(cache_edge(store), provider) as url:
assert call(url).status_code == status
assert len(provider.hits) == 1
with edge(cache_edge(store), provider) as url:
assert call(url).body == response
assert len(provider.hits) == 2
def test_cookie_setting_success_is_reused_without_the_cookie(store: RedisResponseStore, provider: Provider) -> None:
provider.cookie = "__cf_bm=synthetic-bot-management; Path=/; HttpOnly; Secure"
with edge(CacheEdge(store, SECRET), provider) as url:
replies: Final = tuple(call(url) for _ in range(2))
with edge(cache_edge(store), provider) as url:
live: Final = call(url)
with edge(cache_edge(store), provider) as url:
replayed: Final = call(url)
assert len(provider.hits) == 1
assert all(reply.body == SUCCESS and "set-cookie" not in reply.headers for reply in replies)
assert all(reply.body == SUCCESS and "set-cookie" not in reply.headers for reply in (live, replayed))
def test_expiry_does_not_slide(store: RedisResponseStore, provider: Provider) -> None:
short: Final = replace(store, lifetime_ms=250)
with edge(CacheEdge(short, SECRET), provider) as url:
call(url)
call(url)
time.sleep(0.3)
call(url)
call(url)
url: Final = f"http://127.0.0.1:{provider.server_port}/v1/chat/completions"
def drain() -> None:
head = cache_edge(short).forward("openai", "POST", url, dict(HEADERS), BODY, 5)
assert isinstance(head, StreamHead)
assert b"".join(step.data for step in head.steps if isinstance(step, StreamChunk)) == SUCCESS
drain()
assert len(provider.hits) == 1
drain()
assert len(provider.hits) == 1
time.sleep(0.3)
drain()
assert len(provider.hits) == 2
def test_concurrent_requests_publish_atomically(store: RedisResponseStore, provider: Provider) -> None:
def test_concurrent_builds_publish_one_recording_atomically(
store: RedisResponseStore, provider: Provider,
) -> None:
"""Five processes running the same test at the same time all reach slot 0 of
one key, which is the only way the capture lease is contended now that a
repeat inside a single test takes its own slot."""
provider.delay = 0.15
with edge(CacheEdge(store, SECRET), provider) as url:
with ThreadPoolExecutor(max_workers=5) as executor:
replies: Final = tuple(executor.map(lambda _: call(url).body, range(5)))
url: Final = f"http://127.0.0.1:{provider.server_port}/v1/chat/completions"
edges: Final = tuple(cache_edge(store) for _ in range(5))
def drain(cache: CacheEdge) -> bytes:
head = cache.forward("openai", "POST", url, dict(HEADERS), BODY, 5)
assert isinstance(head, StreamHead)
return b"".join(step.data for step in head.steps if isinstance(step, StreamChunk))
with ThreadPoolExecutor(max_workers=5) as executor:
replies: Final = tuple(executor.map(drain, edges))
assert replies == (SUCCESS,) * 5
assert len(provider.hits) == 1
@ -231,9 +335,9 @@ def test_stream_completion_controls_publication(store: RedisResponseStore, provi
provider.stream = True
provider.truncated = truncated
provider.response = b'data: {"choices":[{"index":0,"delta":{"content":"hello"},"finish_reason":"stop"}]}\n\ndata: [DONE]\n\n'
with edge(CacheEdge(store, SECRET), provider) as url:
for _ in range(2):
result: Final = forward("POST", url, headers=HEADERS, body=BODY, timeout=5)
for _ in range(2):
with edge(cache_edge(store), provider) as url:
result = forward("POST", url, headers=HEADERS, body=BODY, timeout=5)
if truncated:
assert isinstance(result, NetworkError)
else:
@ -246,9 +350,9 @@ def test_store_outage_preserves_provider_success(provider: Provider) -> None:
probe.bind(("127.0.0.1", 0))
port: Final = probe.getsockname()[1]
unavailable: Final = redis_store(f"redis://127.0.0.1:{port}/0", "unavailable")
with edge(CacheEdge(unavailable, SECRET), provider) as url:
assert call(url).body == SUCCESS
assert call(url).body == SUCCESS
for _ in range(2):
with edge(cache_edge(unavailable), provider) as url:
assert call(url).body == SUCCESS
assert len(provider.hits) == 2
@ -267,7 +371,8 @@ def test_old_lease_cannot_overwrite_new_owner(store: RedisResponseStore) -> None
def test_identity_preserves_values_and_never_contains_credentials() -> None:
variants: Final = (b'{}', b'{"a":null}', b'{"a":false}', b'{"a":0}', b'{"a":0.0}', b'{"a":"0"}', b' { }', None, b'')
keys: Final = tuple(exact_key(SECRET, "POST", "https://example.invalid/v1/chat/completions", HEADERS, body) for body in variants)
url: Final = "https://example.invalid/v1/chat/completions"
keys: Final = tuple(request_identity(SECRET, TEST_KEY, "POST", url, HEADERS, body) for body in variants)
assert len(set(keys)) == len(variants)
assert all(len(key) == 64 and "synthetic-account" not in key for key in keys)
@ -275,21 +380,22 @@ def test_identity_preserves_values_and_never_contains_credentials() -> None:
@pytest.mark.parametrize("payload", [b"corrupt response", '{"response":"{}","signature":"é"}'.encode()])
def test_corrupt_entry_is_replaced_by_same_successful_request(store: RedisResponseStore, provider: Provider, payload: bytes) -> None:
upstream: Final = f"http://127.0.0.1:{provider.server_port}/v1/chat/completions"
prepared: Final = prepare_forward("POST", upstream, HEADERS, BODY)
assert isinstance(prepared, PreparedForward)
key: Final = exact_key(SECRET, "POST", upstream, prepared.headers, BODY)
key: Final = slot_key(upstream)
lease: Final = store.lookup(key)
assert isinstance(lease, CaptureLease)
assert store.publish(key, lease, payload)
cache: Final = CacheEdge(store, SECRET)
for _ in range(2):
head = cache.forward("POST", upstream, HEADERS, BODY, 5)
caches: Final = tuple(cache_edge(store) for _ in range(2))
for cache in caches:
head = cache.forward("openai", "POST", upstream, dict(HEADERS), BODY, 5)
assert isinstance(head, StreamHead)
assert b"".join(step.data for step in head.steps if isinstance(step, StreamChunk)) == SUCCESS
assert len(provider.hits) == 1
assert dict(cache.counters.counts) == {
"corrupt": 1, "misses": 1, "upstream_attempts": 1, "writes": 1, "hits": 1,
assert dict(caches[0].counters.counts) == {
"corrupt": 1, "mount:openai:corrupt": 1, "misses": 1, "mount:openai:misses": 1,
"upstream_attempts": 1, "mount:openai:upstream_attempts": 1,
"writes": 1, "mount:openai:writes": 1,
}
assert dict(caches[1].counters.counts) == {"hits": 1, "mount:openai:hits": 1}
@pytest.mark.parametrize("payload", [
@ -301,22 +407,242 @@ def test_corrupt_entry_is_replaced_by_same_successful_request(store: RedisRespon
def test_malformed_success_stream_is_never_cached(store: RedisResponseStore, provider: Provider, payload: bytes) -> None:
provider.stream = True
provider.response = payload
with edge(CacheEdge(store, SECRET), provider) as url:
assert call(url).body == payload
assert call(url).body == payload
for _ in range(2):
with edge(cache_edge(store), provider) as url:
assert call(url).body == payload
assert len(provider.hits) == 2
def test_requests_differing_only_by_marker_share_one_recording_per_slot(
store: RedisResponseStore, provider: Provider,
) -> None:
"""The whole point of the canonical key. Every e2e test salts its prompt with
a fresh ``unique_marker()``, so before this the same test could never reuse
anything across builds. The second run mints markers it has never sent, which
is what a later build actually does, and must still serve both from the two
slots the first run recorded."""
with edge(cache_edge(store), provider) as url:
assert call(url, MARKED).body == SUCCESS
assert call(url, marked("f5e4d3c2b1a0")).body == SUCCESS
assert len(provider.hits) == 2
with edge(cache_edge(store), provider) as url:
assert call(url, marked("7c6b5a493827")).body == SUCCESS
assert call(url, marked("1122334455ff")).body == SUCCESS
assert len(provider.hits) == 2
@pytest.mark.parametrize("body", [
b'{"model":"test","messages":[{"role":"user","content":"hello 0a1b2c3d4e5"}]}',
b'{"model":"test","messages":[{"role":"user","content":"hello 0a1b2c3d4e5f0"}]}',
b'{"model":"test","messages":[{"role":"user","content":"hello 0A1B2C3D4E5F"}]}',
b'{"model":"0a1b2c3d4e5f","messages":[{"role":"user","content":"hello"}]}',
])
def test_a_token_that_is_not_a_marker_keeps_its_own_key(
store: RedisResponseStore, provider: Provider, body: bytes,
) -> None:
"""Too short, too long, upper case, or in another field: none of these is the
12-lowercase-hex token ``unique_marker`` mints, so none may fold onto it."""
with edge(cache_edge(store), provider) as url:
call(url, MARKED)
assert len(provider.hits) == 1
with edge(cache_edge(store), provider) as url:
call(url, body)
assert len(provider.hits) == 2
def test_another_test_never_reuses_this_tests_recording(
store: RedisResponseStore, provider: Provider,
) -> None:
with edge(cache_edge(store), provider) as url:
call(url)
assert len(provider.hits) == 1
with edge(cache_edge(store, OTHER_TEST_KEY), provider) as url:
call(url)
assert len(provider.hits) == 2
with edge(cache_edge(store, OTHER_TEST_KEY), provider) as url:
call(url)
assert len(provider.hits) == 2
def test_calls_outside_any_test_are_never_cached(
store: RedisResponseStore, provider: Provider,
) -> None:
url: Final = f"http://127.0.0.1:{provider.server_port}/v1/chat/completions"
cache: Final = CacheEdge(store, SECRET, test_key=lambda: SESSION_TEST_KEY)
for _ in range(2):
head = cache.forward("openai", "POST", url, dict(HEADERS), BODY, 5)
assert isinstance(head, StreamHead)
assert b"".join(step.data for step in head.steps if isinstance(step, StreamChunk)) == SUCCESS
assert len(provider.hits) == 2
assert dict(cache.counters.counts) == {
"bypass": 2, "mount:openai:bypass": 2,
"upstream_attempts": 2, "mount:openai:upstream_attempts": 2,
}
def test_counters_attribute_every_outcome_to_its_mount(
store: RedisResponseStore, provider: Provider,
) -> None:
"""The build report needs per-provider hit counts, and the flat totals cannot
supply them. Anthropic is served a chat-shaped body here, which its validator
rejects, so one mount writes and the other does not."""
upstream: Final = f"http://127.0.0.1:{provider.server_port}"
cache: Final = cache_edge(store)
running: Final = start_provider_edge(cache, mounts={"openai": upstream, "anthropic": upstream})
try:
call(running.edge.api_base("openai") + "/v1/chat/completions")
call(running.edge.api_base("anthropic") + "/v1/messages")
finally:
running.shutdown()
counts: Final = dict(cache.counters.counts)
assert counts["misses"] == 2
assert counts["mount:openai:misses"] == 1 and counts["mount:anthropic:misses"] == 1
assert counts["mount:openai:writes"] == 1 and "mount:anthropic:writes" not in counts
assert counts["mount:anthropic:rejected"] == 1 and "mount:openai:rejected" not in counts
class TestBedrockSigning:
"""Bedrock is the reason the edge could not mount it before: SigV4 covers the
Host header, so forwarding through a rewritten api_base invalidates the
proxy's signature. The edge mints its own over the upstream URL instead."""
def test_the_proxys_signature_is_replaced_not_forwarded(self) -> None:
signer: Final = bedrock_signer("us-east-1", lambda: STATIC_CREDENTIALS)
signed: Final = signer(
"POST",
f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{BEDROCK_MODEL}/converse",
{"content-type": "application/json", "Authorization": "AWS4-HMAC-SHA256 Credential=PROXY/...",
"X-Amz-Date": "19700101T000000Z", "X-Amz-Security-Token": "proxy-session-token"},
BEDROCK_BODY,
)
assert "PROXY" not in str(signed) and "proxy-session-token" not in str(signed)
assert signed["Authorization"].startswith("AWS4-HMAC-SHA256 Credential=AKIAIOSFODNN7EXAMPLE/")
assert "/us-east-1/bedrock/aws4_request" in signed["Authorization"]
assert signed["X-Amz-Date"] != "19700101T000000Z"
assert signed["content-type"] == "application/json"
def test_the_signed_url_reaches_the_wire_byte_for_byte(self) -> None:
"""SigV4 hashes the canonical URI, so if the HTTP layer re-encoded the
colon in an inference-profile id after signing, every call would fail
with a signature mismatch rather than anything that names the cause."""
url: Final = f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{BEDROCK_MODEL}/converse"
signer: Final = bedrock_signer("us-east-1", lambda: STATIC_CREDENTIALS)
prepared: Final = prepare_forward("POST", url, signer("POST", url, dict(HEADERS), BEDROCK_BODY), BEDROCK_BODY)
assert isinstance(prepared, PreparedForward)
assert urlsplit(prepared.url).path == urlsplit(url).path
def test_signature_headers_are_excluded_from_the_key(
self, store: RedisResponseStore, provider: Provider,
) -> None:
"""A real signature is fresh on every call, so keying on it would make
every Bedrock request a permanent miss. The stub signer here varies its
stamp per call on purpose: the real one only varies once a second, which
would let this pass by luck when it should fail."""
provider.response = CONVERSE_SUCCESS
stamps: Final = iter(("20260101T000000Z", "20260102T111111Z"))
def varying(method: str, url: str, headers: Mapping[str, str], body: bytes | None) -> dict[str, str]:
return dict(headers) | {"authorization": f"AWS4-HMAC-SHA256 {url}", "x-amz-date": next(stamps)}
def signing_edge() -> CacheEdge:
return CacheEdge(store, SECRET, test_key=lambda: TEST_KEY, signers={BEDROCK_MOUNT: varying})
for _ in range(2):
with bedrock_edge(signing_edge(), provider) as url:
assert call(url, BEDROCK_BODY).body == CONVERSE_SUCCESS
assert len(provider.hits) == 1
assert provider.authorizations[0] == (
f"AWS4-HMAC-SHA256 http://127.0.0.1:{provider.server_port}/model/{BEDROCK_MODEL}/converse"
), "the signature must cover the upstream URL the edge calls, not the edge URL the proxy called"
def test_a_mount_without_a_signer_still_keys_on_its_credentials(
self, store: RedisResponseStore, provider: Provider,
) -> None:
"""The exclusion is per mount. Dropping authorization globally would let
one OpenAI account read another's recording."""
cache: Final = bedrock_cache_edge(store)
assert "authorization" in SIGNATURE_HEADERS
assert "authorization" in cache.keyed("openai", HEADERS)
assert "authorization" not in cache.keyed(BEDROCK_MOUNT, HEADERS)
with edge(cache, provider) as url:
call(url)
with edge(bedrock_cache_edge(store), provider) as url:
call(url, headers=HEADERS | {"authorization": "Bearer synthetic-account-two"})
assert len(provider.hits) == 2
@pytest.mark.parametrize("action,response", [("converse", CONVERSE_SUCCESS), ("invoke", INVOKE_SUCCESS)])
def test_complete_responses_replay_on_the_next_run(
self, store: RedisResponseStore, provider: Provider, action: str, response: bytes,
) -> None:
provider.response = response
for _ in range(2):
with bedrock_edge(bedrock_cache_edge(store), provider, action) as url:
assert call(url, BEDROCK_BODY).body == response
assert len(provider.hits) == 1
@pytest.mark.parametrize("action,response", [
("converse", b'{"output":{"message":{}}}'),
("converse", b'{"stopReason":"end_turn"}'),
("converse", b'{"message":"The provided model identifier is invalid."}'),
("converse", CONVERSE_SUCCESS[:-20]),
("invoke", b'{"id":"msg_x","type":"message","content":[{"type":"text","text":"hi"}]}'),
("invoke", b'{"id":"msg_x","type":"message","stop_reason":"end_turn"}'),
("invoke", b'{"message":"Too many requests, please wait before trying again."}'),
])
def test_incomplete_or_error_bodies_never_enter_the_cache(
self, store: RedisResponseStore, provider: Provider, action: str, response: bytes,
) -> None:
provider.response = response
for _ in range(2):
with bedrock_edge(bedrock_cache_edge(store), provider, action) as url:
assert call(url, BEDROCK_BODY).body == response
assert len(provider.hits) == 2
@pytest.mark.parametrize("action", ["converse-stream", "invoke-with-response-stream"])
def test_streaming_endpoints_go_live_every_time(
self, store: RedisResponseStore, provider: Provider, action: str,
) -> None:
"""An eventstream's completeness cannot be proven without parsing its
frames, so these bypass rather than risk recording a truncated answer.
They are still signed: a bypass is a forward, not a passthrough."""
provider.response = CONVERSE_SUCCESS
cache: Final = bedrock_cache_edge(store)
for _ in range(2):
with bedrock_edge(cache, provider, action) as url:
assert call(url, BEDROCK_BODY).body == CONVERSE_SUCCESS
assert len(provider.hits) == 2
assert dict(cache.counters.counts)[f"mount:{BEDROCK_MOUNT}:bypass"] == 2
assert all(
sent.startswith("AWS4-HMAC-SHA256 Credential=AKIAIOSFODNN7EXAMPLE/")
for sent in provider.authorizations
), provider.authorizations
@pytest.mark.parametrize("action,cacheable", [
("converse", True), ("invoke", True),
("converse-stream", False), ("invoke-with-response-stream", False),
])
def test_only_the_unary_bedrock_actions_are_cacheable(self, action: str, cacheable: bool) -> None:
url: Final = f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{BEDROCK_MODEL}/{action}"
assert cacheable_endpoint(BEDROCK_MOUNT, "POST", url, BEDROCK_BODY) is cacheable
def test_a_region_mount_resolves_whole(self) -> None:
resolved: Final = resolve_mount(f"/{BEDROCK_MOUNT}/model/{BEDROCK_MODEL}/converse", EDGE_MOUNTS)
assert resolved is not None
assert resolved.mount == BEDROCK_MOUNT
assert resolved.upstream_base == "https://bedrock-runtime.us-east-1.amazonaws.com"
assert resolved.upstream_path == f"model/{BEDROCK_MODEL}/converse"
def test_anthropic_stream_requires_start_finish_and_stop() -> None:
start: Final = b'data: {"type":"message_start","message":{}}\n\n'
finish: Final = b'data: {"type":"message_delta","delta":{"stop_reason":"end_turn"}}\n\n'
stop: Final = b'data: {"type":"message_stop"}\n\n'
url: Final = "https://example.invalid/v1/messages"
headers: Final = {"content-type": "text/event-stream"}
assert successful_response(url, 200, headers, start + finish + stop)
assert not successful_response(url, 200, headers, start + stop)
assert not successful_response(url, 200, headers, finish + stop)
assert not successful_response(url, 200, headers, start + finish)
assert successful_response("anthropic", url, 200, headers, start + finish + stop)
assert not successful_response("anthropic", url, 200, headers, start + stop)
assert not successful_response("anthropic", url, 200, headers, finish + stop)
assert not successful_response("anthropic", url, 200, headers, start + finish)
@pytest.mark.parametrize("provider,suffix", [("openai", "/v1"), ("anthropic", "")])
@ -342,6 +668,53 @@ def test_registration_preserves_unsupported_or_explicit_routes(params: LiteLLMPa
assert route_cache_model(params, unexpected_edge, enabled=True) is params
@pytest.mark.parametrize("model", [
"bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"bedrock/converse/us.anthropic.claude-sonnet-5",
"bedrock/invoke/us.anthropic.claude-haiku-4-5-20251001-v1:0",
])
def test_anthropic_on_bedrock_registers_the_edge_as_its_runtime_endpoint(model: str) -> None:
params: Final = LiteLLMParamsBody(model=model)
routed: Final = route_cache_model(params, lambda mount: f"http://edge.invalid/{mount}", enabled=True)
assert routed.aws_bedrock_runtime_endpoint == "http://edge.invalid/bedrock/us-east-1"
assert routed.api_base is None
assert routed.model_dump(exclude={"aws_bedrock_runtime_endpoint"}) == params.model_dump(
exclude={"aws_bedrock_runtime_endpoint"}
)
@pytest.mark.parametrize("params", [
LiteLLMParamsBody(model="bedrock/amazon.titan-embed-text-v2:0"),
LiteLLMParamsBody(model="bedrock/amazon.nova-canvas-v1:0"),
LiteLLMParamsBody(model="bedrock/amazon.nova-sonic-v1:0"),
LiteLLMParamsBody(model="bedrock/arn:aws:bedrock:us-east-1::foundation-model/cohere.rerank-v3-5:0"),
LiteLLMParamsBody(model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", aws_role_name="arn:aws:iam::1:role/x"),
LiteLLMParamsBody(model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", aws_access_key_id="AKIA"),
LiteLLMParamsBody(model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", api_base="https://custom.invalid"),
LiteLLMParamsBody(
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
aws_bedrock_runtime_endpoint="https://custom.invalid",
),
LiteLLMParamsBody(model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", aws_region_name="eu-west-1"),
])
def test_bedrock_deployments_the_edge_must_not_touch_keep_their_direct_route(params: LiteLLMParamsBody) -> None:
"""Non-Anthropic models the runner role cannot invoke, deployments carrying
their own AWS identity (routing those would replace the assume-role chain the
batch suite exists to prove), explicit endpoints, and unmounted regions."""
routed: Final = route_cache_model(
params, lambda mount: None if mount not in EDGE_MOUNTS else f"http://edge.invalid/{mount}", enabled=True,
)
assert routed is params or routed.aws_bedrock_runtime_endpoint == params.aws_bedrock_runtime_endpoint
@pytest.mark.parametrize("mode", ["batch", "realtime", "image_generation"])
def test_a_bedrock_deployment_with_a_mode_keeps_its_direct_route(mode: ModelMode) -> None:
params: Final = LiteLLMParamsBody(model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0")
assert route_cache_model(
params, lambda mount: f"http://edge.invalid/{mount}", enabled=True, mode=mode,
) is params
def test_rollback_and_live_only_policy_keep_direct_provider_route() -> None:
params: Final = LiteLLMParamsBody(model="openai/test")
assert route_cache_model(params, lambda _: "http://edge.invalid", enabled=False) is params
@ -366,14 +739,16 @@ class PublishOutage:
def test_write_outage_preserves_success_without_hidden_retry(store: RedisResponseStore, provider: Provider) -> None:
unavailable: Final = replace(store, client=PublishOutage(store.client))
cache: Final = CacheEdge(unavailable, SECRET)
cache: Final = cache_edge(unavailable)
with edge(cache, provider) as url:
assert call(url).body == SUCCESS
assert call(url).body == SUCCESS
assert len(provider.hits) == 2
assert dict(cache.counters.counts)["write_failures"] == 2
with edge(CacheEdge(store, SECRET), provider) as url:
with edge(cache_edge(store), provider) as url:
assert call(url).body == SUCCESS
assert len(provider.hits) == 3
with edge(cache_edge(store), provider) as url:
assert call(url).body == SUCCESS
assert len(provider.hits) == 3
@ -382,45 +757,41 @@ def test_connection_failure_releases_capture_lease(store: RedisResponseStore) ->
with socket.socket() as unavailable:
unavailable.bind(("127.0.0.1", 0))
url: Final = f"http://127.0.0.1:{unavailable.getsockname()[1]}/v1/chat/completions"
cache: Final = CacheEdge(store, SECRET)
assert isinstance(cache.forward("POST", url, HEADERS, BODY, 0.2), NetworkError)
prepared: Final = prepare_forward("POST", url, HEADERS, BODY)
assert isinstance(prepared, PreparedForward)
key: Final = exact_key(SECRET, "POST", url, prepared.headers, BODY)
slot: Final = store.lookup(key)
assert isinstance(slot, CaptureLease)
assert store.release(key, slot)
cache: Final = cache_edge(store)
assert isinstance(cache.forward("openai", "POST", url, dict(HEADERS), BODY, 0.2), NetworkError)
key: Final = slot_key(url)
lease: Final = store.lookup(key)
assert isinstance(lease, CaptureLease)
assert store.release(key, lease)
assert dict(cache.counters.counts)["rejected"] == 1
def test_close_before_first_chunk_releases_lease(store: RedisResponseStore, provider: Provider) -> None:
url: Final = f"http://127.0.0.1:{provider.server_port}/v1/chat/completions"
cache: Final = CacheEdge(store, SECRET)
head: Final = cache.forward("POST", url, HEADERS, BODY, 5)
cache: Final = cache_edge(store)
head: Final = cache.forward("openai", "POST", url, dict(HEADERS), BODY, 5)
assert isinstance(head, StreamHead)
head.steps.close()
prepared: Final = prepare_forward("POST", url, HEADERS, BODY)
assert isinstance(prepared, PreparedForward)
key: Final = exact_key(SECRET, "POST", url, prepared.headers, BODY)
slot: Final = store.lookup(key)
assert isinstance(slot, CaptureLease)
assert store.release(key, slot)
key: Final = slot_key(url)
lease: Final = store.lookup(key)
assert isinstance(lease, CaptureLease)
assert store.release(key, lease)
def test_effective_account_change_cannot_reuse_cache(
store: RedisResponseStore, provider: Provider, monkeypatch: pytest.MonkeyPatch, tmp_path,
) -> None:
url: Final = f"http://127.0.0.1:{provider.server_port}/v1/chat/completions"
cache: Final = CacheEdge(store, SECRET)
for account in ("account-a", "account-b", "account-b"):
caches: Final = tuple(cache_edge(store) for _ in range(3))
for account, cache in zip(("account-a", "account-b", "account-b"), caches, strict=True):
netrc = tmp_path / account
netrc.write_text(f"machine 127.0.0.1 login {account} password synthetic\n")
monkeypatch.setenv("NETRC", str(netrc))
head = cache.forward("POST", url, HEADERS, BODY, 5)
head = cache.forward("openai", "POST", url, dict(HEADERS), BODY, 5)
assert isinstance(head, StreamHead)
assert b"".join(step.data for step in head.steps if isinstance(step, StreamChunk)) == SUCCESS
assert len(provider.hits) == 2
assert dict(cache.counters.counts)["hits"] == 1
assert dict(caches[2].counters.counts)["hits"] == 1
def test_enabled_environment_reuses_store_across_fresh_backends(
@ -449,7 +820,7 @@ def test_enabled_environment_reuses_store_across_fresh_backends(
def test_duplicate_headers_bypass_cache_and_count_live_calls(
store: RedisResponseStore, provider: Provider, known_mount: bool,
) -> None:
cache: Final = CacheEdge(store, SECRET)
cache: Final = cache_edge(store)
with edge(cache, provider) as url:
parsed: Final = urlsplit(url)
for _ in range(2):

View file

@ -51,6 +51,9 @@ SECRET_FIELD_SUFFIXES: Final[tuple[str, ...]] = (
)
SECRET_PLACEHOLDER: Final = "<secret>"
MARKER_PATTERN: Final = re.compile(r"(?<![0-9a-fA-F])[0-9a-f]{12}(?![0-9a-fA-F])")
MARKER_PLACEHOLDER: Final = "<marker>"
PLACEHOLDER_RULES: Final[tuple[tuple[re.Pattern[str], str], ...]] = (
(re.compile(r"(?<![0-9a-fA-F])[0-9a-f]{64}(?![0-9a-fA-F])"), "<sha256>"),
(
@ -67,7 +70,7 @@ PLACEHOLDER_RULES: Final[tuple[tuple[re.Pattern[str], str], ...]] = (
re.compile(r"\b(?:chatcmpl|msgbatch|msg|resp|batch|call|req|ftjob|gen|file)[-_][A-Za-z0-9]{8,}\b"),
"<id>",
),
(re.compile(r"(?<![0-9a-fA-F])[0-9a-f]{12}(?![0-9a-fA-F])"), "<marker>"),
(MARKER_PATTERN, MARKER_PLACEHOLDER),
)

View file

@ -951,6 +951,7 @@ class LiteLLMParamsBody(BaseModel):
aws_access_key_id: str | None = None
aws_secret_access_key: str | None = None
aws_region_name: str | None = None
aws_bedrock_runtime_endpoint: str | None = None
vertex_project: str | None = None
vertex_location: str | None = None
vertex_credentials: str | None = None

View file

@ -23,12 +23,18 @@ from e2e_http import (
prepare_forward,
primed_steps,
)
from fixture_canonical import MARKER_PATTERN, MARKER_PLACEHOLDER
from fixture_mode import SESSION_TEST_KEY, current_test_key
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
LIFETIME_SECONDS: Final = 86_400
MAX_REQUEST_BYTES: Final = 256 * 1024
MAX_RESPONSE_BYTES: Final = 8 * 1024 * 1024
UNRECORDED_RESPONSE_HEADERS: Final = frozenset({"set-cookie"})
SIGNATURE_HEADERS: Final = frozenset(
{"authorization", "x-amz-date", "x-amz-security-token", "x-amz-content-sha256"}
)
BEDROCK_MOUNT_PREFIX: Final = "bedrock"
JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
@ -56,6 +62,7 @@ class CacheUnavailable:
type CacheLookup = CacheHit | CaptureLease | CacheBusy | CacheUnavailable
type RequestSigner = Callable[[str, str, Mapping[str, str], bytes | None], dict[str, str]]
class ResponseStore(Protocol):
@ -83,28 +90,51 @@ class SignedResponse(BaseModel):
signature: str
def exact_key(secret: bytes, method: str, url: str, headers: Mapping[str, str], body: bytes | None) -> str:
def canonical_text(value: str) -> str:
return MARKER_PATTERN.sub(MARKER_PLACEHOLDER, value)
def canonical_body(body: bytes) -> bytes:
try:
return canonical_text(body.decode("utf-8")).encode("utf-8")
except UnicodeDecodeError:
return body
def request_identity(
secret: bytes, test_key: str, method: str, url: str, headers: Mapping[str, str], body: bytes | None,
) -> str:
fields: Final = (
b"provider-cache-exact-v1", method.encode(), url.encode(),
b"provider-cache-canonical-v2", test_key.encode(), method.encode(), canonical_text(url).encode(),
*(part.encode() for pair in sorted(headers.items()) for part in pair),
b"no-body" if body is None else b"body", b"" if body is None else body,
b"no-body" if body is None else b"body", b"" if body is None else canonical_body(body),
)
encoded: Final = b"".join(len(part).to_bytes(8, "big") + part for part in fields)
return hmac.new(secret, encoded, hashlib.sha256).hexdigest()
def cacheable_endpoint(method: str, url: str, body: bytes | None) -> bool:
return (
method == "POST"
and urlsplit(url).path in {"/v1/chat/completions", "/v1/messages"}
and body is not None
and len(body) <= MAX_REQUEST_BYTES
)
def slotted_key(secret: bytes, identity: str, slot: int) -> str:
return hmac.new(secret, f"{identity}:{slot}".encode(), hashlib.sha256).hexdigest()
def successful_response(url: str, status: int, headers: Mapping[str, str], body: bytes) -> bool:
def is_bedrock(mount: str) -> bool:
return mount.partition("/")[0] == BEDROCK_MOUNT_PREFIX
def cacheable_endpoint(mount: str, method: str, url: str, body: bytes | None) -> bool:
if method != "POST" or body is None or len(body) > MAX_REQUEST_BYTES:
return False
path: Final = urlsplit(url).path
if is_bedrock(mount):
return path.startswith("/model/") and path.endswith(("/converse", "/invoke"))
return path in {"/v1/chat/completions", "/v1/messages"}
def successful_response(mount: str, url: str, status: int, headers: Mapping[str, str], body: bytes) -> bool:
if not 200 <= status < 300 or len(body) > MAX_RESPONSE_BYTES:
return False
if is_bedrock(mount):
return complete_bedrock_response(url, body)
streaming: Final = "text/event-stream" in headers.get("content-type", "").lower()
if streaming:
try:
@ -147,6 +177,26 @@ def successful_response(url: str, status: int, headers: Mapping[str, str], body:
)
def complete_bedrock_response(url: str, body: bytes) -> bool:
"""Converse answers with ``output`` plus a ``stopReason``; InvokeModel on an
Anthropic model answers the Anthropic message shape. Either way a truncated
or error body is missing the terminator field, which is what makes it safe to
record. The streaming variants never reach here: they are not cacheable."""
try:
value: Final = JSON_VALUE.validate_json(body)
except ValidationError:
return False
if not isinstance(value, dict) or "message" in value:
return False
if urlsplit(url).path.endswith("/converse"):
return isinstance(value.get("output"), dict) and isinstance(value.get("stopReason"), str)
return (
value.get("type") == "message"
and isinstance(value.get("content"), list)
and isinstance(value.get("stop_reason"), str)
)
def complete_chat_stream(values: tuple[JsonValue, ...]) -> bool:
if any(not isinstance(value, dict) or not isinstance(value.get("choices"), list) for value in values):
return False
@ -172,7 +222,7 @@ def encode_response(secret: bytes, response: CachedResponse) -> bytes:
return SignedResponse(response=raw, signature=hmac.new(secret, raw.encode(), hashlib.sha256).hexdigest()).model_dump_json().encode()
def decode_response(secret: bytes, key: str, payload: bytes, url: str) -> CachedResponse | None:
def decode_response(secret: bytes, key: str, payload: bytes, mount: str, url: str) -> CachedResponse | None:
if len(payload) > 2 * MAX_RESPONSE_BYTES:
return None
try:
@ -183,7 +233,9 @@ def decode_response(secret: bytes, key: str, payload: bytes, url: str) -> Cached
chunks: Final = tuple(base64.b64decode(chunk, validate=True) for chunk in response.chunks)
except (ValidationError, ValueError):
return None
if response.request_key != key or not successful_response(url, response.status_code, response.headers, b"".join(chunks)):
if response.request_key != key or not successful_response(
mount, url, response.status_code, response.headers, b"".join(chunks)
):
return None
return response
@ -199,6 +251,24 @@ class CacheCounters:
self.counts = tuple((current | {name: current.get(name, 0) + 1}).items())
@dataclass(slots=True)
class SlotCounter:
"""FIFO position of a request among the canonically identical ones its test
has already sent. Two calls in one test that differ only by ``unique_marker``
canonicalize the same, so without this they would share one recording and the
second would replay the first's provider response id."""
counts: tuple[tuple[str, int], ...] = ()
lock: threading.Lock = field(default_factory=threading.Lock)
def take(self, identity: str) -> int:
with self.lock:
current: Final = dict(self.counts)
taken: Final = current.get(identity, 0)
self.counts = tuple((current | {identity: taken + 1}).items())
return taken
@dataclass(slots=True)
class ResponseCapture:
buffer: io.BytesIO = field(default_factory=io.BytesIO)
@ -231,9 +301,12 @@ class CacheEdge:
store: ResponseStore
secret: bytes = field(repr=False)
counters: CacheCounters = field(default_factory=CacheCounters)
slots: SlotCounter = field(default_factory=SlotCounter)
signers: Mapping[str, RequestSigner] = field(default_factory=dict)
wait_seconds: float = 2.0
clock: Callable[[], float] = time.monotonic
sleep: Callable[[float], None] = time.sleep
test_key: Callable[[], str] = current_test_key
def lookup(self, key: str) -> CacheLookup:
deadline: Final = self.clock() + self.wait_seconds
@ -241,39 +314,71 @@ class CacheEdge:
self.sleep(min(0.05, max(0, deadline - self.clock())))
return result
def forward(self, method: str, url: str, headers: dict[str, str], body: bytes | None, timeout: float) -> StreamHead | NetworkError:
if not cacheable_endpoint(method, url, body):
self.counters.increment("bypass")
self.counters.increment("upstream_attempts")
return forward_stream(method, url, headers=headers, body=body, timeout=timeout)
prepared: Final = prepare_forward(method, url, headers, body)
def count(self, mount: str, name: str) -> None:
self.counters.increment(name)
self.counters.increment(f"mount:{mount}:{name}")
def outbound(self, mount: str, method: str, url: str, headers: dict[str, str], body: bytes | None) -> dict[str, str]:
"""The headers actually sent upstream. A signing mount gets a signature
minted over the upstream URL, because the edge rewrote the Host the proxy
signed and Bedrock verifies it."""
signer: Final = self.signers.get(mount)
return headers if signer is None else signer(method, url, headers, body)
def keyed(self, mount: str, headers: Mapping[str, str]) -> Mapping[str, str]:
"""A signing mount's signature headers are the edge's own and carry a
timestamp, so keying on them would make every request a permanent miss.
Every other mount keys on its headers whole, credentials included, so a
different account can never read another's recording."""
if mount not in self.signers:
return headers
return {name: value for name, value in headers.items() if name.lower() not in SIGNATURE_HEADERS}
def forward(
self, mount: str, method: str, url: str, headers: dict[str, str], body: bytes | None, timeout: float,
) -> StreamHead | NetworkError:
test_key: Final = self.test_key()
if test_key == SESSION_TEST_KEY or not cacheable_endpoint(mount, method, url, body):
self.count(mount, "bypass")
self.count(mount, "upstream_attempts")
return forward_stream(
method, url, headers=self.outbound(mount, method, url, headers, body), body=body, timeout=timeout,
)
prepared: Final = prepare_forward(method, url, self.outbound(mount, method, url, headers, body), body)
if isinstance(prepared, NetworkError):
self.counters.increment("rejected")
self.count(mount, "rejected")
return prepared
key: Final = exact_key(self.secret, method, url, prepared.headers, body)
identity: Final = request_identity(
self.secret, test_key, method, url, self.keyed(mount, prepared.headers), body,
)
key: Final = slotted_key(self.secret, identity, self.slots.take(identity))
found: Final = self.lookup(key)
if isinstance(found, CacheHit):
response: Final = decode_response(self.secret, key, found.payload, url)
response: Final = decode_response(self.secret, key, found.payload, mount, url)
if response is not None and self.clock() < found.valid_until:
self.counters.increment("hits")
self.count(mount, "hits")
return StreamHead(response.status_code, response.headers, response_steps(response))
self.counters.increment("corrupt" if response is None else "expired")
self.count(mount, "corrupt" if response is None else "expired")
self.store.discard(key, found.payload)
capture_slot: Final = self.lookup(key) if isinstance(found, CacheHit) else found
self.counters.increment("misses")
self.count(mount, "misses")
if isinstance(capture_slot, CacheUnavailable):
self.counters.increment("cache_errors")
self.counters.increment("upstream_attempts")
self.count(mount, "cache_errors")
self.count(mount, "upstream_attempts")
head: Final = forward_prepared_stream(prepared, timeout)
if not isinstance(capture_slot, CaptureLease):
return head
if isinstance(head, NetworkError):
self.store.release(key, capture_slot)
self.counters.increment("rejected")
self.count(mount, "rejected")
return head
return StreamHead(head.status_code, head.headers, primed_steps(self.capture(key, capture_slot, url, head)))
return StreamHead(
head.status_code, head.headers, primed_steps(self.capture(mount, key, capture_slot, url, head)),
)
def capture(self, key: str, lease: CaptureLease, url: str, head: StreamHead) -> Generator[StreamStep, None, None]:
def capture(
self, mount: str, key: str, lease: CaptureLease, url: str, head: StreamHead,
) -> Generator[StreamStep, None, None]:
capture: Final = ResponseCapture()
try:
with closing(head.steps):
@ -285,15 +390,15 @@ class CacheEdge:
headers: Final = {
name: value for name, value in head.headers.items() if name.lower() not in UNRECORDED_RESPONSE_HEADERS
}
if not capture.eligible or not successful_response(url, head.status_code, headers, b"".join(chunks)):
self.counters.increment("rejected")
if not capture.eligible or not successful_response(mount, url, head.status_code, headers, b"".join(chunks)):
self.count(mount, "rejected")
return
response: Final = CachedResponse(
request_key=key, status_code=head.status_code, headers=headers,
chunks=tuple(base64.b64encode(chunk).decode("ascii") for chunk in chunks),
)
published: Final = self.store.publish(key, lease, encode_response(self.secret, response))
self.counters.increment("writes" if published else "write_failures")
self.count(mount, "writes" if published else "write_failures")
finally:
self.store.release(key, lease)
capture.buffer.close()

View file

@ -8,14 +8,57 @@ from models import LiteLLMParamsBody, ModelMode
LIVE_PROVIDER_REQUIRED: Final[ContextVar[bool]] = ContextVar("live_provider_required", default=False)
DEFAULT_BEDROCK_REGION: Final = "us-east-1"
BEDROCK_ANTHROPIC_INFIX: Final = "anthropic."
def bedrock_mount(params: LiteLLMParamsBody) -> str | None:
"""The edge mount an Anthropic-on-Bedrock deployment belongs to, or None.
Only the Anthropic models route. The edge validates converse and invoke
bodies by their Anthropic and Converse terminator fields, and the runner role
is allowed to invoke exactly those models, so Bedrock embeddings, image
generation, rerank and realtime keep their existing direct path rather than
reaching an edge that could neither sign nor validate for them."""
route: Final = params.model.partition("/")[2]
model: Final = route.partition("/")[2] or route
if BEDROCK_ANTHROPIC_INFIX not in model:
return None
return f"bedrock/{params.aws_region_name or DEFAULT_BEDROCK_REGION}"
def route_bedrock(
params: LiteLLMParamsBody, base_for: Callable[[str], str | None], mode: ModelMode | None,
) -> LiteLLMParamsBody:
"""Deployments that carry their own AWS identity stay off the edge. The edge
re-signs with the run pod's role, so routing an `aws_role_name` deployment
would quietly replace the very assume-role chain that test exists to prove."""
if mode is not None or params.aws_role_name is not None or params.aws_access_key_id is not None:
return params
if params.api_base is not None or params.aws_bedrock_runtime_endpoint is not None:
return params
mount: Final = bedrock_mount(params)
if mount is None:
return params
base: Final = base_for(mount)
if base is None:
return params
return params.model_copy(update={"aws_bedrock_runtime_endpoint": base})
def route_cache_model(
params: LiteLLMParamsBody, base_for: Callable[[str], str | None], *, enabled: bool, mode: ModelMode | None = None,
) -> LiteLLMParamsBody:
if not enabled or mode == "realtime" or LIVE_PROVIDER_REQUIRED.get() or params.api_base is not None or params.mock_response is not None:
if not enabled or LIVE_PROVIDER_REQUIRED.get() or params.mock_response is not None:
return params
if params.litellm_credential_name is not None:
return params
provider: Final = params.model.partition("/")[0]
if provider not in {"openai", "anthropic"} or params.litellm_credential_name is not None:
if provider == "bedrock":
return route_bedrock(params, base_for, mode)
if mode == "realtime" or params.api_base is not None:
return params
if provider not in {"openai", "anthropic"}:
return params
base: Final = base_for(provider)
if base is None:

View file

@ -48,7 +48,7 @@ import threading
from collections import deque
from collections.abc import Generator, Mapping, Sequence
from contextlib import closing, contextmanager
from dataclasses import dataclass, field
from dataclasses import dataclass, field, replace
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from itertools import islice
from pathlib import Path
@ -94,17 +94,41 @@ from fixture_mode import (
parse_fixture_mode,
)
from fixture_profile import IneligibleRequest, MatchProfile, match_profile, strict_identity
from provider_cache import CacheEdge
from provider_cache import CacheEdge, RequestSigner, is_bedrock
from provider_cache_routing import LIVE_PROVIDER_REQUIRED
from pydantic import JsonValue, TypeAdapter
BEDROCK_REGIONS: Final[tuple[str, ...]] = ("us-east-1",)
EDGE_MOUNTS: Final[Mapping[str, str]] = MappingProxyType(
{
"openai": "https://api.openai.com",
"anthropic": "https://api.anthropic.com",
**{
f"bedrock/{region}": f"https://bedrock-runtime.{region}.amazonaws.com"
for region in BEDROCK_REGIONS
},
}
)
@dataclass(frozen=True, slots=True)
class ResolvedMount:
mount: str
upstream_base: str
upstream_path: str
def resolve_mount(path: str, mounts: Mapping[str, str]) -> ResolvedMount | None:
"""Longest mount prefix wins, so a region-qualified mount such as
``bedrock/us-east-1`` resolves whole instead of leaving the region as the
first segment of the upstream path."""
trimmed: Final = path.lstrip("/")
for mount in sorted(mounts, key=len, reverse=True):
if trimmed == mount or trimmed.startswith(f"{mount}/"):
return ResolvedMount(mount, mounts[mount], trimmed[len(mount):].lstrip("/"))
return None
REPLAY_MISS_STATUS: Final = 599
_HOP_BY_HOP_HEADERS: Final[frozenset[str]] = frozenset(
@ -754,14 +778,14 @@ def _handle_record(
def _handle_live(
method: str, url: str, headers: Mapping[str, str], body: bytes | None, timeout: float,
cache: CacheEdge | None = None,
cache: CacheEdge | None = None, mount: str = "",
) -> EdgeOutcome:
forwarded: Final = {
name: value for name, value in headers.items() if name.lower() not in _REQUEST_DROPPED_HEADERS
}
head: Final = (
forward_stream(method, url, headers=forwarded, body=body, timeout=timeout)
if cache is None else cache.forward(method, url, forwarded, body, timeout)
if cache is None else cache.forward(mount, method, url, forwarded, body, timeout)
)
match head:
case NetworkError(message=message):
@ -796,10 +820,13 @@ def handle_edge_request(
prefix, then record (forward + persist) or replay (serve from the bundle).
Socket-free so unit tests exercise every branch without a server."""
split: Final = urlsplit(raw_path)
mount, _, upstream_path = split.path.lstrip("/").partition("/")
upstream_base: Final = mounts.get(mount)
if upstream_base is None:
return _text_reply(404, f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(mounts))}")
resolved: Final = resolve_mount(split.path, mounts)
if resolved is None:
unknown: Final = split.path.lstrip("/").partition("/")[0]
return _text_reply(404, f"unknown provider mount {unknown!r}; known mounts: {', '.join(sorted(mounts))}")
mount: Final = resolved.mount
upstream_base: Final = resolved.upstream_base
upstream_path: Final = resolved.upstream_path
profile: Final = (
backend.recorder.profile
if isinstance(backend, RecordEdge)
@ -830,7 +857,8 @@ def handle_edge_request(
match backend:
case CacheEdge():
return _handle_live(
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout, backend,
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout,
backend, mount,
)
case LiveEdge():
return _handle_live(
@ -891,7 +919,7 @@ class _EdgeHandler(BaseHTTPRequestHandler):
)
if isinstance(edge_server.backend, CacheEdge) and duplicate_headers:
edge_server.backend.counters.increment("duplicate_header_bypass")
if urlsplit(self.path).path.lstrip("/").partition("/")[0] in edge_server.mounts:
if resolve_mount(urlsplit(self.path).path, edge_server.mounts) is not None:
edge_server.backend.counters.increment("upstream_attempts")
outcome: Final = handle_edge_request(
selected_backend,
@ -1079,6 +1107,8 @@ def provider_edge_api_base(
return _shared_cache_edge(bind_host, advertise_host, forward_timeout).api_base(mount)
return None
case "record" | "replay":
if is_bedrock(mount):
return None
if mount not in EDGE_MOUNTS:
raise ValueError(f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(EDGE_MOUNTS))}")
return _shared_edge(mode, bundle_dir, bind_host, advertise_host, forward_timeout, match_profile()).api_base(
@ -1108,7 +1138,17 @@ def configured_cache_backend() -> CacheEdge | None:
return None
from provider_cache_redis import configured_cache
return configured_cache()
cache: Final = configured_cache()
return None if cache is None else replace(cache, signers=bedrock_signers())
@functools.lru_cache(maxsize=1)
def bedrock_signers() -> Mapping[str, RequestSigner]:
"""One signer per mounted Bedrock region, built lazily so a run that never
mounts Bedrock neither imports botocore nor resolves an AWS identity."""
from provider_edge_bedrock import bedrock_signer
return MappingProxyType({f"bedrock/{region}": bedrock_signer(region) for region in BEDROCK_REGIONS})
@functools.lru_cache(maxsize=8)

View file

@ -0,0 +1,72 @@
"""SigV4 re-signing for Bedrock traffic routed through the provider edge.
Bedrock is the one provider the edge could never mount. SigV4 signs the Host
header, so rewriting ``api_base`` to point at the edge invalidates the proxy's
signature and Bedrock rejects the call before it reaches a model. The edge
therefore has to drop the proxy's signature and mint its own over the upstream
URL it is actually about to call.
The identity it signs with is the run pod's own, from the EKS Pod Identity
association on ServiceAccount ``buildkite-e2e-run``. That role carries Bedrock
invoke and converse on an allowlist of the Anthropic models the suite registers
and nothing else, so a re-signed call can reach exactly the models the suite
already uses. The proxy's own Bedrock credentials are not involved in a routed
deployment, which is why ``aws_role_name`` deployments stay off the edge: their
whole point is to prove the product's assume-role chain.
Signature headers are excluded from the cache key by the caller, and they have
to be: ``x-amz-date`` is a timestamp, so keying on it would make every Bedrock
request a permanent miss.
"""
from __future__ import annotations
import functools
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Final
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
from botocore.session import Session
from provider_cache import SIGNATURE_HEADERS
BEDROCK_SERVICE: Final = "bedrock"
class MissingAwsCredentials(RuntimeError):
"""No AWS identity is resolvable, so the edge cannot sign for Bedrock."""
@dataclass(frozen=True, slots=True)
class BedrockSigner:
region: str
credentials: Callable[[], Credentials]
def __call__(self, method: str, url: str, headers: Mapping[str, str], body: bytes | None) -> dict[str, str]:
unsigned: Final = {
name: value for name, value in headers.items() if name.lower() not in SIGNATURE_HEADERS
}
request: Final = AWSRequest(method=method, url=url, headers=unsigned, data=body or b"")
SigV4Auth(self.credentials(), BEDROCK_SERVICE, self.region).add_auth(request)
return dict(request.headers)
@functools.lru_cache(maxsize=1)
def pod_credentials() -> Credentials:
"""The run pod's own identity, resolved once per process through botocore's
ordinary chain, which reaches Pod Identity at the ``container-role`` link."""
resolved: Final = Session().get_credentials()
if resolved is None:
raise MissingAwsCredentials(
"the provider edge is mounted for Bedrock but no AWS credentials resolve; "
"the run pod gets them from the Pod Identity association on buildkite-e2e-run"
)
return resolved
def bedrock_signer(region: str, credentials: Callable[[], Credentials] = pod_credentials) -> BedrockSigner:
"""Credentials are resolved on the first signed request, not here, so a run
that mounts Bedrock but never calls it needs no AWS identity at all."""
return BedrockSigner(region, credentials)

View file

@ -1279,15 +1279,30 @@ class TestApiBaseSeam:
)
def test_unknown_mount_raises_naming_the_known_mounts(self, tmp_path: Path) -> None:
with pytest.raises(ValueError, match="unknown provider mount 'bedrock'"):
with pytest.raises(ValueError, match="unknown provider mount 'cohere'"):
provider_edge_api_base(
"bedrock",
"cohere",
mode_raw="record",
bundle_dir=tmp_path / "bundle",
bind_host="127.0.0.1",
advertise_host="127.0.0.1",
)
@pytest.mark.parametrize("mode_raw", ["record", "replay"])
def test_bedrock_never_wires_a_bundle_because_the_edge_cannot_sign_into_one(
self, tmp_path: Path, mode_raw: str,
) -> None:
"""Record and replay serve from a bundle without re-signing, so a Bedrock
deployment pointed at that edge would send the proxy's signature over a
rewritten Host. It keeps its direct route in both modes."""
assert provider_edge_api_base(
"bedrock/us-east-1",
mode_raw=mode_raw,
bundle_dir=tmp_path / "bundle",
bind_host="127.0.0.1",
advertise_host="127.0.0.1",
) is None
def test_record_mode_boots_one_shared_edge_and_prepares_the_bundle(self, tmp_path: Path) -> None:
root = tmp_path / "bundle"
first = provider_edge_api_base(