diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 12d04a08392..333bce0b967 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -1125,23 +1125,42 @@ def bedrock_model_accepts_cache_points(model: str | None) -> bool: ``cachePoint`` blocks. Bedrock rejects requests carrying cachePoint blocks for models without prompt caching support ("You invoked an unsupported model or your request did not allow prompt caching"), so a model whose cost-map entry does not declare - ``supports_prompt_caching`` must not receive them. A model absent from the map - (an application inference profile ARN, a model newer than the map) keeps emitting - so existing caching setups never silently degrade. ``litellm.utils.supports_prompt_caching`` - is not reusable here: it returns False for unmapped models, the opposite polarity. + ``supports_prompt_caching`` must not receive them. An explicit + ``supports_prompt_cache_breakpoint`` on the entry wins over that flag: a model can price + cached tokens through implicit caching yet reject the marker on Converse ("This model + doesn't support the cachePoint field", Kimi K3). The router registers a deployment's + ``model_info`` under ``bedrock/`` as configured, route prefix included, while the + Converse transformation sees the model with ``converse/`` or ``converse_like/`` already + stripped, so every registration form is read. That flag set there covers an application + inference profile ARN or a model newer than the map, while only the map decides whether + a model is known: absent a map entry the model keeps emitting so existing caching setups + never silently degrade. ``litellm.utils.supports_prompt_caching`` is not reusable here: + it returns False for unmapped models, the opposite polarity. """ if model is None: return True if _OPENAI_FAMILY_MODEL_RE.search(model): return False - entries: Final = tuple( - entry - for candidate in (model, get_bedrock_base_model(model)) - if (entry := litellm.model_cost.get(candidate)) is not None + map_keys: Final = (model, get_bedrock_base_model(model)) + registered_keys: Final = tuple(f"bedrock/{route}{model}" for route in ("", "converse/", "converse_like/")) + explicit_marker_support: Final = next( + ( + entry.get("supports_prompt_cache_breakpoint") is True + for key in (*registered_keys, *map_keys) + if (entry := litellm.model_cost.get(key)) is not None + and entry.get("supports_prompt_cache_breakpoint") is not None + ), + None, ) - if not entries: + if explicit_marker_support is not None: + return explicit_marker_support + if not any(key in litellm.model_cost for key in map_keys): return True - return any(entry.get("supports_prompt_caching") is True for entry in entries) + return any( + entry.get("supports_prompt_caching") is True + for key in map_keys + if (entry := litellm.model_cost.get(key)) is not None + ) def bedrock_supports_tool_search(model: str) -> bool: diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 3d2acf4e9d9..fd8074f4744 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -76931,6 +76931,7 @@ "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrock/current/us-east-1/index.json", "supports_audio_input": false, "supports_function_calling": true, + "supports_prompt_cache_breakpoint": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, @@ -76952,6 +76953,7 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, + "supports_prompt_cache_breakpoint": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, @@ -76973,6 +76975,7 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, + "supports_prompt_cache_breakpoint": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 3d2acf4e9d9..fd8074f4744 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -76931,6 +76931,7 @@ "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrock/current/us-east-1/index.json", "supports_audio_input": false, "supports_function_calling": true, + "supports_prompt_cache_breakpoint": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, @@ -76952,6 +76953,7 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, + "supports_prompt_cache_breakpoint": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, @@ -76973,6 +76975,7 @@ "source": "https://aws.amazon.com/bedrock/pricing/", "supports_audio_input": false, "supports_function_calling": true, + "supports_prompt_cache_breakpoint": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, diff --git a/tests/integration/providers/test_bedrock_kimi_k3_cache_point_wire.py b/tests/integration/providers/test_bedrock_kimi_k3_cache_point_wire.py new file mode 100644 index 00000000000..107839ed320 --- /dev/null +++ b/tests/integration/providers/test_bedrock_kimi_k3_cache_point_wire.py @@ -0,0 +1,1572 @@ +import asyncio +import base64 +import json +import os +import re +import signal +import socket +import threading +import time +import uuid +from collections.abc import Callable, Iterable, Iterator, Mapping +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal +from urllib.parse import unquote, urlsplit + +import anthropic +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment +from integration._support.database import read_rows +from integration._support.process import OwnedProxy, owned_proxy_process +from integration._support.upstream import _aws_event_frame +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with + +_KIMI_US: Final = "us.moonshotai.kimi-k3" +_KIMI_GLOBAL: Final = "global.moonshotai.kimi-k3" +_KIMI_BASE: Final = "moonshotai.kimi-k3" +_CLAUDE: Final = "us.anthropic.claude-sonnet-5" +_GPT_OSS: Final = "openai.gpt-oss-120b-1:0" +_LLAMA: Final = "us.meta.llama4-maverick-17b-instruct-v1:0" +_PROFILE_ARN_PREFIX: Final = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/" +_YAML_FLAGGED_ARN: Final = f"{_PROFILE_ARN_PREFIX}yaml-flagged-kimi" +_YAML_PLAIN_ARN: Final = f"{_PROFILE_ARN_PREFIX}yaml-plain-kimi" +_KNOWN_MODEL_IDS: Final = frozenset({_KIMI_US, _KIMI_GLOBAL, _KIMI_BASE, _CLAUDE, _GPT_OSS, _LLAMA}) +_EVENT_STREAM: Final = "application/vnd.amazon.eventstream" +_ANSWER: Final = "kimi cache point control" +_REJECTION: Final = "This model doesn't support the cachePoint field. Remove cachePoint and try again." +_UNKNOWN_MODEL: Final = "The provided model identifier is invalid." +_CACHE_USAGE_MARKER: Final = "[peer:cache-usage]" +_REJECT_MARKER: Final = "[peer:reject]" +_SLOW_MARKER: Final = "[peer:slow]" +_HOLD_MARKER: Final = "[peer:hold]" +_TARGET: Final = re.compile( + r"^/model/(?P.+)/(?Pconverse|converse-stream|invoke|invoke-with-response-stream)$" +) +_CONVERSE_LIKE: Final = "converse_like" +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_SIGNING_KEY: Final = os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt") +_JSON: Final = TypeAdapter(dict[str, JsonValue]) +_EPHEMERAL: Final[dict[str, JsonValue]] = {"type": "ephemeral"} +_AWS: Final[dict[str, JsonValue]] = { + "aws_access_key_id": "AKIASCRIPTEDPROVIDER", + "aws_secret_access_key": "scripted-secret", + "aws_region_name": "us-east-1", +} +_EXTRA: Final[dict[str, JsonValue]] = {"num_retries": 0, "cache": {"no-cache": True}} +_SYSTEM_TEXT: Final = "You are terse." +_TOOL_PARAMETERS: Final[dict[str, JsonValue]] = { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], +} + +_NAMES: Final = MappingProxyType( + { + "kimi-us": f"bedrock/{_KIMI_US}", + "kimi-global-converse": f"bedrock/converse/{_KIMI_GLOBAL}", + "kimi-base": f"bedrock/{_KIMI_BASE}", + "kimi-regional": f"bedrock/us-east-1/{_KIMI_US}", + "claude-control": f"bedrock/{_CLAUDE}", + "gpt-oss-control": f"bedrock/{_GPT_OSS}", + "llama-control": f"bedrock/{_LLAMA}", + "arn-yaml-flagged": f"bedrock/{_YAML_FLAGGED_ARN}", + "arn-yaml-plain": f"bedrock/{_YAML_PLAIN_ARN}", + "kimi-inject-message": f"bedrock/{_KIMI_US}", + "kimi-inject-tool": f"bedrock/{_KIMI_US}", + "claude-inject-message": f"bedrock/{_CLAUDE}", + "bedrock/*": "bedrock/*", + } +) +_MODEL_INFO: Final = MappingProxyType({"arn-yaml-flagged": {"supports_prompt_cache_breakpoint": False}}) +_INJECTION: Final = MappingProxyType( + { + "kimi-inject-message": [{"location": "message", "role": "system"}], + "claude-inject-message": [{"location": "message", "role": "system"}], + "kimi-inject-tool": [{"location": "tool_config"}], + } +) + +Endpoint = Literal["chat", "messages", "responses"] +Marker = Literal["system", "user", "tool"] +_ALL_MARKERS: Final = frozenset[Marker]({"system", "user", "tool"}) +_USER_ONLY: Final = frozenset[Marker]({"user"}) +_SYSTEM_AND_USER: Final = frozenset[Marker]({"system", "user"}) +_NO_MARKERS: Final = frozenset[Marker]() + + +def _frame(event_type: str, payload: Mapping[str, JsonValue]) -> bytes: + return _aws_event_frame(event_type, payload, "sc", "u") + + +def _usage(cache_read: int = 0, cache_write: int = 0) -> dict[str, JsonValue]: + cached: Final[dict[str, JsonValue]] = { + **({"cacheReadInputTokens": cache_read} if cache_read else {}), + **({"cacheWriteInputTokens": cache_write} if cache_write else {}), + } + return {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15 + cache_read + cache_write, **cached} + + +def _converse_reply(usage: Mapping[str, JsonValue]) -> bytes: + return json.dumps( + { + "output": {"message": {"role": "assistant", "content": [{"text": _ANSWER}]}}, + "stopReason": "end_turn", + "usage": dict(usage), + "metrics": {"latencyMs": 1}, + } + ).encode() + + +def _text_parts(pieces: int) -> tuple[str, ...]: + words: Final = _ANSWER.split(" ") + assert pieces in (1, len(words)), pieces + if pieces == 1: + return (_ANSWER,) + return tuple(word if index == len(words) - 1 else f"{word} " for index, word in enumerate(words)) + + +def _stream_frames(usage: Mapping[str, JsonValue], pieces: int = 1) -> tuple[bytes, ...]: + deltas: Final = tuple( + _frame("contentBlockDelta", {"delta": {"text": part}, "contentBlockIndex": 0}) for part in _text_parts(pieces) + ) + return ( + _frame("messageStart", {"role": "assistant"}), + *deltas, + _frame("contentBlockStop", {"contentBlockIndex": 0}), + _frame("messageStop", {"stopReason": "end_turn"}), + _frame("metadata", {"usage": dict(usage)}), + ) + + +def _known(model: str) -> bool: + return model in _KNOWN_MODEL_IDS or "application-inference-profile" in model + + +def _error(message: str) -> bytes: + return json.dumps({"message": message}).encode() + + +def _anthropic_message(usage: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return { + "id": f"msg_{uuid.uuid4().hex}", + "type": "message", + "role": "assistant", + "model": _CLAUDE, + "content": [{"type": "text", "text": _ANSWER}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": usage["inputTokens"], "output_tokens": usage["outputTokens"]}, + } + + +def _invoke_chunk(event: Mapping[str, JsonValue]) -> bytes: + encoded: Final = base64.b64encode(json.dumps(event, separators=(",", ":")).encode()).decode() + return _frame("chunk", {"bytes": encoded}) + + +def _invoke_stream(usage: Mapping[str, JsonValue]) -> bytes: + message: Final = { + **_anthropic_message(usage), + "content": [], + "usage": {"input_tokens": usage["inputTokens"], "output_tokens": 1}, + } + events: Final[tuple[dict[str, JsonValue], ...]] = ( + {"type": "message_start", "message": message}, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _ANSWER}}, + {"type": "content_block_stop", "index": 0}, + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": usage["outputTokens"]}, + }, + {"type": "message_stop"}, + ) + return b"".join(_invoke_chunk(event) for event in events) + + +@dataclass(frozen=True, slots=True) +class _Target: + model: str + action: str + + +def _target(raw: str) -> _Target: + path: Final = unquote(raw) + if path == "/": + return _Target(_CONVERSE_LIKE, "converse") + found: Final = _TARGET.match(path) + assert found is not None, raw + return _Target(found["model"], found["action"]) + + +def _peer(hold: threading.Event | None = None, held: SimpleQueue[str] | None = None) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + target: Final = _target(request.target) + streaming: Final = target.action in ("converse-stream", "invoke-with-response-stream") + text: Final = request.body.decode(errors="replace") + if target.model != _CONVERSE_LIKE and not _known(target.model): + return Reply(status=400, body=_error(_UNKNOWN_MODEL)) + if _REJECT_MARKER in text and not streaming: + return Reply(status=400, body=_error(_REJECTION)) + if _REJECT_MARKER in text: + return Reply(body=_frame("validationException", {"message": _REJECTION}), content_type=_EVENT_STREAM) + usage: Final = _usage(100, 50) if _CACHE_USAGE_MARKER in text else _usage() + if target.action == "invoke": + return Reply(body=json.dumps(_anthropic_message(usage)).encode()) + if target.action == "invoke-with-response-stream": + return Reply(body=_invoke_stream(usage), content_type=_EVENT_STREAM) + if not streaming: + return Reply(body=_converse_reply(usage)) + if _HOLD_MARKER in text and hold is not None: + if held is not None: + held.put(target.model) + return Reply(content_type=_EVENT_STREAM, chunks=_stream_frames(usage), gate_after_first=hold) + if _SLOW_MARKER in text: + return Reply(content_type=_EVENT_STREAM, chunks=_stream_frames(usage, 4), pause_between_chunks=0.2) + return Reply(body=b"".join(_stream_frames(usage)), content_type=_EVENT_STREAM) + + return respond + + +@dataclass(frozen=True, slots=True) +class _Received: + model: str + action: str + streaming: bool + body: dict[str, JsonValue] + cache_points: int + cache_controls: int + + +def _message_blocks(messages: JsonValue) -> Iterator[JsonValue]: + if not isinstance(messages, list): + return + for message in messages: + content: Final = message.get("content") if isinstance(message, dict) else None + if isinstance(content, list): + yield from content + + +def _tool_blocks(body: Mapping[str, JsonValue]) -> tuple[JsonValue, ...]: + tool_config: Final = body.get("toolConfig") + tools: Final = tool_config.get("tools") if isinstance(tool_config, dict) else None + return tuple(tools) if isinstance(tools, list) else () + + +def _cache_points(body: Mapping[str, JsonValue]) -> int: + system: Final = body.get("system") + blocks: Final = ( + *(system if isinstance(system, list) else ()), + *_message_blocks(body.get("messages")), + *_tool_blocks(body), + ) + return sum(1 for block in blocks if isinstance(block, dict) and "cachePoint" in block) + + +def _cache_controls(value: JsonValue) -> int: + if isinstance(value, dict): + return sum(_cache_controls(item) for item in value.values()) + (1 if "cache_control" in value else 0) + if isinstance(value, list): + return sum(_cache_controls(item) for item in value) + return 0 + + +def _parse(request: Request) -> _Received: + target: Final = _target(request.target) + body: Final = _JSON.validate_python(json.loads(request.body)) + streaming: Final = target.action in ("converse-stream", "invoke-with-response-stream") + return _Received(target.model, target.action, streaming, body, _cache_points(body), _cache_controls(body)) + + +def _received(wire: Wire) -> tuple[_Received, ...]: + return tuple(_parse(request) for request in wire.drain()) + + +def _only_received(wire: Wire) -> _Received: + (received,) = _received(wire) + return received + + +def _prompt(*markers: str) -> str: + return " ".join((f"Reply with the control sentence {uuid.uuid4().hex}", *markers)) + + +def _text_block(text: str, cached: bool, kind: str = "text") -> dict[str, JsonValue]: + return {"type": kind, "text": text, **({"cache_control": dict(_EPHEMERAL)} if cached else {})} + + +def _chat_tools(cached: bool) -> dict[str, JsonValue]: + return { + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Weather for a city", + "parameters": _TOOL_PARAMETERS, + }, + **({"cache_control": dict(_EPHEMERAL)} if cached else {}), + } + ], + } + + +def _anthropic_tools(cached: bool) -> dict[str, JsonValue]: + return { + "tools": [ + { + "name": "get_weather", + "description": "Weather for a city", + "input_schema": _TOOL_PARAMETERS, + **({"cache_control": dict(_EPHEMERAL)} if cached else {}), + } + ] + } + + +def _chat_body( + model: str, prompt: str, markers: frozenset[Marker], *, stream: bool = False, with_tools: bool = False +) -> dict[str, JsonValue]: + return { + "model": model, + "max_tokens": 16, + "stream": stream, + "messages": [ + {"role": "system", "content": [_text_block(_SYSTEM_TEXT, "system" in markers)]}, + {"role": "user", "content": [_text_block(prompt, "user" in markers)]}, + ], + **(_chat_tools("tool" in markers) if with_tools or "tool" in markers else {}), + **_EXTRA, + } + + +def _messages_body( + model: str, prompt: str, markers: frozenset[Marker], *, stream: bool = False, with_tools: bool = False +) -> dict[str, JsonValue]: + return { + "model": model, + "max_tokens": 16, + "stream": stream, + "system": [_text_block(_SYSTEM_TEXT, "system" in markers)], + "messages": [{"role": "user", "content": [_text_block(prompt, "user" in markers)]}], + **(_anthropic_tools("tool" in markers) if with_tools or "tool" in markers else {}), + **_EXTRA, + } + + +def _responses_body( + model: str, prompt: str, markers: frozenset[Marker], *, stream: bool = False +) -> dict[str, JsonValue]: + return { + "model": model, + "max_output_tokens": 16, + "stream": stream, + "input": [{"role": "user", "content": [_text_block(prompt, "user" in markers, "input_text")]}], + **_EXTRA, + } + + +def _body( + endpoint: Endpoint, model: str, prompt: str, markers: frozenset[Marker], *, stream: bool = False +) -> dict[str, JsonValue]: + match endpoint: + case "chat": + return _chat_body(model, prompt, markers, stream=stream) + case "messages": + return _messages_body(model, prompt, markers, stream=stream) + case "responses": + return _responses_body(model, prompt, markers, stream=stream) + + +def _path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "messages": + return "/v1/messages" + case "responses": + return "/v1/responses" + + +@dataclass(frozen=True, slots=True) +class _Outcome: + status: int + call_id: str + response_id: str + text: str + headers: Mapping[str, str] + raw: str + + +def _sse_payloads(lines: Iterable[str]) -> tuple[dict[str, JsonValue], ...]: + return tuple(json.loads(line[6:]) for line in lines if line.startswith("data: ") and line != "data: [DONE]") + + +def _first_choice(chunk: Mapping[str, JsonValue]) -> dict[str, JsonValue] | None: + choices: Final = chunk.get("choices") + return _JSON.validate_python(choices[0]) if isinstance(choices, list) and choices else None + + +def _chat_stream_text(chunks: Iterable[dict[str, JsonValue]]) -> str: + choices: Final = tuple(choice for choice in map(_first_choice, chunks) if choice is not None) + return "".join(str(_JSON.validate_python(choice["delta"]).get("content") or "") for choice in choices) + + +def _chat_stream_id(chunks: Iterable[dict[str, JsonValue]]) -> str: + (identity,) = {str(chunk["id"]) for chunk in chunks if "id" in chunk} + return identity + + +def _message_stream_id(payloads: Iterable[dict[str, JsonValue]]) -> str: + (started,) = tuple(payload for payload in payloads if payload.get("type") == "message_start") + return str(_JSON.validate_python(started["message"])["id"]) + + +def _messages_stream_text(payloads: Iterable[dict[str, JsonValue]]) -> str: + deltas: Final = tuple(payload for payload in payloads if payload.get("type") == "content_block_delta") + return "".join(str(_JSON.validate_python(delta["delta"]).get("text") or "") for delta in deltas) + + +def _responses_stream_text(events: Iterable[dict[str, JsonValue]]) -> str: + return "".join( + str(event.get("delta") or "") for event in events if event.get("type") == "response.output_text.delta" + ) + + +def _inner_response_id(identity: str) -> str: + managed: Final = decrypt_if_encrypted_with(identity.removeprefix("resp_"), _SIGNING_KEY) + assert managed is not None, identity + issued: Final = managed.split(";", 1)[0].rsplit("response_id:", 1)[1] + decoded: Final = base64.b64decode(issued.removeprefix("resp_")).decode() + return decoded.rsplit("response_id:", 1)[1] + + +def _completed_response_id(events: Iterable[dict[str, JsonValue]]) -> str: + (completed,) = tuple(event for event in events if event.get("type") == "response.completed") + return _inner_response_id(str(_JSON.validate_python(completed["response"])["id"])) + + +def _chat_text(body: Mapping[str, JsonValue]) -> str: + choices: Final = body.get("choices") + assert isinstance(choices, list) and choices, body + return str(_JSON.validate_python(_JSON.validate_python(choices[0])["message"]).get("content") or "") + + +def _messages_text(body: Mapping[str, JsonValue]) -> str: + content: Final = body.get("content") + assert isinstance(content, list), body + return "".join(str(block.get("text") or "") for block in content if isinstance(block, dict)) + + +def _responses_text(body: Mapping[str, JsonValue]) -> str: + output: Final = body.get("output") + assert isinstance(output, list), body + return "".join( + str(part.get("text") or "") + for item in output + if isinstance(item, dict) + for part in ( + item.get("content") if isinstance(item.get("content"), list) else () + ) # comprehension-ok: nested response items + if isinstance(part, dict) + ) + + +def _outcome_of(endpoint: Endpoint, stream: bool, response: httpx.Response, lines: tuple[str, ...]) -> _Outcome: + call_id: Final = response.headers.get("x-litellm-call-id", "") + raw: Final = "\n".join(lines) + if response.status_code != 200: + return _Outcome(response.status_code, call_id, "", "", response.headers, raw) + if stream: + payloads: Final = _sse_payloads(lines) + match endpoint: + case "chat": + return _Outcome( + 200, call_id, _chat_stream_id(payloads), _chat_stream_text(payloads), response.headers, raw + ) + case "messages": + return _Outcome( + 200, call_id, _message_stream_id(payloads), _messages_stream_text(payloads), response.headers, raw + ) + case "responses": + return _Outcome( + 200, + call_id, + _completed_response_id(payloads), + _responses_stream_text(payloads), + response.headers, + raw, + ) + body: Final = _JSON.validate_python(json.loads(raw)) + match endpoint: + case "chat": + return _Outcome(200, call_id, str(body["id"]), _chat_text(body), response.headers, raw) + case "messages": + return _Outcome(200, call_id, str(body["id"]), _messages_text(body), response.headers, raw) + case "responses": + return _Outcome( + 200, call_id, _inner_response_id(str(body["id"])), _responses_text(body), response.headers, raw + ) + + +def _send(gateway: Gateway, endpoint: Endpoint, body: Mapping[str, JsonValue], *, key: str | None = None) -> _Outcome: + stream: Final = body.get("stream") is True + headers: Final = {"Authorization": f"Bearer {gateway.key if key is None else key}"} + with gateway.client.stream("POST", _path(endpoint), json=body, headers=headers, timeout=60) as response: + lines: Final = tuple(line for line in response.iter_lines() if line) + return _outcome_of(endpoint, stream, response, lines) + + +def _spend_rows(request_ids: frozenset[str], *, expected: int, seconds: float = 90) -> tuple[dict[str, JsonValue], ...]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, status, model, spend FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(string_to_array(%s, %s))', + (",".join(sorted(request_ids)), ","), + ), + lambda found: len(found) >= expected, + seconds=seconds, + ) + return tuple(rows) + + +def _success_row(request_id: str) -> dict[str, JsonValue]: + assert request_id, "No id to look the spend row up by" + (row,) = _spend_rows(frozenset({request_id}), expected=1) + assert row["status"] == "success", row + return row + + +def _failure_row(call_id: str) -> dict[str, JsonValue]: + assert call_id, "No call id to look the failure row up by" + (row,) = _spend_rows(frozenset({call_id}), expected=1) + assert row["status"] == "failure", row + return row + + +def _assert_answered(outcome: _Outcome) -> None: + assert outcome.status == 200, (outcome.status, outcome.raw) + assert outcome.text == _ANSWER, outcome.raw + + +def _assert_cache_points(received: _Received, emits: bool, model: str) -> None: + assert received.model == model, (received.model, model) + assert (received.cache_points > 0) == emits, (emits, received.body) + + +def _head_strips() -> bool: + return os.environ.get("INTEGRATION_LEG", "head") == "head" + + +def _owned_config(wire: Wire, directory: Path, *, names: Iterable[str] = tuple(_NAMES)) -> Path: + base: Final = _JSON.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + config: Final[dict[str, JsonValue]] = { + **base, + "model_list": [ + { + "model_name": name, + "litellm_params": { + "model": _NAMES[name], + "api_base": wire.url, + "num_retries": 0, + **_AWS, + **({"cache_control_injection_points": _INJECTION[name]} if name in _INJECTION else {}), + }, + **({"model_info": dict(_MODEL_INFO[name])} if name in _MODEL_INFO else {}), + } + for name in names + ], + "router_settings": {**_JSON.validate_python(base["router_settings"]), "num_retries": 0}, + } + path: Final = directory / f"kimi-k3-cache-point-{uuid.uuid4().hex}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@dataclass(frozen=True, slots=True) +class _Rig: + gateway: Gateway + wire: Wire + owned: OwnedProxy + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]: + directory: Final = tmp_path_factory.mktemp("kimi-k3-cache-point") + with gateway_from_environment() as environment, wire_server(_peer()) as wire: + config: Final = _owned_config(wire, directory) + with owned_proxy_process(environment, directory, {}, config=config, workers=2) as owned: + eventually( + lambda: len(_STARTED_WORKER.findall(owned.log.read_text())), lambda count: count == 2, seconds=60 + ) + wire.drain() + yield _Rig(owned.gateway, wire, owned) + + +def _observe(rig: _Rig, endpoint: Endpoint, body: Mapping[str, JsonValue]) -> tuple[_Outcome, _Received]: + rig.wire.drain() + outcome: Final = _send(rig.gateway, endpoint, body) + return outcome, _only_received(rig.wire) + + +_KIMI_FORMS: Final = ("kimi-us", "kimi-global-converse", "kimi-base", "kimi-regional", "arn-yaml-flagged") +_EMITTING_CONTROLS: Final = ("claude-control", "arn-yaml-plain") +_SILENT_CONTROLS: Final = ("gpt-oss-control", "llama-control") + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("name", _KIMI_FORMS, ids=tuple(f"m-{name}" for name in _KIMI_FORMS)) +def test_m01_to_m05_every_kimi_k3_form_sends_converse_without_a_cache_point(rig: _Rig, name: str) -> None: + outcome, received = _observe(rig, "chat", _chat_body(name, _prompt(), _ALL_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 0, received.body + assert received.body.get("toolConfig") is not None, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("name", _EMITTING_CONTROLS, ids=tuple(f"m-{name}" for name in _EMITTING_CONTROLS)) +def test_m08_m09_models_that_take_cache_points_still_get_all_three(rig: _Rig, name: str) -> None: + outcome, received = _observe(rig, "chat", _chat_body(name, _prompt(), _ALL_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 3, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("name", _SILENT_CONTROLS, ids=tuple(f"m-{name}" for name in _SILENT_CONTROLS)) +def test_m10_m11_models_without_caching_never_got_cache_points(rig: _Rig, name: str) -> None: + outcome, received = _observe(rig, "chat", _chat_body(name, _prompt(), _ALL_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 0, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_m14_a_wildcard_deployment_resolves_the_caller_model_before_deciding(rig: _Rig) -> None: + kimi, kimi_received = _observe(rig, "chat", _chat_body(f"bedrock/{_KIMI_US}", _prompt(), _SYSTEM_AND_USER)) + _assert_answered(kimi) + assert kimi_received.model == _KIMI_US, kimi_received.model + assert kimi_received.cache_points == 0, kimi_received.body + claude, claude_received = _observe(rig, "chat", _chat_body(f"bedrock/{_CLAUDE}", _prompt(), _SYSTEM_AND_USER)) + _assert_answered(claude) + assert claude_received.model == _CLAUDE, claude_received.model + assert claude_received.cache_points == 2, claude_received.body + _success_row(kimi.response_id) + _success_row(claude.response_id) + + +_RAW_CELLS: Final = ( + ("chat", False), + ("chat", True), + ("messages", False), + ("messages", True), + ("responses", False), + ("responses", True), +) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize( + ("endpoint", "stream"), + _RAW_CELLS, + ids=tuple(f"e-{endpoint}-{'stream' if stream else 'plain'}" for endpoint, stream in _RAW_CELLS), +) +def test_e01_to_e06_every_endpoint_reaches_kimi_without_a_cache_point( + rig: _Rig, endpoint: Endpoint, stream: bool +) -> None: + outcome, received = _observe(rig, endpoint, _body(endpoint, "kimi-us", _prompt(), _SYSTEM_AND_USER, stream=stream)) + _assert_answered(outcome) + assert received.streaming == stream, received + assert received.model == _KIMI_US, received.model + assert received.cache_points == 0, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize( + ("endpoint", "stream"), + _RAW_CELLS, + ids=tuple(f"e-{endpoint}-{'stream' if stream else 'plain'}" for endpoint, stream in _RAW_CELLS), +) +def test_e07_to_e12_every_endpoint_still_sends_claude_its_cache_points( + rig: _Rig, endpoint: Endpoint, stream: bool +) -> None: + outcome, received = _observe( + rig, endpoint, _body(endpoint, "claude-control", _prompt(), _SYSTEM_AND_USER, stream=stream) + ) + _assert_answered(outcome) + assert received.model == _CLAUDE, received.model + assert received.streaming == stream, received + if endpoint == "messages": + assert received.action.startswith("invoke"), received.action + assert received.cache_controls == 2 and received.cache_points == 0, received.body + else: + assert received.action.startswith("converse"), received.action + expected: Final = 1 if endpoint == "responses" else 2 + assert received.cache_points == expected, (endpoint, received.body) + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize( + ("endpoint", "stream"), + (("messages", False), ("messages", True)), + ids=("e-messages-arn-plain", "e-messages-arn-stream"), +) +def test_e13_e14_messages_keeps_cache_control_for_an_arn_and_the_flag_decides( + rig: _Rig, endpoint: Endpoint, stream: bool +) -> None: + flagged, flagged_received = _observe( + rig, endpoint, _body(endpoint, "arn-yaml-flagged", _prompt(), _SYSTEM_AND_USER, stream=stream) + ) + _assert_answered(flagged) + assert flagged_received.model == _YAML_FLAGGED_ARN, flagged_received.model + assert flagged_received.cache_points == 0, flagged_received.body + plain, plain_received = _observe( + rig, endpoint, _body(endpoint, "arn-yaml-plain", _prompt(), _SYSTEM_AND_USER, stream=stream) + ) + _assert_answered(plain) + assert plain_received.model == _YAML_PLAIN_ARN, plain_received.model + assert plain_received.cache_points == 2, plain_received.body + _success_row(flagged.response_id) + _success_row(plain.response_id) + + +def _openai_client(rig: _Rig) -> openai.OpenAI: + return openai.OpenAI(base_url=f"{_proxy_url(rig.gateway)}/v1", api_key=rig.gateway.key, max_retries=0) + + +def _async_openai_client(rig: _Rig) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI(base_url=f"{_proxy_url(rig.gateway)}/v1", api_key=rig.gateway.key, max_retries=0) + + +def _anthropic_client(rig: _Rig) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=_proxy_url(rig.gateway), api_key=rig.gateway.key, max_retries=0) + + +def _async_anthropic_client(rig: _Rig) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic(base_url=_proxy_url(rig.gateway), api_key=rig.gateway.key, max_retries=0) + + +def _proxy_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def _sdk_messages(prompt: str) -> list[dict[str, JsonValue]]: + return [ + {"role": "system", "content": [_text_block(_SYSTEM_TEXT, True)]}, + {"role": "user", "content": [_text_block(prompt, True)]}, + ] + + +@pytest.mark.timeout(600) +def test_k01_openai_sdk_sync_chat_reaches_kimi_without_a_cache_point(rig: _Rig) -> None: + rig.wire.drain() + completion: Final = _openai_client(rig).chat.completions.create( + model="kimi-us", + messages=_sdk_messages(_prompt()), # pyright: ignore[reportArgumentType] # cache_control rides as an extra key + max_tokens=16, + extra_body=_EXTRA, + ) + assert completion.choices[0].message.content == _ANSWER, completion + received: Final = _only_received(rig.wire) + assert received.model == _KIMI_US and not received.streaming, received + assert received.cache_points == 0, received.body + _success_row(completion.id) + + +@pytest.mark.timeout(600) +async def test_k02_openai_sdk_async_chat_stream_reaches_kimi_without_a_cache_point(rig: _Rig) -> None: + rig.wire.drain() + stream: Final = await _async_openai_client(rig).chat.completions.create( + model="kimi-us", + messages=_sdk_messages(_prompt()), # pyright: ignore[reportArgumentType] # cache_control rides as an extra key + max_tokens=16, + stream=True, + extra_body=_EXTRA, + ) + chunks: Final = [chunk async for chunk in stream] + text: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + assert text == _ANSWER, chunks + (identity,) = {chunk.id for chunk in chunks} + received: Final = _only_received(rig.wire) + assert received.model == _KIMI_US and received.streaming, received + assert received.cache_points == 0, received.body + _success_row(identity) + + +_ANTHROPIC_SDK_TARGETS: Final = (("kimi-us", _KIMI_US), ("arn-yaml-flagged", _YAML_FLAGGED_ARN)) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize(("name", "model_id"), _ANTHROPIC_SDK_TARGETS, ids=("k03-kimi", "k03-flagged-arn")) +def test_k03_anthropic_sdk_sync_messages_reaches_the_model_without_a_cache_point( + rig: _Rig, name: str, model_id: str +) -> None: + rig.wire.drain() + message: Final = _anthropic_client(rig).messages.create( + model=name, + max_tokens=16, + system=[{"type": "text", "text": _SYSTEM_TEXT, "cache_control": {"type": "ephemeral"}}], + messages=[ + {"role": "user", "content": [{"type": "text", "text": _prompt(), "cache_control": {"type": "ephemeral"}}]} + ], + extra_body=_EXTRA, + ) + assert "".join(block.text for block in message.content if block.type == "text") == _ANSWER, message + received: Final = _only_received(rig.wire) + assert received.model == model_id and not received.streaming, received + assert received.cache_points == 0, received.body + _success_row(message.id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize(("name", "model_id"), _ANTHROPIC_SDK_TARGETS, ids=("k04-kimi", "k04-flagged-arn")) +async def test_k04_anthropic_sdk_async_stream_reaches_the_model_without_a_cache_point( + rig: _Rig, name: str, model_id: str +) -> None: + rig.wire.drain() + async with _async_anthropic_client(rig).messages.stream( + model=name, + max_tokens=16, + system=[{"type": "text", "text": _SYSTEM_TEXT, "cache_control": {"type": "ephemeral"}}], + messages=[ + {"role": "user", "content": [{"type": "text", "text": _prompt(), "cache_control": {"type": "ephemeral"}}]} + ], + extra_body=_EXTRA, + ) as stream: + final: Final = await stream.get_final_message() + assert "".join(block.text for block in final.content if block.type == "text") == _ANSWER, final + received: Final = _only_received(rig.wire) + assert received.model == model_id and received.streaming, received + assert received.cache_points == 0, received.body + _success_row(final.id) + + +@pytest.mark.timeout(600) +def test_k05_openai_sdk_sync_responses_reaches_kimi_without_a_cache_point(rig: _Rig) -> None: + rig.wire.drain() + response: Final = _openai_client(rig).responses.create( + model="kimi-us", + input=[ + { + "role": "user", + "content": [{"type": "input_text", "text": _prompt(), "cache_control": {"type": "ephemeral"}}], + } + ], # pyright: ignore[reportArgumentType] # cache_control rides as an extra key + max_output_tokens=16, + extra_body=_EXTRA, + ) + assert response.output_text == _ANSWER, response + received: Final = _only_received(rig.wire) + assert received.model == _KIMI_US and not received.streaming, received + assert received.cache_points == 0, received.body + _success_row(_inner_response_id(response.id)) + + +@pytest.mark.timeout(600) +async def test_k06_openai_sdk_async_responses_stream_reaches_kimi_without_a_cache_point(rig: _Rig) -> None: + rig.wire.drain() + stream: Final = await _async_openai_client(rig).responses.create( + model="kimi-us", + input=[ + { + "role": "user", + "content": [{"type": "input_text", "text": _prompt(), "cache_control": {"type": "ephemeral"}}], + } + ], # pyright: ignore[reportArgumentType] # cache_control rides as an extra key + max_output_tokens=16, + stream=True, + extra_body=_EXTRA, + ) + completed: Final = [event.response async for event in stream if event.type == "response.completed"] + assert len(completed) == 1 and completed[0].output_text == _ANSWER, completed + received: Final = _only_received(rig.wire) + assert received.model == _KIMI_US and received.streaming, received + assert received.cache_points == 0, received.body + _success_row(_inner_response_id(completed[0].id)) + + +_LOCATIONS: Final = ( + ("system", frozenset[Marker]({"system"})), + ("user", frozenset[Marker]({"user"})), + ("tool", frozenset[Marker]({"tool"})), +) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize(("location", "markers"), _LOCATIONS, ids=tuple(f"l-{location}" for location, _ in _LOCATIONS)) +def test_l01_to_l03_each_marker_location_is_dropped_for_kimi_and_kept_for_claude( + rig: _Rig, location: str, markers: frozenset[Marker] +) -> None: + kimi, kimi_received = _observe(rig, "chat", _chat_body("kimi-us", _prompt(), markers, with_tools=True)) + _assert_answered(kimi) + assert kimi_received.cache_points == 0, (location, kimi_received.body) + claude, claude_received = _observe(rig, "chat", _chat_body("claude-control", _prompt(), markers, with_tools=True)) + _assert_answered(claude) + assert claude_received.cache_points == 1, (location, claude_received.body) + _success_row(kimi.response_id) + _success_row(claude.response_id) + + +@pytest.mark.timeout(600) +def test_l04_l05_gateway_injection_points_are_dropped_for_kimi_and_kept_for_claude(rig: _Rig) -> None: + kimi_message, kimi_message_received = _observe( + rig, "chat", _chat_body("kimi-inject-message", _prompt(), _NO_MARKERS, with_tools=True) + ) + _assert_answered(kimi_message) + assert kimi_message_received.cache_points == 0, kimi_message_received.body + kimi_tool, kimi_tool_received = _observe( + rig, "chat", _chat_body("kimi-inject-tool", _prompt(), _NO_MARKERS, with_tools=True) + ) + _assert_answered(kimi_tool) + assert kimi_tool_received.cache_points == 0, kimi_tool_received.body + claude, claude_received = _observe( + rig, "chat", _chat_body("claude-inject-message", _prompt(), _NO_MARKERS, with_tools=True) + ) + _assert_answered(claude) + assert claude_received.cache_points == 1, claude_received.body + system: Final = claude_received.body.get("system") + assert isinstance(system, list) and "cachePoint" in _JSON.validate_python(system[-1]), claude_received.body + for outcome in (kimi_message, kimi_tool, claude): + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_l06_a_request_without_markers_was_never_touched(rig: _Rig) -> None: + for name in ("kimi-us", "claude-control"): + outcome, received = _observe(rig, "chat", _chat_body(name, _prompt(), _NO_MARKERS, with_tools=True)) + _assert_answered(outcome) + assert received.cache_points == 0, (name, received.body) + _success_row(outcome.response_id) + + +def _rates(gateway: Gateway, name: str) -> dict[str, float]: + entries: Final = gateway.get("/model/info")["data"] + assert isinstance(entries, list), entries + (entry,) = tuple( + _JSON.validate_python(item) for item in entries if _JSON.validate_python(item)["model_name"] == name + ) + info: Final = _JSON.validate_python(entry["model_info"]) + return { + key: float(str(info[key])) + for key in ( + "input_cost_per_token", + "output_cost_per_token", + "cache_read_input_token_cost", + "cache_creation_input_token_cost", + ) + } + + +@pytest.mark.timeout(600) +def test_p01_kimi_cache_usage_is_still_priced_from_its_own_row(rig: _Rig) -> None: + outcome, received = _observe(rig, "chat", _chat_body("kimi-us", _prompt(_CACHE_USAGE_MARKER), _NO_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 0, received.body + body: Final = _JSON.validate_python(json.loads(outcome.raw)) + usage: Final = _JSON.validate_python(body["usage"]) + assert usage["prompt_tokens"] == 161 and usage["completion_tokens"] == 4, usage + details: Final = _JSON.validate_python(usage["prompt_tokens_details"]) + assert details["cached_tokens"] == 100, usage + assert usage["cache_creation_input_tokens"] == 50, usage + rates: Final = _rates(rig.gateway, "kimi-us") + expected: Final = ( + 11 * rates["input_cost_per_token"] + + 100 * rates["cache_read_input_token_cost"] + + 50 * rates["cache_creation_input_token_cost"] + + 4 * rates["output_cost_per_token"] + ) + assert abs(float(outcome.headers["x-litellm-response-cost"]) - expected) < 1e-12, (outcome.headers, rates) + row: Final = _success_row(outcome.response_id) + assert abs(float(str(row["spend"])) - expected) < 1e-12, (row, rates) + + +@pytest.mark.timeout(600) +def test_p02_a_plain_kimi_reply_is_priced_from_its_own_row(rig: _Rig) -> None: + outcome, received = _observe(rig, "chat", _chat_body("kimi-us", _prompt(), _NO_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 0, received.body + rates: Final = _rates(rig.gateway, "kimi-us") + expected: Final = 11 * rates["input_cost_per_token"] + 4 * rates["output_cost_per_token"] + assert abs(float(outcome.headers["x-litellm-response-cost"]) - expected) < 1e-12, (outcome.headers, rates) + row: Final = _success_row(outcome.response_id) + assert abs(float(str(row["spend"])) - expected) < 1e-12, (row, rates) + + +@pytest.mark.timeout(600) +def test_s07_a_5kb_caller_model_under_the_wildcard_is_a_bounded_4xx_and_liveliness_stays_up(rig: _Rig) -> None: + rig.wire.drain() + started: Final = time.monotonic() + outcome: Final = _send(rig.gateway, "chat", _chat_body(f"bedrock/{'a' * 5120}", _prompt(), _SYSTEM_AND_USER)) + elapsed: Final = time.monotonic() - started + liveliness: Final = rig.gateway.client.get("/health/liveliness", timeout=5) + assert liveliness.status_code == 200, liveliness.text + assert outcome.status in (400, 404), (outcome.status, outcome.raw[:300]) + error: Final = _JSON.validate_python(_JSON.validate_python(json.loads(outcome.raw))["error"]) + assert "a" * 5120 in str(error["message"]), outcome.raw[:300] + assert elapsed < 10, elapsed + assert rig.wire.drain() == () + _failure_row(outcome.call_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("stream", (False, True), ids=("s08-plain", "s08-stream")) +def test_s08_a_peer_rejection_on_kimi_is_a_400_with_the_message_and_a_failure_row(rig: _Rig, stream: bool) -> None: + rig.wire.drain() + outcome: Final = _send( + rig.gateway, "chat", _chat_body("kimi-us", _prompt(_REJECT_MARKER), _NO_MARKERS, stream=stream) + ) + received: Final = _only_received(rig.wire) + assert received.cache_points == 0, received.body + assert outcome.status == 400, (outcome.status, outcome.raw) + assert _REJECTION in outcome.raw.replace('\\"', '"'), outcome.raw + _failure_row(outcome.call_id) + + +@pytest.mark.timeout(600) +def test_s09_an_unauthenticated_kimi_request_never_reaches_the_peer(rig: _Rig) -> None: + rig.wire.drain() + outcome: Final = _send( + rig.gateway, "chat", _chat_body("kimi-us", _prompt(), _SYSTEM_AND_USER), key="sk-integration-bogus" + ) + assert outcome.status == 401, outcome.raw + assert rig.wire.drain() == () + + +@pytest.mark.timeout(600) +def test_x02_two_identical_uncached_requests_are_two_peer_calls_and_two_rows(rig: _Rig) -> None: + body: Final = _chat_body("kimi-us", _prompt(), _NO_MARKERS) + rig.wire.drain() + first: Final = _send(rig.gateway, "chat", body) + second: Final = _send(rig.gateway, "chat", body) + _assert_answered(first) + _assert_answered(second) + assert first.response_id != second.response_id, (first.response_id, second.response_id) + received: Final = _received(rig.wire) + assert len(received) == 2 and all(item.cache_points == 0 for item in received), received + rows: Final = _spend_rows(frozenset({first.response_id, second.response_id}), expected=2) + assert {str(row["request_id"]) for row in rows} == {first.response_id, second.response_id}, rows + + +@pytest.mark.timeout(600) +def test_x03_a_response_cache_hit_still_answers_after_fewer_peer_calls_than_sends(rig: _Rig) -> None: + body: Final = {key: value for key, value in _chat_body("kimi-us", _prompt(), _NO_MARKERS).items() if key != "cache"} + rig.wire.drain() + first: Final = _send(rig.gateway, "chat", body) + _assert_answered(first) + sends: Final[list[_Outcome]] = [] + + def resend() -> _Outcome: + served: Final = _send(rig.gateway, "chat", body) + sends.append(served) + return served + + hit: Final = eventually( + resend, lambda served: served.status == 200 and served.response_id == first.response_id, seconds=15 + ) + _assert_answered(hit) + received: Final = _received(rig.wire) + assert all(item.cache_points == 0 for item in received), received + assert len(received) == len(sends), (len(received), len(sends)) + _success_row(first.response_id) + + +@pytest.mark.timeout(600) +def test_x05_an_arn_whose_profile_name_starts_with_openai_dot_is_classified_by_the_pre_existing_family_rule( + gateway: Gateway, +) -> None: + arn: Final = f"{_PROFILE_ARN_PREFIX}openai.custom-{uuid.uuid4().hex[:8]}" + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment(gateway, scenario, wire, f"bedrock/{arn}") + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.model == arn, received.model + assert received.cache_points == 0, received.body + _success_row(outcome.response_id) + + +def _deployment( + gateway: Gateway, + scenario: Scenario, + wire: Wire, + model: str, + model_info: Mapping[str, JsonValue] | None = None, + *, + cleanup: bool = True, +) -> str: + name: Final = f"kimi-audit-{uuid.uuid4().hex}" + created: Final = gateway.post( + "/model/new", + { + "model_name": name, + "litellm_params": {"model": model, "api_base": wire.url, "num_retries": 0, **_AWS}, + "model_info": dict(model_info) if model_info is not None else {}, + }, + ) + identity: Final = str(_JSON.validate_python(created["model_info"])["id"]) + if cleanup: + scenario.cleanups.callback(gateway.post, "/model/delete", {"id": identity}) + _settled(gateway, name, wire) + return name + + +def _settled(gateway: Gateway, name: str, wire: Wire) -> None: + eventually( + lambda: tuple( + gateway.request("POST", "/v1/chat/completions", _chat_body(name, _prompt(), _NO_MARKERS)).status_code + for _ in range(12) + ), + lambda codes: all(code == 200 for code in codes), + seconds=60, + ) + wire.drain() + + +def _observe_at( + gateway: Gateway, wire: Wire, endpoint: Endpoint, body: Mapping[str, JsonValue] +) -> tuple[_Outcome, _Received]: + wire.drain() + outcome: Final = _send(gateway, endpoint, body) + return outcome, _only_received(wire) + + +def _arn(label: str) -> str: + return f"{_PROFILE_ARN_PREFIX}{label}-{uuid.uuid4().hex[:12]}" + + +_ROUTES: Final = ("", "converse/", "converse_like/") + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("route", _ROUTES, ids=("m05-plain-route", "m06-converse-route", "m07-converse-like-route")) +def test_m05_to_m07_a_deployment_flag_false_on_an_arn_strips_cache_points_on_every_route( + gateway: Gateway, route: str +) -> None: + arn: Final = _arn("flagged") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{route}{arn}", {"supports_prompt_cache_breakpoint": False} + ) + expected_model: Final = _CONVERSE_LIKE if route == "converse_like/" else arn + for endpoint in ("chat", "messages", "responses"): + outcome, received = _observe_at(gateway, wire, endpoint, _body(endpoint, name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.model == expected_model, (endpoint, received.model) + assert received.action == "converse", (endpoint, received.action) + assert received.cache_points == 0, (endpoint, received.body) + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_m08_an_arn_without_the_flag_still_gets_cache_points(gateway: Gateway) -> None: + arn: Final = _arn("plain") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment(gateway, scenario, wire, f"bedrock/{arn}") + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _ALL_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 3, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_m12_a_deployment_flag_true_on_kimi_overrides_the_cost_map_row(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{_KIMI_BASE}", {"supports_prompt_cache_breakpoint": True} + ) + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.model == _KIMI_BASE, received.model + assert received.cache_points == 2, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_m13_a_null_deployment_flag_on_kimi_falls_back_to_the_cost_map_row(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{_KIMI_US}", {"supports_prompt_cache_breakpoint": None} + ) + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.model == _KIMI_US, received.model + assert received.cache_points == 0, received.body + _success_row(outcome.response_id) + + +_ODD_FLAGS: Final[tuple[tuple[str, JsonValue], ...]] = ( + ("s01-string-true", "true"), + ("s02-int-one", 1), + ("s03-int-zero", 0), + ("s04-empty-list", []), + ("s05-empty-string", ""), + ("s06-5kb-string", "x" * 5120), +) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize(("label", "flag"), _ODD_FLAGS, ids=tuple(label for label, _ in _ODD_FLAGS)) +def test_s01_to_s06_an_odd_typed_flag_is_read_as_false_never_as_true( + gateway: Gateway, label: str, flag: JsonValue +) -> None: + arn: Final = _arn(label) + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment(gateway, scenario, wire, f"bedrock/{arn}", {"supports_prompt_cache_breakpoint": flag}) + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.cache_points == 0, (label, received.body) + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_s10_s11_two_deployments_of_one_arn_share_the_flag_and_a_delete_leaves_it_in_place(gateway: Gateway) -> None: + arn: Final = _arn("shared") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + flagged: Final = _deployment( + gateway, scenario, wire, f"bedrock/{arn}", {"supports_prompt_cache_breakpoint": False}, cleanup=False + ) + plain: Final = _deployment(gateway, scenario, wire, f"bedrock/{arn}") + for name in (flagged, plain): + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.cache_points == 0, (name, received.body) + _success_row(outcome.response_id) + flagged_identity: Final = _model_id(flagged) + gateway.post("/model/delete", {"id": flagged_identity}) + eventually( + lambda: tuple( + gateway.request("POST", "/v1/chat/completions", _chat_body(flagged, _prompt(), _NO_MARKERS)).status_code + for _ in range(12) + ), + lambda codes: all(code != 200 for code in codes), + seconds=60, + ) + wire.drain() + after, after_received = _observe_at(gateway, wire, "chat", _chat_body(plain, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(after) + assert after_received.cache_points == 0, after_received.body + _success_row(after.response_id) + + +def _model_id(name: str) -> str: + (row,) = read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_name=%s', (name,)) + return str(row["model_id"]) + + +def _stored_flag(identity: str) -> JsonValue: + (row,) = read_rows('SELECT model_info FROM "LiteLLM_ProxyModelTable" WHERE model_id=%s', (identity,)) + stored: Final = row["model_info"] + info: Final = _JSON.validate_python(stored if isinstance(stored, dict) else json.loads(str(stored))) + return info.get("supports_prompt_cache_breakpoint") + + +def _six_cache_point_counts(gateway: Gateway, wire: Wire, name: str) -> tuple[int, ...]: + wire.drain() + outcomes: Final = tuple(_send(gateway, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) for _ in range(6)) + assert all(item.status == 200 for item in outcomes), outcomes + return tuple(item.cache_points for item in _received(wire)) + + +@pytest.mark.timeout(600) +def test_x01_flipping_the_flag_to_true_through_a_patch_update_turns_cache_points_back_on(gateway: Gateway) -> None: + arn: Final = _arn("flip") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{arn}", {"supports_prompt_cache_breakpoint": False} + ) + before, before_received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(before) + assert before_received.cache_points == 0, before_received.body + identity: Final = _model_id(name) + patched: Final = gateway.request( + "PATCH", f"/model/{identity}/update", {"model_info": {"supports_prompt_cache_breakpoint": True}} + ) + assert patched.status_code == 200, patched.text + assert _stored_flag(identity) is True + points: Final = eventually( + lambda: _six_cache_point_counts(gateway, wire, name), + lambda counts: len(counts) == 6 and all(count == 2 for count in counts), + seconds=60, + ) + assert points == (2,) * 6, points + _success_row(before.response_id) + + +@pytest.mark.timeout(600) +def test_x06_the_legacy_post_update_answers_200_and_leaves_the_stored_flag_alone(gateway: Gateway) -> None: + arn: Final = _arn("legacy") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{arn}", {"supports_prompt_cache_breakpoint": False} + ) + before, before_received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(before) + identity: Final = _model_id(name) + updated: Final = gateway.request( + "POST", + "/model/update", + { + "model_name": name, + "litellm_params": {"model": f"bedrock/{arn}", "api_base": wire.url, "num_retries": 0, **_AWS}, + "model_info": {"id": identity, "supports_prompt_cache_breakpoint": True}, + }, + ) + assert updated.status_code == 200, updated.text + assert _stored_flag(identity) is False + _settled(gateway, name, wire) + after, after_received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(after) + assert after_received.cache_points == before_received.cache_points, (before_received.body, after_received.body) + _success_row(before.response_id) + _success_row(after.response_id) + + +@dataclass(frozen=True, slots=True) +class _Call: + endpoint: Endpoint + index: int + stream: bool + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + outcome: _Outcome + + +def _burst_calls(count: int) -> tuple[_Call, ...]: + endpoints: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses") + return tuple(_Call(endpoints[index % 3], index, index % 2 == 0) for index in range(count)) + + +async def _send_async(client: httpx.AsyncClient, key: str, model: str, call: _Call, prompt: str) -> _Served: + body: Final = _body(call.endpoint, model, f"{prompt} {call.index}", _SYSTEM_AND_USER, stream=call.stream) + async with client.stream( + "POST", _path(call.endpoint), json=body, headers={"Authorization": f"Bearer {key}"} + ) as response: + raw: Final = (await response.aread()).decode() + lines: Final = tuple(line for line in raw.splitlines() if line) + return _Served(call, _outcome_of(call.endpoint, call.stream, response, lines)) + + +async def _burst( + base_url: str, + key: str, + model: str, + calls: tuple[_Call, ...], + prompt: str, + *, + tolerate_transport_errors: bool = False, +) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_send_async(client, key, model, call, prompt) for call in calls), + return_exceptions=tolerate_transport_errors, + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Served)) + + +def _assert_rows_once(served: tuple[_Served, ...]) -> None: + ids: Final = frozenset(item.outcome.response_id for item in served) + assert len(ids) == len(served), ids + rows: Final = _spend_rows(ids, expected=len(ids)) + assert sorted(str(row["request_id"]) for row in rows) == sorted(ids), rows + assert all(row["status"] == "success" for row in rows), rows + + +def _reserved_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] + + +@pytest.mark.timeout(900) +async def test_c01_a_peer_outage_mid_traffic_fails_cleanly_and_recovery_lands_every_id_once( + gateway: Gateway, tmp_path: Path +) -> None: + port: Final = _reserved_port() + prompt: Final = _prompt() + with wire_server(_peer(), port=port) as first_peer: + config: Final = _owned_config(first_peer, tmp_path, names=("kimi-us", "claude-control")) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate = owned.gateway # rebind-ok: the same name covers the restarted proxy below + url = _proxy_url(candidate) # rebind-ok: the same name covers the restarted proxy below + served: Final = await _burst(url, candidate.key, "kimi-us", _burst_calls(30), prompt) + assert len(served) == 30 + for item in served: + _assert_answered(item.outcome) + received: Final = _received(first_peer) + assert len(received) == 30 and all(item.cache_points == 0 for item in received), received + _assert_rows_once(served) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate = owned.gateway + url = _proxy_url(candidate) + first_peer_down: Final = await _burst(url, candidate.key, "kimi-us", _burst_calls(30), prompt) + assert len(first_peer_down) == 30 + for item in first_peer_down: + assert item.outcome.status >= 500, (item.call, item.outcome.status, item.outcome.raw) + liveliness: Final = candidate.client.get("/health/liveliness", timeout=5) + assert liveliness.status_code == 200, liveliness.text + failed_ids: Final = frozenset( + item.outcome.call_id + for item in first_peer_down + if item.call.endpoint != "messages" and item.outcome.call_id + ) + assert len(failed_ids) == 20, failed_ids + failure_rows: Final = _spend_rows(failed_ids, expected=20) + assert all(row["status"] == "failure" for row in failure_rows), failure_rows + with wire_server(_peer(), port=port) as second_peer: + recovered: Final = await _burst(url, candidate.key, "kimi-us", _burst_calls(30), prompt) + assert len(recovered) == 30 + for item in recovered: + _assert_answered(item.outcome) + again: Final = _received(second_peer) + assert len(again) == 30 and all(item.cache_points == 0 for item in again), again + _assert_rows_once(recovered) + control: Final = await _burst(url, candidate.key, "claude-control", _burst_calls(6), prompt) + assert all(item.outcome.status == 200 for item in control), control + control_received: Final = _received(second_peer) + assert len(control_received) == 6, control_received + assert all(item.cache_points + item.cache_controls > 0 for item in control_received), control_received + + +def _open_peer_connections(pid: int, peer_url: str) -> int: + port: Final = urlsplit(peer_url).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +async def _held_burst(url: str, key: str, count: int, prompt: str) -> tuple[_Served, ...]: + calls: Final = tuple(_Call("chat", index, True) for index in range(count)) + return await _burst(url, key, "kimi-us", calls, f"{prompt} {_HOLD_MARKER}", tolerate_transport_errors=True) + + +@pytest.mark.timeout(900) +async def test_c02_a_worker_killed_mid_traffic_leaves_the_survivor_and_the_respawn_stripping( + gateway: Gateway, tmp_path: Path +) -> None: + hold: Final = threading.Event() + held: Final[SimpleQueue[str]] = SimpleQueue() + prompt: Final = _prompt() + with wire_server(_peer(hold, held)) as wire: + config: Final = _owned_config(wire, tmp_path, names=("kimi-us",)) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + url: Final = _proxy_url(candidate) + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=60, + ) + burst: Final = asyncio.create_task(_held_burst(url, candidate.key, 20, prompt)) + await asyncio.to_thread(eventually, held.qsize, lambda size: size == 20, 60) + held_by: Final = MappingProxyType({pid: _open_peer_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + psutil.Process(victim_pid).send_signal(signal.SIGKILL) + hold.set() + served: Final = await burst + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for item in served: + _assert_answered(item.outcome) + received: Final = _received(wire) + assert len(received) == 20 and all(item.cache_points == 0 for item in received), received + respawned: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 3, + seconds=60, + )[-1] + hold.clear() + for _ in range(20): + for _ in range(held.qsize()): + held.get_nowait() + follow_up_task: Final = asyncio.create_task(_held_burst(url, candidate.key, 6, prompt)) + await asyncio.to_thread(eventually, held.qsize, lambda size: size == 6, 60) + on_respawn: Final = _open_peer_connections(respawned, wire.url) + hold.set() + follow_up: Final = await follow_up_task + hold.clear() + assert len(follow_up) == 6, follow_up + for item in follow_up: + _assert_answered(item.outcome) + later: Final = _received(wire) + assert len(later) == 6 and all(item.cache_points == 0 for item in later), later + if on_respawn > 0: + break + else: + raise AssertionError("The respawned worker never took a held stream") + _assert_rows_once(served) + + +@pytest.mark.timeout(900) +def test_c03_a_restart_re_registers_a_stored_deployment_flag_from_the_database( + gateway: Gateway, tmp_path: Path +) -> None: + arn: Final = _arn("restart") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + config: Final = _owned_config(wire, tmp_path, names=("claude-control",)) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + name: Final = _deployment( + owned.gateway, + scenario, + wire, + f"bedrock/{arn}", + {"supports_prompt_cache_breakpoint": False}, + cleanup=False, + ) + before, before_received = _observe_at( + owned.gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER) + ) + _assert_answered(before) + assert before_received.cache_points == 0, before_received.body + _success_row(before.response_id) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as restarted: + _settled(restarted.gateway, name, wire) + after, after_received = _observe_at( + restarted.gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER) + ) + _assert_answered(after) + assert after_received.model == arn, after_received.model + assert after_received.cache_points == 0, after_received.body + _success_row(after.response_id) + control, control_received = _observe_at( + restarted.gateway, wire, "chat", _chat_body("claude-control", _prompt(), _SYSTEM_AND_USER) + ) + _assert_answered(control) + assert control_received.cache_points == 2, control_received.body + restarted.gateway.post("/model/delete", {"id": _model_id(name)}) + + +@pytest.mark.timeout(900) +async def test_c04_a_slow_peer_under_ten_concurrent_streams_completes_every_call_once(rig: _Rig) -> None: + rig.wire.drain() + calls: Final = tuple(_Call("chat", index, True) for index in range(10)) + started: Final = time.monotonic() + served: Final = await _burst( + _proxy_url(rig.gateway), rig.gateway.key, "kimi-us", calls, f"{_prompt()} {_SLOW_MARKER}" + ) + elapsed: Final = time.monotonic() - started + assert len(served) == 10 + for item in served: + _assert_answered(item.outcome) + received: Final = _received(rig.wire) + assert len(received) == 10 and all(item.cache_points == 0 for item in received), received + assert elapsed < 30, elapsed + _assert_rows_once(served) diff --git a/tests/unit/llms/bedrock/chat/test_converse_transformation.py b/tests/unit/llms/bedrock/chat/test_converse_transformation.py index 07c54eee395..e4a50317190 100644 --- a/tests/unit/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/unit/llms/bedrock/chat/test_converse_transformation.py @@ -5769,15 +5769,20 @@ def test_cache_control_injection_tool_config_drops_ttl_for_unsupported_model(): pytest.param("global.openai.gpt-6-astra", False, id="openai-family-implicit-caching-only"), pytest.param("openai.gpt-oss-120b-1:0", False, id="openai-gpt-oss"), pytest.param("us.openai.gpt-99-unmapped", False, id="unmapped-openai-family-still-suppressed"), + pytest.param("us.moonshotai.kimi-k3", False, id="kimi-k3-prices-cached-tokens-but-rejects-cachepoint"), + pytest.param("global.moonshotai.kimi-k3", False, id="kimi-k3-global-profile"), + pytest.param("us-east-1/us.moonshotai.kimi-k3", False, id="kimi-k3-regional-route-resolves-through-profile"), ], ) def test_cache_points_emitted_only_for_models_that_support_prompt_caching(model, expects_cache_points, monkeypatch): """Bedrock rejects cachePoint blocks for models without prompt caching support - ("You invoked an unsupported model or your request did not allow prompt caching"), - and clients like Claude Code attach cache_control to every request, so a map-known - model without the capability must not receive them. Unmapped ids (application - inference profile ARNs, models newer than the map) keep emitting so existing - caching setups never silently degrade.""" + ("You invoked an unsupported model or your request did not allow prompt caching") + and for models that price cached tokens yet take the marker only on their native + endpoints ("This model doesn't support the cachePoint field", Kimi K3), and clients + like Claude Code attach cache_control to every request, so a map-known model without + the capability must not receive them on system, message, or tool blocks. Unmapped ids + (application inference profile ARNs, models newer than the map) keep emitting so + existing caching setups never silently degrade.""" monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) @@ -5787,14 +5792,25 @@ def test_cache_points_emitted_only_for_models_that_support_prompt_caching(model, {"role": "system", "content": [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}]}, {"role": "user", "content": [{"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}}]}, ], - optional_params={}, + optional_params={ + "tools": [ + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object", "properties": {}}}, + "cache_control": {"type": "ephemeral"}, + } + ] + }, litellm_params={}, headers={}, ) - assert ("cachePoint" in json.dumps(body)) is expects_cache_points + assert ("cachePoint" in json.dumps(body["system"])) is expects_cache_points + assert ("cachePoint" in json.dumps(body["messages"])) is expects_cache_points + assert ("cachePoint" in json.dumps(body["toolConfig"])) is expects_cache_points assert body["system"][0]["text"] == "sys" assert body["messages"][0]["content"][0]["text"] == "hi" + assert body["toolConfig"]["tools"][0]["toolSpec"]["name"] == "get_weather" def test_tool_config_cachepoint_not_placed_or_credited_for_model_without_prompt_caching(monkeypatch): diff --git a/tests/unit/llms/bedrock/test_bedrock_common_utils.py b/tests/unit/llms/bedrock/test_bedrock_common_utils.py index 22e7d354be7..d4c182dc952 100644 --- a/tests/unit/llms/bedrock/test_bedrock_common_utils.py +++ b/tests/unit/llms/bedrock/test_bedrock_common_utils.py @@ -479,6 +479,75 @@ def test_capability_lookups_fall_back_to_base_model_when_regional_entry_lacks_fi assert bedrock_converse_supports_parallel_tool_use_config(regional) is True +@pytest.mark.parametrize( + ("entry", "expected"), + [ + pytest.param( + {"supports_prompt_caching": True, "supports_prompt_cache_breakpoint": False}, + False, + id="priced-cached-tokens-but-rejects-the-explicit-marker", + ), + pytest.param( + {"supports_prompt_caching": False, "supports_prompt_cache_breakpoint": True}, + True, + id="explicit-marker-flag-wins-over-the-caching-flag", + ), + pytest.param({"supports_prompt_caching": True}, True, id="caching-flag-alone-keeps-emitting"), + pytest.param({"supports_prompt_caching": False}, False, id="no-caching-and-no-marker-flag"), + ], +) +def test_bedrock_model_accepts_cache_points_prefers_the_explicit_breakpoint_flag(monkeypatch, entry, expected): + import litellm + from litellm.llms.bedrock.common_utils import bedrock_model_accepts_cache_points + + base = "vendor.breakpoint-flag-test" + monkeypatch.setitem(litellm.model_cost, f"us.{base}", {"input_cost_per_token": 1e-06}) + monkeypatch.setitem(litellm.model_cost, base, entry) + + assert bedrock_model_accepts_cache_points(f"us.{base}") is expected + + +@pytest.mark.parametrize("model", ["moonshotai.kimi-k3", "us.moonshotai.kimi-k3", "global.moonshotai.kimi-k3"]) +def test_kimi_k3_keeps_cached_token_pricing_while_refusing_converse_cache_points(model, local_model_cost_map): + import litellm + from litellm.llms.bedrock.common_utils import bedrock_model_accepts_cache_points + + assert bedrock_model_accepts_cache_points(model) is False + assert litellm.utils.supports_prompt_caching(model=model, custom_llm_provider="bedrock") is True + assert litellm.model_cost[model]["cache_read_input_token_cost"] > 0 + + +def test_deployment_model_info_breakpoint_flag_covers_an_unmapped_arn(local_model_cost_map): + from litellm import Router + from litellm.llms.bedrock.common_utils import bedrock_model_accepts_cache_points + + flagged_arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/flagged" + unflagged_arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/unflagged" + converse_arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/converse" + Router( + model_list=[ + { + "model_name": "kimi-k3-profile-converse", + "litellm_params": {"model": f"bedrock/converse/{converse_arn}", "aws_region_name": "us-east-1"}, + "model_info": {"supports_prompt_cache_breakpoint": False}, + }, + { + "model_name": "kimi-k3-profile", + "litellm_params": {"model": f"bedrock/{flagged_arn}", "aws_region_name": "us-east-1"}, + "model_info": {"supports_prompt_cache_breakpoint": False}, + }, + { + "model_name": "kimi-k3-profile-unflagged", + "litellm_params": {"model": f"bedrock/{unflagged_arn}", "aws_region_name": "us-east-1"}, + }, + ] + ) + + assert bedrock_model_accepts_cache_points(flagged_arn) is False + assert bedrock_model_accepts_cache_points(converse_arn) is False + assert bedrock_model_accepts_cache_points(unflagged_arn) is True + + def test_merge_bedrock_aws_request_params_strips_caller_identity_when_deployment_has_static_credentials(): from litellm.llms.bedrock.common_utils import merge_bedrock_aws_request_params