mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
a8979fe054
commit
2d40254b57
8 changed files with 768 additions and 118 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
72
tests/e2e/provider_edge_bedrock.py
Normal file
72
tests/e2e/provider_edge_bedrock.py
Normal 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)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue