mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 5ac3f5edfe into 635085ac14
This commit is contained in:
commit
5355beb676
5 changed files with 670 additions and 2 deletions
|
|
@ -31,6 +31,8 @@ CLAUDE_CODE_REASONING_BETAS: Final = (
|
|||
"interleaved-thinking-2025-05-14",
|
||||
"thinking-token-count-2026-05-13",
|
||||
)
|
||||
CLAUDE_CODE_CACHING_BETAS: Final = ("extended-cache-ttl-2025-04-11", "prompt-caching-scope-2026-01-05")
|
||||
CACHING_CLI_BETA: Final = f"{CLI_BETA},extended-cache-ttl-2025-04-11"
|
||||
REASONING_FIELDS: Final = ("thinking", "output_config", "reasoning_effort", "temperature")
|
||||
_OUTPUT_USAGE_KEYS: Final = frozenset({"output_tokens", "output_tokens_details"})
|
||||
METADATA_USER_ID: Final = json.dumps(
|
||||
|
|
@ -731,6 +733,114 @@ def forwarded(sent: Mapping[str, JsonValue], request: Request) -> Forwarded:
|
|||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CacheForwarded:
|
||||
breakpoints: dict[str, JsonValue]
|
||||
other_changes: dict[str, JsonValue]
|
||||
caching_betas: tuple[str, ...]
|
||||
|
||||
|
||||
def _marked(prefix: str, blocks: JsonValue) -> tuple[tuple[str, JsonValue], ...]:
|
||||
if not isinstance(blocks, list):
|
||||
return ()
|
||||
return tuple(
|
||||
(f"{prefix}[{index}]", block["cache_control"])
|
||||
for index, block in enumerate(blocks)
|
||||
if isinstance(block, dict) and "cache_control" in block
|
||||
)
|
||||
|
||||
|
||||
def _marked_message(index: int, message: JsonValue) -> tuple[tuple[str, JsonValue], ...]:
|
||||
return _marked(f"messages[{index}].content", message.get("content")) if isinstance(message, dict) else ()
|
||||
|
||||
|
||||
def _marked_messages(messages: JsonValue) -> tuple[tuple[str, JsonValue], ...]:
|
||||
if not isinstance(messages, list):
|
||||
return ()
|
||||
return tuple(chain.from_iterable(_marked_message(index, message) for index, message in enumerate(messages)))
|
||||
|
||||
|
||||
def cache_breakpoints(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
top_level: Final = (("cache_control", body["cache_control"]),) if "cache_control" in body else ()
|
||||
return dict(
|
||||
(
|
||||
*top_level,
|
||||
*_marked("system", body.get("system")),
|
||||
*_marked("tools", body.get("tools")),
|
||||
*_marked_messages(body.get("messages")),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _unmarked_block(block: JsonValue) -> JsonValue:
|
||||
if not isinstance(block, dict):
|
||||
return block
|
||||
return {key: value for key, value in block.items() if key != "cache_control"}
|
||||
|
||||
|
||||
def _unmarked_blocks(blocks: JsonValue) -> JsonValue:
|
||||
return [_unmarked_block(block) for block in blocks] if isinstance(blocks, list) else blocks
|
||||
|
||||
|
||||
def _unmarked_message(message: JsonValue) -> JsonValue:
|
||||
if not isinstance(message, dict) or "content" not in message:
|
||||
return message
|
||||
return {**message, "content": _unmarked_blocks(message["content"])}
|
||||
|
||||
|
||||
def _unmarked_value(key: str, value: JsonValue) -> JsonValue:
|
||||
match key, value:
|
||||
case ("system" | "tools", _):
|
||||
return _unmarked_blocks(value)
|
||||
case ("messages", list()):
|
||||
return [_unmarked_message(message) for message in value]
|
||||
case _:
|
||||
return value
|
||||
|
||||
|
||||
def with_block_breakpoint(
|
||||
body: Mapping[str, JsonValue], field: str, index: int, control: Mapping[str, JsonValue]
|
||||
) -> dict[str, JsonValue]:
|
||||
blocks: Final = body[field]
|
||||
assert isinstance(blocks, list), body
|
||||
return {
|
||||
**body,
|
||||
field: [
|
||||
{**block, "cache_control": dict(control)} if position == index and isinstance(block, dict) else block
|
||||
for position, block in enumerate(blocks)
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def without_cache_breakpoints(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return {key: _unmarked_value(key, value) for key, value in body.items() if key != "cache_control"}
|
||||
|
||||
|
||||
def caching_betas(anthropic_beta: str) -> tuple[str, ...]:
|
||||
return tuple(sorted(beta for beta in anthropic_beta.split(",") if beta in CLAUDE_CODE_CACHING_BETAS))
|
||||
|
||||
|
||||
def _without_model(body: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return {key: value for key, value in without_cache_breakpoints(body).items() if key != "model"}
|
||||
|
||||
|
||||
def cache_forwarded(sent: Mapping[str, JsonValue], request: Request) -> CacheForwarded:
|
||||
body: Final = JSON_OBJECT.validate_json(request.body)
|
||||
return CacheForwarded(
|
||||
breakpoints=cache_breakpoints(body),
|
||||
other_changes=body_diff(_without_model(sent), _without_model(body)),
|
||||
caching_betas=caching_betas(request.headers.get("anthropic-beta", "")),
|
||||
)
|
||||
|
||||
|
||||
def streamed_start_usage(stream: str) -> JsonValue:
|
||||
return next(
|
||||
JSON_OBJECT.validate_python(data["message"]).get("usage")
|
||||
for event, data in _client_events(stream)
|
||||
if event == "message_start"
|
||||
)
|
||||
|
||||
|
||||
def _appended(block: Mapping[str, JsonValue], key: str, text: str) -> dict[str, JsonValue]:
|
||||
return {**block, key: f"{block.get(key) or ''}{text}"}
|
||||
|
||||
|
|
@ -848,7 +958,11 @@ def message_reply(
|
|||
|
||||
|
||||
def message_stream(
|
||||
identity: str, model: str, content: tuple[dict[str, JsonValue], ...], usage: dict[str, JsonValue]
|
||||
identity: str,
|
||||
model: str,
|
||||
content: tuple[dict[str, JsonValue], ...],
|
||||
usage: dict[str, JsonValue],
|
||||
final_usage: Mapping[str, JsonValue] | None = None,
|
||||
) -> tuple[bytes, ...]:
|
||||
start: Final = sse_frame(
|
||||
"message_start",
|
||||
|
|
@ -871,7 +985,7 @@ def message_stream(
|
|||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
|
||||
"usage": _delta_usage(usage),
|
||||
"usage": _delta_usage(usage) if final_usage is None else dict(final_usage),
|
||||
},
|
||||
)
|
||||
blocks: Final = chain.from_iterable(_block_frames(index, block) for index, block in enumerate(content))
|
||||
|
|
|
|||
|
|
@ -0,0 +1,234 @@
|
|||
import uuid
|
||||
from collections.abc import Callable, Mapping
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support import claude_code as cc
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
_MODEL: Final = "claude-sonnet-4-5"
|
||||
_FIVE_MINUTES: Final = {"type": "ephemeral"}
|
||||
_PONG: Final[tuple[dict[str, JsonValue], ...]] = ({"type": "text", "text": "PONG"},)
|
||||
_SUBAGENT_BILLING: Final = (
|
||||
"x-anthropic-billing-header: cc_version=2.1.283.00; cc_entrypoint=sdk-cli; cc_is_subagent=true;"
|
||||
)
|
||||
_MAIN_AGENT_BILLING: Final = "x-anthropic-billing-header: cc_version=2.1.283.00; cc_entrypoint=sdk-cli;"
|
||||
|
||||
|
||||
def _unmarked_claude_code_turn() -> dict[str, JsonValue]:
|
||||
return cc.without_cache_breakpoints({**cc.claude_code_request(f"cache-bust-{uuid.uuid4().hex}"), "stream": False})
|
||||
|
||||
|
||||
def _one_shot_turn(billing: str) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"system": [{"type": "text", "text": billing}],
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": f"Summarize {uuid.uuid4().hex}"}]}],
|
||||
"max_tokens": 512,
|
||||
"stream": False,
|
||||
}
|
||||
|
||||
|
||||
def _forward(
|
||||
gateway: Gateway,
|
||||
sent: Mapping[str, JsonValue],
|
||||
key_fields: Mapping[str, JsonValue],
|
||||
deployment: Mapping[str, JsonValue],
|
||||
) -> tuple[cc.CacheForwarded, dict[str, JsonValue], dict[str, JsonValue]]:
|
||||
identity: Final = f"msg_{uuid.uuid4().hex}"
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
return Reply(body=cc.message_reply(identity, _MODEL, _PONG, {"input_tokens": 12, "output_tokens": 4}))
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=cc.ANTHROPIC_API_KEY, **deployment
|
||||
)
|
||||
key: Final = scenario.key(**key_fields)
|
||||
response: Final = gateway.request(
|
||||
"POST", "/v1/messages", {**sent, "model": model}, key=key, headers=cc.cli_headers(key, cc.CACHING_CLI_BETA)
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
upstream: Final = wire.drain()
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
"SELECT model_id, metadata->>'litellm_gateway_injected_cache' AS injected_for "
|
||||
'FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
|
||||
(identity,),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
assert len(upstream) == 1, upstream
|
||||
return cc.cache_forwarded(sent, upstream[0]), cc.JSON_OBJECT.validate_json(upstream[0].body), rows[0]
|
||||
|
||||
|
||||
def test_prompt_caching_key_adds_breakpoints_to_the_system_prompt_and_last_message(gateway: Gateway) -> None:
|
||||
sent: Final = _unmarked_claude_code_turn()
|
||||
forwarded, _, spend_row = _forward(gateway, sent, {"enable_prompt_caching": True}, {})
|
||||
assert cc.cache_breakpoints(sent) == {}
|
||||
assert forwarded.breakpoints == {"system[2]": _FIVE_MINUTES, "messages[0].content[7]": _FIVE_MINUTES}
|
||||
assert forwarded.other_changes == {}
|
||||
assert forwarded.caching_betas == ("extended-cache-ttl-2025-04-11", "prompt-caching-scope-2026-01-05")
|
||||
assert spend_row["injected_for"] == spend_row["model_id"], spend_row
|
||||
|
||||
|
||||
def test_prompt_caching_key_turns_a_string_system_prompt_into_a_cached_block(gateway: Gateway) -> None:
|
||||
sent: Final = {**_unmarked_claude_code_turn(), "system": "Synthetic agent identity system prompt."}
|
||||
forwarded, _, _ = _forward(gateway, sent, {"enable_prompt_caching": True}, {})
|
||||
assert forwarded.breakpoints == {"system[0]": _FIVE_MINUTES, "messages[0].content[7]": _FIVE_MINUTES}
|
||||
assert forwarded.other_changes == {
|
||||
"system": {
|
||||
"expected": "Synthetic agent identity system prompt.",
|
||||
"upstream": [{"type": "text", "text": "Synthetic agent identity system prompt."}],
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_key_without_prompt_caching_adds_no_breakpoints(gateway: Gateway) -> None:
|
||||
sent: Final = _unmarked_claude_code_turn()
|
||||
forwarded, _, spend_row = _forward(gateway, sent, {}, {})
|
||||
assert forwarded.breakpoints == {}
|
||||
assert forwarded.other_changes == {}
|
||||
assert spend_row["injected_for"] is None, spend_row
|
||||
|
||||
|
||||
def test_prompt_caching_key_adds_nothing_when_the_client_marked_any_breakpoint(gateway: Gateway) -> None:
|
||||
sent: Final = cc.with_block_breakpoint(_unmarked_claude_code_turn(), "tools", 23, _FIVE_MINUTES)
|
||||
forwarded, _, spend_row = _forward(gateway, sent, {"enable_prompt_caching": True}, {})
|
||||
assert forwarded.breakpoints == {"tools[23]": _FIVE_MINUTES}
|
||||
assert forwarded.other_changes == {}
|
||||
assert spend_row["injected_for"] is None, spend_row
|
||||
|
||||
|
||||
def test_prompt_caching_key_adds_nothing_when_extra_body_unmarks_the_client_tool_breakpoint(gateway: Gateway) -> None:
|
||||
marked: Final = cc.with_block_breakpoint(_unmarked_claude_code_turn(), "tools", 23, _FIVE_MINUTES)
|
||||
envelope: Final = {"tools": cc.without_cache_breakpoints(marked)["tools"]}
|
||||
sent: Final = {**marked, "extra_body": envelope}
|
||||
forwarded, _, spend_row = _forward(gateway, sent, {"enable_prompt_caching": True}, {})
|
||||
assert forwarded.breakpoints == {"tools[23]": _FIVE_MINUTES}
|
||||
assert forwarded.other_changes == {"extra_body": {"expected": envelope, "upstream": None}}
|
||||
assert spend_row["injected_for"] is None, spend_row
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("billing", "received"),
|
||||
(
|
||||
pytest.param(_SUBAGENT_BILLING, {}, id="subagent-gets-no-breakpoints"),
|
||||
pytest.param(
|
||||
_MAIN_AGENT_BILLING,
|
||||
{"system[0]": _FIVE_MINUTES, "messages[0].content[0]": _FIVE_MINUTES},
|
||||
id="main-agent-gets-breakpoints",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_prompt_caching_key_skips_one_shot_claude_code_subagent_requests(
|
||||
gateway: Gateway, billing: str, received: dict[str, JsonValue]
|
||||
) -> None:
|
||||
sent: Final = _one_shot_turn(billing)
|
||||
forwarded, _, _ = _forward(gateway, sent, {"enable_prompt_caching": True}, {})
|
||||
assert forwarded.breakpoints == received
|
||||
assert forwarded.other_changes == {}
|
||||
|
||||
|
||||
def _three_client_breakpoints() -> dict[str, JsonValue]:
|
||||
system_marked: Final = cc.with_block_breakpoint(_unmarked_claude_code_turn(), "system", 1, _FIVE_MINUTES)
|
||||
return {**cc.with_block_breakpoint(system_marked, "tools", 23, _FIVE_MINUTES), "cache_control": dict(_FIVE_MINUTES)}
|
||||
|
||||
|
||||
def _four_client_breakpoints() -> dict[str, JsonValue]:
|
||||
return cc.with_block_breakpoint(_three_client_breakpoints(), "system", 2, _FIVE_MINUTES)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("sent", "received"),
|
||||
(
|
||||
pytest.param(
|
||||
_three_client_breakpoints,
|
||||
{
|
||||
"cache_control": _FIVE_MINUTES,
|
||||
"system[1]": _FIVE_MINUTES,
|
||||
"tools[23]": _FIVE_MINUTES,
|
||||
"messages[0].content[7]": _FIVE_MINUTES,
|
||||
},
|
||||
id="three-client-breakpoints-plus-last-message",
|
||||
),
|
||||
pytest.param(
|
||||
_four_client_breakpoints,
|
||||
{
|
||||
"cache_control": _FIVE_MINUTES,
|
||||
"system[1]": _FIVE_MINUTES,
|
||||
"system[2]": _FIVE_MINUTES,
|
||||
"tools[23]": _FIVE_MINUTES,
|
||||
},
|
||||
id="four-client-breakpoints-leave-no-room",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_configured_injection_points_stop_at_anthropics_four_breakpoint_limit(
|
||||
gateway: Gateway, sent: Callable[[], dict[str, JsonValue]], received: dict[str, JsonValue]
|
||||
) -> None:
|
||||
client_body: Final = sent()
|
||||
forwarded, _, _ = _forward(
|
||||
gateway,
|
||||
client_body,
|
||||
{},
|
||||
{
|
||||
"cache_control_injection_points": [
|
||||
{"location": "message", "role": "system"},
|
||||
{"location": "message", "index": -1},
|
||||
]
|
||||
},
|
||||
)
|
||||
assert forwarded.breakpoints == received
|
||||
assert forwarded.other_changes == {}
|
||||
|
||||
|
||||
def test_fallback_deployment_gets_its_own_breakpoints_and_the_injection_credit(gateway: Gateway) -> None:
|
||||
identity: Final = f"msg_{uuid.uuid4().hex}"
|
||||
sent: Final = _unmarked_claude_code_turn()
|
||||
|
||||
def fail(request: Request) -> Reply:
|
||||
return Reply(status=500, body=b'{"type":"error","error":{"type":"api_error","message":"primary down"}}')
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
return Reply(body=cc.message_reply(identity, _MODEL, _PONG, {"input_tokens": 12, "output_tokens": 4}))
|
||||
|
||||
with wire_server(fail) as primary_wire, wire_server(respond) as fallback_wire, gateway.scenario() as scenario:
|
||||
primary: Final = scenario.model(
|
||||
model=f"anthropic/{_MODEL}", api_base=primary_wire.url, api_key=cc.ANTHROPIC_API_KEY
|
||||
)
|
||||
fallback: Final = scenario.model(
|
||||
model=f"anthropic/{_MODEL}", api_base=fallback_wire.url, api_key=cc.ANTHROPIC_API_KEY
|
||||
)
|
||||
key: Final = scenario.key(enable_prompt_caching=True)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/messages",
|
||||
{**sent, "model": primary, "fallbacks": [fallback]},
|
||||
key=key,
|
||||
headers=cc.cli_headers(key, cc.CACHING_CLI_BETA),
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
primary_legs: Final = primary_wire.drain()
|
||||
fallback_legs: Final = fallback_wire.drain()
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
"SELECT spend.model_id, deployment.model_name, "
|
||||
"spend.metadata->>'litellm_gateway_injected_cache' AS injected_for "
|
||||
'FROM "LiteLLM_SpendLogs" spend JOIN "LiteLLM_ProxyModelTable" deployment '
|
||||
"ON deployment.model_id = spend.model_id WHERE spend.request_id=%s",
|
||||
(identity,),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
injected: Final = {"system[2]": _FIVE_MINUTES, "messages[0].content[7]": _FIVE_MINUTES}
|
||||
legs: Final = [cc.cache_forwarded(sent, leg) for leg in (*primary_legs, *fallback_legs)]
|
||||
assert len(primary_legs) == 1, primary_legs
|
||||
assert len(fallback_legs) == 1, fallback_legs
|
||||
assert [(leg.breakpoints, leg.other_changes) for leg in legs] == [(injected, {}), (injected, {})], legs
|
||||
assert rows[0]["model_name"] == fallback, rows
|
||||
assert rows[0]["injected_for"] == rows[0]["model_id"], rows
|
||||
|
|
@ -0,0 +1,86 @@
|
|||
import uuid
|
||||
from collections.abc import Callable
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support import claude_code as cc
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
_MODEL: Final = "claude-sonnet-4-5"
|
||||
_FIVE_MINUTES: Final = {"type": "ephemeral"}
|
||||
_ONE_HOUR: Final = {"type": "ephemeral", "ttl": "1h"}
|
||||
_ONE_HOUR_GLOBAL: Final = {"type": "ephemeral", "ttl": "1h", "scope": "global"}
|
||||
_CLAUDE_CODE_BREAKPOINTS: Final = {
|
||||
"system[1]": _FIVE_MINUTES,
|
||||
"system[2]": _FIVE_MINUTES,
|
||||
"messages[0].content[7]": _FIVE_MINUTES,
|
||||
}
|
||||
|
||||
|
||||
def _claude_code_default(body: dict[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return body
|
||||
|
||||
|
||||
def _one_hour_global_system_prompt(body: dict[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return cc.with_block_breakpoint(body, "system", 2, _ONE_HOUR_GLOBAL)
|
||||
|
||||
|
||||
def _one_hour_tool_breakpoint(body: dict[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return cc.with_block_breakpoint(body, "tools", 23, _ONE_HOUR)
|
||||
|
||||
|
||||
def _top_level_automatic_caching(body: dict[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return {**body, "cache_control": dict(_FIVE_MINUTES)}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("mark", "received"),
|
||||
(
|
||||
pytest.param(_claude_code_default, _CLAUDE_CODE_BREAKPOINTS, id="claude-code-default"),
|
||||
pytest.param(
|
||||
_one_hour_global_system_prompt,
|
||||
{**_CLAUDE_CODE_BREAKPOINTS, "system[2]": _ONE_HOUR_GLOBAL},
|
||||
id="1h-global-scope-on-system-prompt",
|
||||
),
|
||||
pytest.param(
|
||||
_one_hour_tool_breakpoint,
|
||||
{**_CLAUDE_CODE_BREAKPOINTS, "tools[23]": _ONE_HOUR},
|
||||
id="1h-on-last-tool",
|
||||
),
|
||||
pytest.param(
|
||||
_top_level_automatic_caching,
|
||||
{"cache_control": _FIVE_MINUTES, **_CLAUDE_CODE_BREAKPOINTS},
|
||||
id="top-level-automatic-caching",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_client_cache_breakpoints_reach_anthropic_unchanged(
|
||||
gateway: Gateway, mark: Callable[[dict[str, JsonValue]], dict[str, JsonValue]], received: dict[str, JsonValue]
|
||||
) -> None:
|
||||
sent: Final = mark({**cc.claude_code_request(f"cache-bust-{uuid.uuid4().hex}"), "stream": False})
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
return Reply(
|
||||
body=cc.message_reply(
|
||||
f"msg_{uuid.uuid4().hex}",
|
||||
_MODEL,
|
||||
({"type": "text", "text": "PONG"},),
|
||||
{"input_tokens": 12, "output_tokens": 4},
|
||||
)
|
||||
)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=cc.ANTHROPIC_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST", "/v1/messages", {**sent, "model": model}, headers=cc.cli_headers(gateway.key, cc.CACHING_CLI_BETA)
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
upstream: Final = wire.drain()
|
||||
assert len(upstream) == 1, upstream
|
||||
forwarded: Final = cc.cache_forwarded(sent, upstream[0])
|
||||
assert cc.cache_breakpoints(sent) == received
|
||||
assert forwarded.breakpoints == received
|
||||
assert forwarded.other_changes == {}
|
||||
assert forwarded.caching_betas == ("extended-cache-ttl-2025-04-11", "prompt-caching-scope-2026-01-05")
|
||||
|
|
@ -0,0 +1,159 @@
|
|||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support import claude_code as cc
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue
|
||||
|
||||
_MODEL: Final = "claude-sonnet-4-5"
|
||||
_FAST_MODE_MODEL: Final = "claude-opus-5-5"
|
||||
_INPUT_RATE: Final = 1e-6
|
||||
_OUTPUT_RATE: Final = 2e-6
|
||||
_CACHE_READ_RATE: Final = 1e-7
|
||||
_CACHE_WRITE_5M_RATE: Final = 1.25e-6
|
||||
_CACHE_WRITE_1H_RATE: Final = 2e-6
|
||||
_FAST_MULTIPLIER: Final = 6.0
|
||||
_RATES: Final = {
|
||||
"input_cost_per_token": _INPUT_RATE,
|
||||
"output_cost_per_token": _OUTPUT_RATE,
|
||||
"cache_read_input_token_cost": _CACHE_READ_RATE,
|
||||
"cache_creation_input_token_cost": _CACHE_WRITE_5M_RATE,
|
||||
"cache_creation_input_token_cost_above_1hr": _CACHE_WRITE_1H_RATE,
|
||||
}
|
||||
_CONTENT: Final[tuple[dict[str, JsonValue], ...]] = ({"type": "text", "text": "PONG"},)
|
||||
_USAGE: Final[dict[str, JsonValue]] = {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"cache_read_input_tokens": 400,
|
||||
"cache_creation_input_tokens": 300,
|
||||
"cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200},
|
||||
}
|
||||
|
||||
|
||||
class _BilledRow(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
spend: float
|
||||
prompt_tokens: int
|
||||
cache_read_cost: float
|
||||
cache_creation_cost: float
|
||||
|
||||
|
||||
_STREAM_IDS: Final = (pytest.param(False, id="non-streamed"), pytest.param(True, id="streamed"))
|
||||
|
||||
|
||||
def _reply(identity: str, model: str, usage: dict[str, JsonValue], stream: bool) -> Reply:
|
||||
if stream:
|
||||
return Reply(chunks=cc.message_stream(identity, model, _CONTENT, usage), content_type="text/event-stream")
|
||||
return Reply(body=cc.message_reply(identity, model, _CONTENT, usage))
|
||||
|
||||
|
||||
def _billed_row(
|
||||
gateway: Gateway,
|
||||
model_name: str,
|
||||
usage: dict[str, JsonValue],
|
||||
stream: bool,
|
||||
request_fields: Mapping[str, JsonValue],
|
||||
model_info: Mapping[str, JsonValue] | None,
|
||||
) -> _BilledRow:
|
||||
identity: Final = f"msg_{uuid.uuid4().hex}"
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
return _reply(identity, model_name, usage, stream)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model=f"anthropic/{model_name}",
|
||||
api_base=wire.url,
|
||||
api_key=cc.ANTHROPIC_API_KEY,
|
||||
model_info=model_info,
|
||||
**_RATES,
|
||||
)
|
||||
body: Final = {
|
||||
**cc.claude_code_request(f"cache-bust-{uuid.uuid4().hex}"),
|
||||
**request_fields,
|
||||
"stream": stream,
|
||||
"model": model,
|
||||
}
|
||||
response: Final = gateway.request("POST", "/v1/messages", body)
|
||||
assert response.status_code == 200, response.text
|
||||
assert len(wire.drain()) == 1
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
"SELECT spend, prompt_tokens, "
|
||||
"(metadata->'cost_breakdown'->>'cache_read_cost')::float AS cache_read_cost, "
|
||||
"(metadata->'cost_breakdown'->>'cache_creation_cost')::float AS cache_creation_cost "
|
||||
'FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
|
||||
(identity,),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
return _BilledRow.model_validate(rows[0])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAM_IDS)
|
||||
def test_cache_reads_and_5m_and_1h_cache_writes_are_each_billed_at_their_own_rate(
|
||||
gateway: Gateway, stream: bool
|
||||
) -> None:
|
||||
row: Final = _billed_row(gateway, _MODEL, _USAGE, stream, {}, None)
|
||||
input_tokens, output_tokens, cache_read, cache_write_5m, cache_write_1h = 100, 50, 400, 100, 200
|
||||
assert row.prompt_tokens == input_tokens + cache_read + cache_write_5m + cache_write_1h, row
|
||||
assert row.cache_read_cost == pytest.approx(cache_read * _CACHE_READ_RATE), row
|
||||
assert row.cache_creation_cost == pytest.approx(
|
||||
cache_write_5m * _CACHE_WRITE_5M_RATE + cache_write_1h * _CACHE_WRITE_1H_RATE
|
||||
), row
|
||||
assert row.spend == pytest.approx(
|
||||
input_tokens * _INPUT_RATE
|
||||
+ output_tokens * _OUTPUT_RATE
|
||||
+ cache_read * _CACHE_READ_RATE
|
||||
+ cache_write_5m * _CACHE_WRITE_5M_RATE
|
||||
+ cache_write_1h * _CACHE_WRITE_1H_RATE
|
||||
), row
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAM_IDS)
|
||||
def test_fast_mode_multiplies_cache_reads_and_writes_along_with_input_and_output(
|
||||
gateway: Gateway, stream: bool
|
||||
) -> None:
|
||||
row: Final = _billed_row(
|
||||
gateway,
|
||||
_FAST_MODE_MODEL,
|
||||
{**_USAGE, "speed": "fast"},
|
||||
stream,
|
||||
{"speed": "fast"},
|
||||
{"provider_specific_entry": {"fast": _FAST_MULTIPLIER}},
|
||||
)
|
||||
input_tokens, output_tokens, cache_read, cache_write_5m, cache_write_1h = 100, 50, 400, 100, 200
|
||||
assert row.spend == pytest.approx(
|
||||
_FAST_MULTIPLIER
|
||||
* (
|
||||
input_tokens * _INPUT_RATE
|
||||
+ output_tokens * _OUTPUT_RATE
|
||||
+ cache_read * _CACHE_READ_RATE
|
||||
+ cache_write_5m * _CACHE_WRITE_5M_RATE
|
||||
+ cache_write_1h * _CACHE_WRITE_1H_RATE
|
||||
)
|
||||
), row
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAM_IDS)
|
||||
def test_fast_mode_cost_breakdown_reports_the_multiplied_cache_costs(gateway: Gateway, stream: bool) -> None:
|
||||
pytest.skip("BUG: fast mode bills 6x cache costs but cost_breakdown.cache_read_cost/cache_creation_cost stay 1x")
|
||||
row: Final = _billed_row(
|
||||
gateway,
|
||||
_FAST_MODE_MODEL,
|
||||
{**_USAGE, "speed": "fast"},
|
||||
stream,
|
||||
{"speed": "fast"},
|
||||
{"provider_specific_entry": {"fast": _FAST_MULTIPLIER}},
|
||||
)
|
||||
cache_read, cache_write_5m, cache_write_1h = 400, 100, 200
|
||||
assert row.cache_read_cost == pytest.approx(_FAST_MULTIPLIER * cache_read * _CACHE_READ_RATE), row
|
||||
assert row.cache_creation_cost == pytest.approx(
|
||||
_FAST_MULTIPLIER * (cache_write_5m * _CACHE_WRITE_5M_RATE + cache_write_1h * _CACHE_WRITE_1H_RATE)
|
||||
), row
|
||||
|
|
@ -0,0 +1,75 @@
|
|||
import uuid
|
||||
from typing import Final
|
||||
|
||||
from integration._support import claude_code as cc
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue
|
||||
|
||||
_MODEL: Final = "claude-sonnet-4-5"
|
||||
_ANTHROPIC_USAGE: Final = {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"cache_read_input_tokens": 400,
|
||||
"cache_creation_input_tokens": 300,
|
||||
"cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200},
|
||||
}
|
||||
|
||||
|
||||
def _client_body(stream: bool) -> dict[str, JsonValue]:
|
||||
return {**cc.claude_code_request(f"cache-bust-{uuid.uuid4().hex}"), "stream": stream}
|
||||
|
||||
|
||||
def test_streamed_cache_read_and_write_tokens_reach_the_client_unchanged(gateway: Gateway) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
return Reply(
|
||||
chunks=cc.message_stream(
|
||||
f"msg_{uuid.uuid4().hex}",
|
||||
_MODEL,
|
||||
({"type": "text", "text": "PONG"},),
|
||||
_ANTHROPIC_USAGE,
|
||||
final_usage=_ANTHROPIC_USAGE,
|
||||
),
|
||||
content_type="text/event-stream",
|
||||
)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=cc.ANTHROPIC_API_KEY)
|
||||
response: Final = gateway.request("POST", "/v1/messages", {**_client_body(stream=True), "model": model})
|
||||
assert response.status_code == 200, response.text
|
||||
assert len(wire.drain()) == 1
|
||||
assert cc.streamed_start_usage(response.text) == {
|
||||
"input_tokens": 100,
|
||||
"cache_read_input_tokens": 400,
|
||||
"cache_creation_input_tokens": 300,
|
||||
"cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200},
|
||||
}, response.text
|
||||
assert cc.streamed_usage(response.text) == {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"cache_read_input_tokens": 400,
|
||||
"cache_creation_input_tokens": 300,
|
||||
"cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200},
|
||||
}, response.text
|
||||
|
||||
|
||||
def test_non_streamed_cache_read_and_write_tokens_reach_the_client_unchanged(gateway: Gateway) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
return Reply(
|
||||
body=cc.message_reply(
|
||||
f"msg_{uuid.uuid4().hex}", _MODEL, ({"type": "text", "text": "PONG"},), _ANTHROPIC_USAGE
|
||||
)
|
||||
)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=cc.ANTHROPIC_API_KEY)
|
||||
response: Final = gateway.request("POST", "/v1/messages", {**_client_body(stream=False), "model": model})
|
||||
assert response.status_code == 200, response.text
|
||||
assert len(wire.drain()) == 1
|
||||
assert response.json()["usage"] == {
|
||||
"input_tokens": 100,
|
||||
"output_tokens": 50,
|
||||
"cache_read_input_tokens": 400,
|
||||
"cache_creation_input_tokens": 300,
|
||||
"cache_creation": {"ephemeral_5m_input_tokens": 100, "ephemeral_1h_input_tokens": 200},
|
||||
}, response.text
|
||||
Loading…
Add table
Reference in a new issue