From bc2df7cdcd615be87de15d0cf14a2f8ea682f9b4 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 02:31:53 +0000 Subject: [PATCH] fix(gemini): stop replaying thinking block signatures to Gemini (#44661) * fix(gemini): stop replaying thinking block signatures to Gemini A thinking block's signature has no provenance, and LiteLLM never fills it from a Gemini response (Google signs text and functionCall parts, which ride provider_specific_fields and the tool call id), so a Claude signature replayed through a mixed model group reached Gemini as a thoughtSignature and Google answered 400 Invalid thought signature on every later Gemini-served turn. The same replay also sent the thinking text a second time as a plain text part. The thinking text now goes out once, as the thought part built from reasoning_content, and no part is built from thinking_blocks * test(gemini): type the parts helper and split its comprehension * test(gemini): cover thinking signature replay on the integration rig Two integration files from the audit of the foreign thought signature fix: 62 wire cells asserting the model turn Google receives on chat, messages and responses across gemini and vertex_ai, streaming and not, SDK and httpx clients, the sad shapes of thinking_blocks, context caching through cachedContents, and 3 chaos cells (a concurrent burst across endpoints, upstream stream drops, a worker SIGKILL mid burst) * test(gemini): read the integration salt from the environment The wire test decrypted Responses ids with a literal salt; tests/integration/_support/process.py boots the proxy with LITELLM_SALT_KEY when it is set, so the test now reads the same variable with the same default --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../llms/vertex_ai/gemini/transformation.py | 22 - .../test_gemini_thinking_replay_chaos.py | 344 ++++++ .../test_gemini_thinking_replay_wire.py | 990 ++++++++++++++++++ .../test_vertex_ai_gemini_transformation.py | 59 +- 4 files changed, 1392 insertions(+), 23 deletions(-) create mode 100644 tests/integration/providers/test_gemini_thinking_replay_chaos.py create mode 100644 tests/integration/providers/test_gemini_thinking_replay_wire.py diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index e3cc3bbb2dc..2521a462fc0 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -4,7 +4,6 @@ Transformation logic from OpenAI format to Gemini format. Why separate file? Make it easy to see how transformation works """ -import json import os import re from collections.abc import Mapping @@ -881,29 +880,8 @@ def _gemini_convert_messages_with_history( assistant_msg = ChatCompletionAssistantMessage(**msg_dict) _message_content = assistant_msg.get("content", None) reasoning_content = assistant_msg.get("reasoning_content", None) - thinking_blocks = assistant_msg.get("thinking_blocks") if reasoning_content is not None: assistant_content.append(PartType(thought=True, text=reasoning_content)) - if thinking_blocks is not None: - for block in thinking_blocks: - if block["type"] == "thinking": - block_thinking_str = block.get("thinking") - block_signature = block.get("signature") - if block_thinking_str is not None and block_signature is not None: - try: - assistant_content.append( - PartType( - thoughtSignature=block_signature, - **json.loads(block_thinking_str), - ) - ) - except Exception: - assistant_content.append( - PartType( - thoughtSignature=block_signature, - text=block_thinking_str, - ) - ) if _message_content is not None and isinstance(_message_content, list): _parts = [] for element in _message_content: diff --git a/tests/integration/providers/test_gemini_thinking_replay_chaos.py b/tests/integration/providers/test_gemini_thinking_replay_chaos.py new file mode 100644 index 00000000000..57058eb5fbe --- /dev/null +++ b/tests/integration/providers/test_gemini_thinking_replay_chaos.py @@ -0,0 +1,344 @@ +import asyncio +import json +import re +import signal +import threading +import uuid +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal, TypeAlias +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gemini-2.5-flash" +_API_KEY: Final = "synthetic-gemini-key" +_CONFIG_MODEL: Final = "gemini-thinking-replay-chaos" +_CLAUDE_SIGNATURE: Final = "CAQSyAsKEAgSGAI4AUIIdGhpbmtpbmcSDAlY-synthetic-claude-signature" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_CONTENTS: Final = TypeAdapter(list[dict[str, JsonValue]]) +_PARTS: Final = TypeAdapter(list[dict[str, JsonValue]]) +_MARKER: Final = re.compile(r"marker-([0-9a-f]{32})") +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_GENERATE: Final = f"/models/{_BACKEND}:generateContent" +_STREAM_GENERATE: Final = f"/models/{_BACKEND}:streamGenerateContent?alt=sse" +_USAGE: Final = {"promptTokenCount": 30, "candidatesTokenCount": 5, "totalTokenCount": 35} + +Endpoint: TypeAlias = Literal["chat", "messages", "responses"] + + +@dataclass(frozen=True, slots=True) +class _Call: + endpoint: Endpoint + stream: bool + marker: str + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + text: str + + +def _thought(marker: str) -> str: + return f"private thought for {marker}" + + +def _answer(marker: str) -> str: + return f"answer marker-{marker}" + + +def _response_id(marker: str) -> str: + return f"gemini-reply-{marker}" + + +def _path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "messages": + return "/v1/messages" + case "responses": + return "/v1/responses" + + +def _block(marker: str) -> Mapping[str, JsonValue]: + return {"type": "thinking", "thinking": _thought(marker), "signature": _CLAUDE_SIGNATURE} + + +def _body(model: str, call: _Call) -> Mapping[str, JsonValue]: + question: Final = f"Question marker-{call.marker}" + common: Final[Mapping[str, JsonValue]] = { + "model": model, + "stream": call.stream, + "num_retries": 0, + "cache": {"no-cache": True}, + } + match call.endpoint: + case "chat": + return { + **common, + "messages": [ + {"role": "user", "content": question}, + { + "role": "assistant", + "content": "Working on it.", + "reasoning_content": _thought(call.marker), + "thinking_blocks": [_block(call.marker)], + }, + {"role": "user", "content": "Go on."}, + ], + } + case "messages": + return { + **common, + "max_tokens": 64, + "messages": [ + {"role": "user", "content": question}, + {"role": "assistant", "content": [_block(call.marker), {"type": "text", "text": "Working on it."}]}, + {"role": "user", "content": "Go on."}, + ], + } + case "responses": + return { + **common, + "input": [ + {"role": "user", "content": question}, + { + "id": f"rs_{call.marker}", + "type": "reasoning", + "summary": [{"type": "summary_text", "text": _thought(call.marker)}], + "encrypted_content": json.dumps([_block(call.marker)]), + }, + { + "type": "message", + "role": "assistant", + "id": f"msg_{call.marker}", + "status": "completed", + "content": [{"type": "output_text", "text": "Working on it.", "annotations": []}], + }, + {"role": "user", "content": "Go on."}, + ], + } + + +def _frame(marker: str, text: str, finished: bool) -> Mapping[str, JsonValue]: + return { + "responseId": _response_id(marker), + "candidates": [ + { + "content": {"role": "model", "parts": [{"text": text}]}, + "index": 0, + **({"finishReason": "STOP"} if finished else {}), + } + ], + "usageMetadata": dict(_USAGE), + "modelVersion": _BACKEND, + } + + +def _gemini_reply(marker: str, stream: bool, abort_after: int | None = None, pause: float = 0) -> Reply: + if not stream: + return Reply(body=json.dumps(_frame(marker, _answer(marker), finished=True)).encode()) + frames: Final = (_frame(marker, "answer ", finished=False), _frame(marker, f"marker-{marker}", finished=True)) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"data: {json.dumps(frame)}\n\n".encode() for frame in frames), + abort_after=abort_after, + pause_between_chunks=pause, + ) + + +def _marker_of(request: Request) -> str: + found: Final = _MARKER.search(request.body.decode()) + assert found is not None, request.body + return found.group(1) + + +def _echo(request: Request) -> Reply: + assert request.target in (_GENERATE, _STREAM_GENERATE), request.target + return _gemini_reply(_marker_of(request), stream=request.target == _STREAM_GENERATE) + + +def _forwarded_thought(request: Request) -> tuple[str, JsonValue]: + body: Final = _JSON_OBJECT.validate_json(request.body) + assert "thoughtSignature" not in request.body.decode(), request.body + model_turn: Final = _CONTENTS.validate_python(body["contents"])[1] + assert model_turn["role"] == "model", model_turn + parts: Final = _PARTS.validate_python(model_turn["parts"]) + assert [part.get("thought") for part in parts] == [True, None], parts + return _marker_of(request), parts[0]["text"] + + +def _assert_no_bleed(received: tuple[Request, ...], markers: frozenset[str]) -> None: + forwarded: Final = [_forwarded_thought(request) for request in received] + assert sorted(marker for marker, _ in forwarded) == sorted(markers) + assert all(thought == _thought(marker) for marker, thought in forwarded), forwarded + + +def _spend_request_ids(model: str, expected: int) -> frozenset[str]: + rows: Final = eventually( + lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda found: len(found) >= expected, + seconds=60, + ) + assert [row["status"] for row in rows] == ["success"] * len(rows), rows + assert len(rows) == expected, rows + return frozenset(str(row["request_id"]) for row in rows) + + +async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> _Served: + async with client.stream( + "POST", + _path(call.endpoint), + json=_body(model, call), + headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"}, + ) as response: + raw: Final = await response.aread() + return _Served(call=call, status=response.status_code, text=raw.decode()) + + +async def _burst( + base_url: str, key: str, model: str, calls: tuple[_Call, ...], *, 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(client, key, model, call) 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 _calls(count: int, endpoints: tuple[Endpoint, ...], stream: Callable[[int], bool]) -> tuple[_Call, ...]: + return tuple( + _Call(endpoint=endpoints[index % len(endpoints)], stream=stream(index), marker=uuid.uuid4().hex) + for index in range(count) + ) + + +def _assert_answered_with_its_own_marker(served: _Served) -> None: + assert served.status == 200, served.text + assert set(_MARKER.findall(served.text)) == {served.call.marker}, served.text + + +async def test_concurrent_thinking_replays_across_endpoints_reach_gemini_signature_free(gateway: Gateway) -> None: + calls: Final = _calls(30, ("chat", "messages", "responses"), lambda index: index % 2 == 0) + with wire_server(_echo) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls) + assert len(served) == 30 + for item in served: + _assert_answered_with_its_own_marker(item) + _assert_no_bleed(wire.drain(), frozenset(call.marker for call in calls)) + landed: Final = _spend_request_ids(model, 30) + assert landed >= { + _response_id(call.marker) for call in calls if not (call.endpoint == "messages" and call.stream) + } + + +async def test_upstream_stream_drops_reach_callers_and_later_replays_still_reach_gemini(gateway: Gateway) -> None: + calls: Final = _calls(12, ("chat", "messages", "responses"), lambda _: True) + dropped: Final = frozenset(call.marker for index, call in enumerate(calls) if index % 3 == 0) + + def respond(request: Request) -> Reply: + marker: Final = _marker_of(request) + return _gemini_reply(marker, stream=True, abort_after=0 if marker in dropped else None) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls) + assert len(served) == 12 + for item in served: + if item.call.marker in dropped: + assert item.status == 500, item.text + assert "marker-" not in item.text, item.text + else: + _assert_answered_with_its_own_marker(item) + recovery: Final = _Call(endpoint="chat", stream=True, marker=uuid.uuid4().hex) + (recovered,) = await _burst(str(gateway.client.base_url), gateway.key, model, (recovery,)) + _assert_answered_with_its_own_marker(recovered) + assert recovered.text.rstrip().endswith("data: [DONE]"), recovered.text + _assert_no_bleed(wire.drain(), frozenset(call.marker for call in (*calls, recovery))) + + +def _chaos_config(wire: Wire, tmp_path: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["model_list"] = [ + { + "model_name": _CONFIG_MODEL, + "litellm_params": {"model": f"gemini/{_BACKEND}", "api_base": wire.url, "api_key": _API_KEY}, + } + ] + path: Final = tmp_path / "gemini-thinking-replay-chaos.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).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 + ) + + +@pytest.mark.timeout(180) +async def test_worker_sigkill_mid_burst_leaves_the_sibling_replaying_thinking_signature_free( + gateway: Gateway, tmp_path: Path +) -> None: + calls: Final = _calls(20, ("chat", "messages", "responses"), lambda _: False) + release: Final = threading.Event() + held_markers: Final[SimpleQueue[str]] = SimpleQueue() + + def held(request: Request) -> Reply: + held_markers.put(_marker_of(request)) + assert release.wait(timeout=60), "The burst was never released" + return _echo(request) + + with wire_server(held) as wire: + path: Final = _chaos_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned: + candidate: Final = owned.gateway + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + burst: Final = asyncio.create_task( + _burst( + str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, calls, tolerate_transport_errors=True + ) + ) + await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60) + held_by: Final = MappingProxyType({pid: _open_upstream_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__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + served: Final = await burst + assert held_by[survivor_pid] >= 10, held_by + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for item in served: + _assert_answered_with_its_own_marker(item) + follow_up: Final = _Call(endpoint="messages", stream=False, marker=uuid.uuid4().hex) + (answered,) = await _burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, (follow_up,)) + _assert_answered_with_its_own_marker(answered) + received: Final = wire.drain() + assert {(request.method, request.target) for request in received} == {("POST", _GENERATE)}, received + _assert_no_bleed(received, frozenset(call.marker for call in (*calls, follow_up))) diff --git a/tests/integration/providers/test_gemini_thinking_replay_wire.py b/tests/integration/providers/test_gemini_thinking_replay_wire.py new file mode 100644 index 00000000000..571274d20e2 --- /dev/null +++ b/tests/integration/providers/test_gemini_thinking_replay_wire.py @@ -0,0 +1,990 @@ +import base64 +import itertools +import json +import os +import uuid +from collections.abc import Iterator, Mapping, Sequence +from collections.abc import Set as AbstractSet +from dataclasses import dataclass +from typing import Final, Literal, TypeAlias + +import anthropic +import httpx +import openai +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, Scenario, eventually +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, Wire, wire_server +from openai.types.responses import ResponseCompletedEvent +from pydantic import JsonValue, TypeAdapter + +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with + +_BACKEND: Final = "gemini-2.5-flash" +_API_KEY: Final = "synthetic-gemini-key" +_PROJECT: Final = "scripted-project" +_LOCATION: Final = "us-central1" +_MODEL_PATH: Final = f"/v1/projects/{_PROJECT}/locations/{_LOCATION}/publishers/google/models/{_BACKEND}" +_SALT: Final = os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt") +_CLAUDE_SIGNATURE: Final = "CAQSyAsKEAgSGAI4AUIIdGhpbmtpbmcSDAlY-synthetic-claude-signature" +_GEMINI_SIGNATURE: Final = "synthetic-gemini-thought-signature" +_REASONING: Final = "private thought about fruit" +_QUESTION: Final = "How many apples are left?" +_PRIOR_ANSWER: Final = "Two apples." +_FOLLOW_UP: Final = "And after eating one?" +_ANSWER: Final = "One left." +_CACHE_BUST: Final[Mapping[str, JsonValue]] = {"cache": {"no-cache": True}} +_USAGE: Final = {"promptTokenCount": 20, "candidatesTokenCount": 5, "totalTokenCount": 25} +_FUNCTION: Final = { + "name": "count_fruit", + "description": "Count fruit", + "parameters": {"type": "object", "properties": {"kind": {"type": "string"}}, "required": ["kind"]}, +} +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_JSON_LIST: Final = TypeAdapter(list[JsonValue]) + +Provider: TypeAlias = Literal["gemini", "vertex_ai"] +Endpoint: TypeAlias = Literal["chat", "messages", "responses"] +Client: TypeAlias = Literal["sdk_sync", "sdk_async", "httpx"] +_PROVIDERS: Final[tuple[Provider, ...]] = ("gemini", "vertex_ai") +_ENDPOINTS: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses") +_CLIENTS: Final[tuple[Client, ...]] = ("sdk_sync", "sdk_async", "httpx") + + +@dataclass(frozen=True, slots=True) +class _Observed: + response_id: str + answer: str + + +def _thinking_block(signature: JsonValue, thinking: str = _REASONING) -> Mapping[str, JsonValue]: + return {"type": "thinking", "thinking": thinking, "signature": signature} + + +def _user(text: str) -> Mapping[str, JsonValue]: + return {"role": "user", "parts": [{"text": text}]} + + +def _model_turn(*parts: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]: + return {"role": "model", "parts": list(parts)} + + +def _thought(text: str = _REASONING) -> Mapping[str, JsonValue]: + return {"thought": True, "text": text} + + +_REPLAYED_TURN: Final = _model_turn(_thought(), {"text": _PRIOR_ANSWER}) + + +def _chat_history(assistant_fields: Mapping[str, JsonValue]) -> Sequence[Mapping[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + {"role": "assistant", "content": _PRIOR_ANSWER, **assistant_fields}, + {"role": "user", "content": _FOLLOW_UP}, + ] + + +_CHAT_HISTORY: Final = _chat_history( + {"reasoning_content": _REASONING, "thinking_blocks": [_thinking_block(_CLAUDE_SIGNATURE)]} +) +_CONTROL_HISTORY: Final = _chat_history({}) + + +def _messages_history(*assistant_blocks: Mapping[str, JsonValue]) -> Sequence[Mapping[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + {"role": "assistant", "content": [*assistant_blocks, {"type": "text", "text": _PRIOR_ANSWER}]}, + {"role": "user", "content": _FOLLOW_UP}, + ] + + +_MESSAGES_HISTORY: Final = _messages_history(_thinking_block(_CLAUDE_SIGNATURE)) + + +def _responses_input(*blocks: Mapping[str, JsonValue]) -> Sequence[Mapping[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + { + "type": "reasoning", + "id": "rs_prior", + "summary": [{"type": "summary_text", "text": _REASONING}], + "encrypted_content": json.dumps(list(blocks)), + }, + { + "type": "message", + "role": "assistant", + "id": "msg_prior", + "status": "completed", + "content": [{"type": "output_text", "text": _PRIOR_ANSWER, "annotations": []}], + }, + {"role": "user", "content": _FOLLOW_UP}, + ] + + +_RESPONSES_INPUT: Final = _responses_input(_thinking_block(_CLAUDE_SIGNATURE)) + + +def _service_account_json(token_url: str) -> str: + private_key: Final = ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() + ) + return json.dumps( + { + "type": "service_account", + "project_id": _PROJECT, + "private_key_id": "scripted", + "private_key": private_key, + "client_email": f"scripted@{_PROJECT}.iam.gserviceaccount.com", + "client_id": "0", + "auth_uri": f"{token_url}/_oauth/authorize", + "token_uri": f"{token_url}/_oauth/token", + } + ) + + +def _register(gateway: Gateway, scenario: Scenario, provider: Provider, wire: Wire) -> str: + if provider == "gemini": + return scenario.model(model=f"gemini/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + return scenario.model( + model=f"vertex_ai/{_BACKEND}", + api_base=f"{wire.url}{_MODEL_PATH}", + api_key=None, + vertex_project=_PROJECT, + vertex_location=_LOCATION, + vertex_credentials=_service_account_json(gateway.upstream_url.rstrip("/")), + ) + + +def _target(provider: Provider, stream: bool) -> str: + prefix: Final = f"/models/{_BACKEND}" if provider == "gemini" else _MODEL_PATH + return f"{prefix}:streamGenerateContent?alt=sse" if stream else f"{prefix}:generateContent" + + +def _expected_auth(provider: Provider) -> tuple[str, str]: + return ("x-goog-api-key", _API_KEY) if provider == "gemini" else ("authorization", "Bearer scripted-token") + + +def _frame(response_id: str, parts: Sequence[Mapping[str, JsonValue]], finished: bool) -> Mapping[str, JsonValue]: + return { + "responseId": response_id, + "candidates": [ + { + "content": {"role": "model", "parts": list(parts)}, + "index": 0, + **({"finishReason": "STOP"} if finished else {}), + } + ], + "usageMetadata": dict(_USAGE), + "modelVersion": _BACKEND, + } + + +def _reply(response_id: str, stream: bool, parts: Sequence[Mapping[str, JsonValue]] = ({"text": _ANSWER},)) -> Reply: + if not stream: + return Reply(body=json.dumps(_frame(response_id, parts, finished=True)).encode()) + frames: Final = (_frame(response_id, parts[:-1], finished=False), _frame(response_id, parts[-1:], finished=True)) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"data: {json.dumps(frame)}\n\n".encode() for frame in frames if frame["candidates"]), + ) + + +def _single(values: AbstractSet[str]) -> str: + assert len(values) == 1, values + return next(iter(values)) + + +def _openai(gateway: Gateway) -> openai.OpenAI: + return openai.OpenAI( + base_url=f"{gateway.client.base_url}/v1", + api_key=gateway.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=60), + ) + + +def _async_openai(gateway: Gateway) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI( + base_url=f"{gateway.client.base_url}/v1", + api_key=gateway.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=60), + ) + + +def _anthropic(gateway: Gateway) -> anthropic.Anthropic: + return anthropic.Anthropic( + base_url=str(gateway.client.base_url), + api_key=gateway.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=60), + ) + + +def _async_anthropic(gateway: Gateway) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic( + base_url=str(gateway.client.base_url), + api_key=gateway.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=60), + ) + + +def _chat_sync(gateway: Gateway, model: str, stream: bool) -> _Observed: + client: Final = _openai(gateway) + if not stream: + reply: Final = client.chat.completions.create(model=model, messages=_CHAT_HISTORY, extra_body=dict(_CACHE_BUST)) + return _Observed(reply.id, reply.choices[0].message.content or "") + chunks: Final = list( + client.chat.completions.create(model=model, messages=_CHAT_HISTORY, stream=True, extra_body=dict(_CACHE_BUST)) + ) + return _Observed( + _single({chunk.id for chunk in chunks}), + "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices), + ) + + +async def _chat_async(gateway: Gateway, model: str, stream: bool) -> _Observed: + client: Final = _async_openai(gateway) + if not stream: + reply: Final = await client.chat.completions.create( + model=model, messages=_CHAT_HISTORY, extra_body=dict(_CACHE_BUST) + ) + return _Observed(reply.id, reply.choices[0].message.content or "") + chunks: Final = [ + chunk + async for chunk in await client.chat.completions.create( + model=model, messages=_CHAT_HISTORY, stream=True, extra_body=dict(_CACHE_BUST) + ) + ] + return _Observed( + _single({chunk.id for chunk in chunks}), + "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices), + ) + + +def _messages_sync(gateway: Gateway, model: str, stream: bool) -> _Observed: + client: Final = _anthropic(gateway) + if not stream: + reply: Final = client.messages.create( + model=model, max_tokens=64, messages=_MESSAGES_HISTORY, extra_body=dict(_CACHE_BUST) + ) + return _Observed(reply.id, "".join(block.text for block in reply.content if block.type == "text")) + with client.messages.stream( + model=model, max_tokens=64, messages=_MESSAGES_HISTORY, extra_body=dict(_CACHE_BUST) + ) as streamed: + final: Final = streamed.get_final_message() + return _Observed(final.id, "".join(block.text for block in final.content if block.type == "text")) + + +async def _messages_async(gateway: Gateway, model: str, stream: bool) -> _Observed: + client: Final = _async_anthropic(gateway) + if not stream: + reply: Final = await client.messages.create( + model=model, max_tokens=64, messages=_MESSAGES_HISTORY, extra_body=dict(_CACHE_BUST) + ) + return _Observed(reply.id, "".join(block.text for block in reply.content if block.type == "text")) + async with client.messages.stream( + model=model, max_tokens=64, messages=_MESSAGES_HISTORY, extra_body=dict(_CACHE_BUST) + ) as streamed: + final: Final = await streamed.get_final_message() + return _Observed(final.id, "".join(block.text for block in final.content if block.type == "text")) + + +def _responses_sync(gateway: Gateway, model: str, stream: bool) -> _Observed: + client: Final = _openai(gateway) + if not stream: + reply: Final = client.responses.create(model=model, input=_RESPONSES_INPUT, extra_body=dict(_CACHE_BUST)) + return _Observed(reply.id, reply.output_text) + events: Final = list( + client.responses.create(model=model, input=_RESPONSES_INPUT, stream=True, extra_body=dict(_CACHE_BUST)) + ) + completed: Final = events[-1] + assert isinstance(completed, ResponseCompletedEvent), events + return _Observed(completed.response.id, completed.response.output_text) + + +async def _responses_async(gateway: Gateway, model: str, stream: bool) -> _Observed: + client: Final = _async_openai(gateway) + if not stream: + reply: Final = await client.responses.create(model=model, input=_RESPONSES_INPUT, extra_body=dict(_CACHE_BUST)) + return _Observed(reply.id, reply.output_text) + events: Final = [ + event + async for event in await client.responses.create( + model=model, input=_RESPONSES_INPUT, stream=True, extra_body=dict(_CACHE_BUST) + ) + ] + completed: Final = events[-1] + assert isinstance(completed, ResponseCompletedEvent), events + return _Observed(completed.response.id, completed.response.output_text) + + +def _sse_payloads(text: str) -> tuple[Mapping[str, JsonValue], ...]: + lines: Final = tuple(line for line in text.splitlines() if line.startswith("data: ") and line != "data: [DONE]") + return tuple(_JSON_OBJECT.validate_json(line.removeprefix("data: ").encode()) for line in lines) + + +def _httpx_chat(gateway: Gateway, model: str, stream: bool) -> _Observed: + response: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _CHAT_HISTORY, "stream": stream, **_CACHE_BUST} + ) + assert response.status_code == 200, response.text + if not stream: + payload: Final = _JSON_OBJECT.validate_json(response.content) + choice: Final = _JSON_OBJECT.validate_python(_JSON_LIST.validate_python(payload["choices"])[0]) + return _Observed(str(payload["id"]), str(_JSON_OBJECT.validate_python(choice["message"])["content"])) + assert response.text.rstrip().endswith("data: [DONE]"), response.text + chunks: Final = _sse_payloads(response.text) + deltas: Final = tuple( + _JSON_OBJECT.validate_python( + _JSON_OBJECT.validate_python(_JSON_LIST.validate_python(chunk["choices"])[0])["delta"] + ) + for chunk in chunks + ) + return _Observed( + _single({str(chunk["id"]) for chunk in chunks}), + "".join(str(delta["content"]) for delta in deltas if delta.get("content")), + ) + + +def _httpx_messages(gateway: Gateway, model: str, stream: bool) -> _Observed: + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 64, "messages": _MESSAGES_HISTORY, "stream": stream, **_CACHE_BUST}, + headers={"anthropic-version": "2023-06-01"}, + ) + assert response.status_code == 200, response.text + if not stream: + payload: Final = _JSON_OBJECT.validate_json(response.content) + blocks: Final = tuple( + _JSON_OBJECT.validate_python(block) for block in _JSON_LIST.validate_python(payload["content"]) + ) + return _Observed(str(payload["id"]), "".join(str(block["text"]) for block in blocks if block["type"] == "text")) + events: Final = _sse_payloads(response.text) + starts: Final = tuple(event for event in events if event["type"] == "message_start") + assert events[-1]["type"] == "message_stop", events + deltas: Final = tuple( + _JSON_OBJECT.validate_python(event["delta"]) for event in events if event["type"] == "content_block_delta" + ) + return _Observed( + str(_JSON_OBJECT.validate_python(starts[0]["message"])["id"]), + "".join(str(delta["text"]) for delta in deltas if delta["type"] == "text_delta"), + ) + + +def _httpx_responses(gateway: Gateway, model: str, stream: bool) -> _Observed: + response: Final = gateway.request( + "POST", "/v1/responses", {"model": model, "input": _RESPONSES_INPUT, "stream": stream, **_CACHE_BUST} + ) + assert response.status_code == 200, response.text + if not stream: + payload: Final = _JSON_OBJECT.validate_json(response.content) + return _Observed(str(payload["id"]), _output_text(payload)) + events: Final = _sse_payloads(response.text) + assert events[-1]["type"] == "response.completed", events + completed: Final = _JSON_OBJECT.validate_python(events[-1]["response"]) + return _Observed(str(completed["id"]), _output_text(completed)) + + +def _content_parts(item: Mapping[str, JsonValue]) -> Iterator[Mapping[str, JsonValue]]: + return (_JSON_OBJECT.validate_python(part) for part in _JSON_LIST.validate_python(item["content"])) + + +def _output_text(payload: Mapping[str, JsonValue]) -> str: + items: Final = tuple(_JSON_OBJECT.validate_python(item) for item in _JSON_LIST.validate_python(payload["output"])) + messages: Final = tuple(item for item in items if item["type"] == "message") + parts: Final = tuple(itertools.chain.from_iterable(_content_parts(item) for item in messages)) + return "".join(str(part["text"]) for part in parts if part["type"] == "output_text") + + +async def _observe(gateway: Gateway, endpoint: Endpoint, client: Client, model: str, stream: bool) -> _Observed: + match (endpoint, client): + case ("chat", "sdk_sync"): + return _chat_sync(gateway, model, stream) + case ("chat", "sdk_async"): + return await _chat_async(gateway, model, stream) + case ("chat", "httpx"): + return _httpx_chat(gateway, model, stream) + case ("messages", "sdk_sync"): + return _messages_sync(gateway, model, stream) + case ("messages", "sdk_async"): + return await _messages_async(gateway, model, stream) + case ("messages", "httpx"): + return _httpx_messages(gateway, model, stream) + case ("responses", "sdk_sync"): + return _responses_sync(gateway, model, stream) + case ("responses", "sdk_async"): + return await _responses_async(gateway, model, stream) + case ("responses", "httpx"): + return _httpx_responses(gateway, model, stream) + raise AssertionError((endpoint, client)) + + +def _upstream_response_id(endpoint: Endpoint, caller_id: str) -> str: + if endpoint != "responses": + return caller_id + decrypted: Final = decrypt_if_encrypted_with(caller_id.removeprefix("resp_"), _SALT) + assert decrypted is not None, caller_id + issued: Final = decrypted.split(";")[0].split("response_id:")[-1] + return base64.b64decode(issued.removeprefix("resp_")).decode().split(";")[-1].removeprefix("response_id:") + + +def _call_type(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "acompletion" + case "messages": + return "anthropic_messages" + case "responses": + return "aresponses" + + +def _spend_row(*request_ids: str) -> Mapping[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT model_group, call_type, status, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" ' + "WHERE request_id = ANY(%s)", + (list(request_ids),), # pyright: ignore[reportArgumentType] # psycopg adapts the list to a text array + ), + lambda found: len(found) == 1, + seconds=70, + ) + return rows[0] + + +def _assert_caller_id(endpoint: Endpoint, stream: bool, caller_id: str, response_id: str) -> None: + if endpoint == "messages" and stream: + assert caller_id.startswith("msg_"), caller_id + return + assert _upstream_response_id(endpoint, caller_id) == response_id, caller_id + + +def _expected_body(endpoint: Endpoint, *contents: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]: + return { + "contents": list(contents), + **({"generationConfig": {"max_output_tokens": 64}} if endpoint == "messages" else {}), + } + + +def _only_request(wire: Wire, provider: Provider, stream: bool) -> Mapping[str, JsonValue]: + received: Final = wire.drain() + assert [(request.method, request.target) for request in received] == [("POST", _target(provider, stream))], received + header, value = _expected_auth(provider) + assert received[0].headers[header] == value, dict(received[0].headers) + assert "thoughtSignature" not in received[0].body.decode(), received[0].body + return _JSON_OBJECT.validate_json(received[0].body) + + +def _happy_cells() -> tuple[pytest.ParameterSet, ...]: + return tuple( + pytest.param( + provider, endpoint, stream, client, id=f"{provider}-{endpoint}-{'stream' if stream else 'sync'}-{client}" + ) + for provider, endpoint, stream, client in itertools.product(_PROVIDERS, _ENDPOINTS, (False, True), _CLIENTS) + ) + + +@pytest.mark.parametrize(("provider", "endpoint", "stream", "client"), _happy_cells()) +async def test_claude_thinking_replay_reaches_gemini_as_a_thought_part_without_its_signature( + gateway: Gateway, provider: Provider, endpoint: Endpoint, stream: bool, client: Client +) -> None: + response_id: Final = f"gemini-reply-{uuid.uuid4().hex}" + + def respond(request: Request) -> Reply: + return _reply(response_id, stream) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _register(gateway, scenario, provider, wire) + observed: Final = await _observe(gateway, endpoint, client, model, stream) + assert observed.answer == _ANSWER, observed + _assert_caller_id(endpoint, stream, observed.response_id, response_id) + assert _only_request(wire, provider, stream) == _expected_body( + endpoint, _user(_QUESTION), _REPLAYED_TURN, _user(_FOLLOW_UP) + ) + assert _spend_row(observed.response_id, response_id) == { + "model_group": model, + "call_type": _call_type(endpoint), + "status": "success", + "prompt_tokens": 20, + "completion_tokens": 5, + } + + +_GEMINI_TOOL_REPLY_PARTS: Final = ( + _thought("I should count."), + {"functionCall": {"name": "count_fruit", "args": {"kind": "apple"}}, "thoughtSignature": _GEMINI_SIGNATURE}, +) +_GEMINI_TOOL_REPLAYED_TURN: Final = _model_turn( + _thought("I should count."), + {"function_call": {"name": "count_fruit", "args": {"kind": "apple"}}, "thoughtSignature": _GEMINI_SIGNATURE}, +) +_FUNCTION_RESPONSE_TURN: Final = { + "role": "user", + "parts": [{"function_response": {"name": "count_fruit", "response": {"content": "12"}}}], +} +_DECLARATIONS: Final = {"tools": [{"function_declarations": [_FUNCTION]}]} + + +def _tool_round_one(endpoint: Endpoint, model: str) -> Mapping[str, JsonValue]: + match endpoint: + case "chat": + return { + "model": model, + "tools": [{"type": "function", "function": _FUNCTION}], + "messages": [{"role": "user", "content": "Count the apples."}], + **_CACHE_BUST, + } + case "messages": + return { + "model": model, + "max_tokens": 64, + "tools": [ + {"name": "count_fruit", "description": "Count fruit", "input_schema": _FUNCTION["parameters"]} + ], + "messages": [{"role": "user", "content": "Count the apples."}], + **_CACHE_BUST, + } + case "responses": + return { + "model": model, + "tools": [{"type": "function", **_FUNCTION}], + "input": [{"role": "user", "content": "Count the apples."}], + **_CACHE_BUST, + } + + +def _tool_round_two( + endpoint: Endpoint, first: Mapping[str, JsonValue], payload: Mapping[str, JsonValue] +) -> Mapping[str, JsonValue]: + match endpoint: + case "chat": + choice: Final = _JSON_OBJECT.validate_python(_JSON_LIST.validate_python(payload["choices"])[0]) + message: Final = _JSON_OBJECT.validate_python(choice["message"]) + call: Final = _JSON_OBJECT.validate_python(_JSON_LIST.validate_python(message["tool_calls"])[0]) + return { + **first, + "messages": [ + *_JSON_LIST.validate_python(first["messages"]), + message, + {"role": "tool", "tool_call_id": call["id"], "content": "12"}, + ], + } + case "messages": + blocks: Final = tuple( + _JSON_OBJECT.validate_python(block) for block in _JSON_LIST.validate_python(payload["content"]) + ) + tool_use: Final = next(block for block in blocks if block["type"] == "tool_use") + return { + **first, + "messages": [ + *_JSON_LIST.validate_python(first["messages"]), + {"role": "assistant", "content": list(blocks)}, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": tool_use["id"], "content": "12"}], + }, + ], + } + case "responses": + items: Final = tuple( + _JSON_OBJECT.validate_python(item) for item in _JSON_LIST.validate_python(payload["output"]) + ) + function_call: Final = next(item for item in items if item["type"] == "function_call") + return { + **first, + "input": [ + *_JSON_LIST.validate_python(first["input"]), + *items, + {"type": "function_call_output", "call_id": function_call["call_id"], "output": "12"}, + ], + } + + +def _path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "messages": + return "/v1/messages" + case "responses": + return "/v1/responses" + + +@pytest.mark.parametrize("endpoint", _ENDPOINTS) +def test_gemini_own_tool_call_signature_is_still_replayed_on_the_function_call_part( + gateway: Gateway, endpoint: Endpoint +) -> None: + rounds: Final = (f"gemini-reply-{uuid.uuid4().hex}", f"gemini-reply-{uuid.uuid4().hex}") + calls: Final = itertools.count() + + def respond(request: Request) -> Reply: + if next(calls) == 0: + return _reply(rounds[0], stream=False, parts=_GEMINI_TOOL_REPLY_PARTS) + return _reply(rounds[1], stream=False, parts=({"text": "Twelve apples."},)) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _register(gateway, scenario, "gemini", wire) + first: Final = _tool_round_one(endpoint, model) + headers: Final = {"anthropic-version": "2023-06-01"} + opened: Final = gateway.request("POST", _path(endpoint), first, headers=headers) + assert opened.status_code == 200, opened.text + assert _GEMINI_SIGNATURE in opened.text, opened.text + second: Final = _tool_round_two(endpoint, first, _JSON_OBJECT.validate_json(opened.content)) + closed: Final = gateway.request("POST", _path(endpoint), second, headers=headers) + assert closed.status_code == 200, closed.text + assert "Twelve apples." in closed.text, closed.text + received: Final = wire.drain() + assert [(request.method, request.target) for request in received] == [("POST", _target("gemini", False))] * 2 + replay: Final = _JSON_OBJECT.validate_json(received[1].body) + assert replay == { + **_expected_body(endpoint, _user("Count the apples."), _GEMINI_TOOL_REPLAYED_TURN, _FUNCTION_RESPONSE_TURN), + **_DECLARATIONS, + }, replay + assert _spend_row(rounds[1])["status"] == "success" + + +_FIVE_KB: Final = "s" * 5000 +_ASSISTANT_SHAPES: Final = ( + pytest.param({"reasoning_content": _REASONING, "thinking_blocks": 5}, _REPLAYED_TURN, id="blocks-int"), + pytest.param( + {"reasoning_content": _REASONING, "thinking_blocks": _FIVE_KB}, _REPLAYED_TURN, id="blocks-5kb-string" + ), + pytest.param({"reasoning_content": _REASONING, "thinking_blocks": ""}, _REPLAYED_TURN, id="blocks-empty-string"), + pytest.param( + {"reasoning_content": _REASONING, "thinking_blocks": [{"type": "thinking"}]}, + _REPLAYED_TURN, + id="block-without-thinking", + ), + pytest.param( + {"reasoning_content": _REASONING, "thinking_blocks": [{"signature": "s"}]}, + _REPLAYED_TURN, + id="block-without-type", + ), + pytest.param( + {"reasoning_content": _REASONING, "thinking_blocks": [_thinking_block(5)]}, _REPLAYED_TURN, id="signature-int" + ), + pytest.param({"reasoning_content": _REASONING, "thinking_blocks": ["str"]}, _REPLAYED_TURN, id="block-string"), + pytest.param( + {"reasoning_content": _REASONING, "thinking_blocks": [_thinking_block(_FIVE_KB)]}, + _REPLAYED_TURN, + id="signature-5kb", + ), + pytest.param( + {"reasoning_content": "", "thinking_blocks": [_thinking_block(_CLAUDE_SIGNATURE)]}, + _model_turn(_thought(""), {"text": _PRIOR_ANSWER}), + id="reasoning-empty-with-blocks", + ), + pytest.param( + {"reasoning_content": None, "thinking_blocks": None}, _model_turn({"text": _PRIOR_ANSWER}), id="both-null" + ), + pytest.param({"reasoning_content": _REASONING, "thinking_blocks": []}, _REPLAYED_TURN, id="blocks-empty-list"), + pytest.param({}, _model_turn({"text": _PRIOR_ANSWER}), id="both-missing"), + pytest.param( + {"thinking_blocks": [_thinking_block(_CLAUDE_SIGNATURE)]}, + _model_turn({"text": _PRIOR_ANSWER}), + id="blocks-without-reasoning-content", + ), + pytest.param( + { + "reasoning_content": _REASONING, + "thinking_blocks": [_thinking_block(_CLAUDE_SIGNATURE, json.dumps(_thought()))], + }, + _REPLAYED_TURN, + id="json-era-block", + ), +) + + +@pytest.mark.parametrize(("assistant_fields", "expected_turn"), _ASSISTANT_SHAPES) +def test_chat_assistant_thinking_shapes_never_crash_or_leak_a_signature( + gateway: Gateway, assistant_fields: Mapping[str, JsonValue], expected_turn: Mapping[str, JsonValue] +) -> None: + response_ids: Final = (f"gemini-reply-{uuid.uuid4().hex}", f"gemini-reply-{uuid.uuid4().hex}") + calls: Final = itertools.count() + + def respond(request: Request) -> Reply: + return _reply(response_ids[next(calls)], stream=False) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _register(gateway, scenario, "gemini", wire) + shaped: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": _chat_history(assistant_fields), **_CACHE_BUST}, + ) + assert shaped.status_code == 200, shaped.text + assert _JSON_OBJECT.validate_json(shaped.content)["id"] == response_ids[0], shaped.text + control: Final = gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": _CONTROL_HISTORY, **_CACHE_BUST} + ) + assert control.status_code == 200, control.text + assert _JSON_OBJECT.validate_json(control.content)["id"] == response_ids[1], control.text + received: Final = wire.drain() + assert [(request.method, request.target) for request in received] == [("POST", _target("gemini", False))] * 2 + assert "thoughtSignature" not in received[0].body.decode(), received[0].body + assert _JSON_OBJECT.validate_json(received[0].body) == _expected_body( + "chat", _user(_QUESTION), expected_turn, _user(_FOLLOW_UP) + ) + assert _JSON_OBJECT.validate_json(received[1].body) == _expected_body( + "chat", _user(_QUESTION), _model_turn({"text": _PRIOR_ANSWER}), _user(_FOLLOW_UP) + ) + assert [_spend_row(response_id)["status"] for response_id in response_ids] == ["success", "success"] + + +def test_chat_same_thinking_replay_sent_twice_lands_one_spend_row_per_request(gateway: Gateway) -> None: + response_ids: Final = (f"gemini-reply-{uuid.uuid4().hex}", f"gemini-reply-{uuid.uuid4().hex}") + calls: Final = itertools.count() + + def respond(request: Request) -> Reply: + return _reply(response_ids[next(calls)], stream=False) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _register(gateway, scenario, "gemini", wire) + body: Final = {"model": model, "messages": _CHAT_HISTORY, **_CACHE_BUST} + answers: Final = tuple(gateway.request("POST", "/v1/chat/completions", body) for _ in response_ids) + assert [answer.status_code for answer in answers] == [200, 200], [answer.text for answer in answers] + assert [_JSON_OBJECT.validate_json(answer.content)["id"] for answer in answers] == list(response_ids) + received: Final = wire.drain() + assert [_JSON_OBJECT.validate_json(request.body) for request in received] == [ + _expected_body("chat", _user(_QUESTION), _REPLAYED_TURN, _user(_FOLLOW_UP)) + ] * 2 + assert [_spend_row(response_id)["status"] for response_id in response_ids] == ["success", "success"] + + +def test_chat_two_consecutive_thinking_turns_merge_into_one_signature_free_model_turn(gateway: Gateway) -> None: + response_id: Final = f"gemini-reply-{uuid.uuid4().hex}" + with wire_server(lambda request: _reply(response_id, stream=False)) as wire, gateway.scenario() as scenario: + model: Final = _register(gateway, scenario, "gemini", wire) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [ + {"role": "user", "content": _QUESTION}, + { + "role": "assistant", + "content": "Two.", + "reasoning_content": "first thought", + "thinking_blocks": [_thinking_block(_CLAUDE_SIGNATURE, "first thought")], + }, + { + "role": "assistant", + "content": "Wait, three.", + "reasoning_content": "second thought", + "thinking_blocks": [_thinking_block(_CLAUDE_SIGNATURE, "second thought")], + }, + {"role": "user", "content": _FOLLOW_UP}, + ], + **_CACHE_BUST, + }, + ) + assert response.status_code == 200, response.text + assert _JSON_OBJECT.validate_json(response.content)["id"] == response_id, response.text + assert _only_request(wire, "gemini", False) == _expected_body( + "chat", + _user(_QUESTION), + _model_turn( + _thought("first thought"), {"text": "Two."}, _thought("second thought"), {"text": "Wait, three."} + ), + _user(_FOLLOW_UP), + ) + assert _spend_row(response_id)["status"] == "success" + + +def test_chat_claude_thinking_beside_a_tool_call_replays_the_thought_and_the_call_only(gateway: Gateway) -> None: + response_id: Final = f"gemini-reply-{uuid.uuid4().hex}" + with wire_server(lambda request: _reply(response_id, stream=False)) as wire, gateway.scenario() as scenario: + model: Final = _register(gateway, scenario, "gemini", wire) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "tools": [{"type": "function", "function": _FUNCTION}], + "messages": [ + {"role": "user", "content": "Count the apples."}, + { + "role": "assistant", + "content": None, + "reasoning_content": _REASONING, + "thinking_blocks": [_thinking_block(_CLAUDE_SIGNATURE)], + "tool_calls": [ + { + "id": "call_prior", + "type": "function", + "function": {"name": "count_fruit", "arguments": json.dumps({"kind": "apple"})}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_prior", "content": "12"}, + ], + **_CACHE_BUST, + }, + ) + assert response.status_code == 200, response.text + assert _JSON_OBJECT.validate_json(response.content)["id"] == response_id, response.text + assert _only_request(wire, "gemini", False) == { + **_expected_body( + "chat", + _user("Count the apples."), + _model_turn(_thought(), {"function_call": {"name": "count_fruit", "args": {"kind": "apple"}}}), + _FUNCTION_RESPONSE_TURN, + ), + **_DECLARATIONS, + } + assert _spend_row(response_id)["status"] == "success" + + +_MESSAGES_SHAPES: Final = ( + pytest.param(_thinking_block(5), _REPLAYED_TURN, id="signature-int"), + pytest.param(_thinking_block(""), _REPLAYED_TURN, id="signature-empty"), + pytest.param( + {"type": "redacted_thinking", "data": "synthetic-redacted-data"}, + _model_turn({"text": _PRIOR_ANSWER}), + id="redacted-thinking", + ), +) + + +@pytest.mark.parametrize(("assistant_block", "expected_turn"), _MESSAGES_SHAPES) +def test_messages_assistant_thinking_shapes_never_leak_a_signature( + gateway: Gateway, assistant_block: Mapping[str, JsonValue], expected_turn: Mapping[str, JsonValue] +) -> None: + response_id: Final = f"gemini-reply-{uuid.uuid4().hex}" + with wire_server(lambda request: _reply(response_id, stream=False)) as wire, gateway.scenario() as scenario: + model: Final = _register(gateway, scenario, "gemini", wire) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 64, "messages": _messages_history(assistant_block), **_CACHE_BUST}, + headers={"anthropic-version": "2023-06-01"}, + ) + assert response.status_code == 200, response.text + assert _JSON_OBJECT.validate_json(response.content)["id"] == response_id, response.text + assert _only_request(wire, "gemini", False) == _expected_body( + "messages", _user(_QUESTION), expected_turn, _user(_FOLLOW_UP) + ) + assert _spend_row(response_id)["status"] == "success" + + +def test_responses_encrypted_block_with_an_int_signature_still_replays_the_summary_only(gateway: Gateway) -> None: + response_id: Final = f"gemini-reply-{uuid.uuid4().hex}" + with wire_server(lambda request: _reply(response_id, stream=False)) as wire, gateway.scenario() as scenario: + model: Final = _register(gateway, scenario, "gemini", wire) + response: Final = gateway.request( + "POST", "/v1/responses", {"model": model, "input": _responses_input(_thinking_block(5)), **_CACHE_BUST} + ) + assert response.status_code == 200, response.text + caller_id: Final = str(_JSON_OBJECT.validate_json(response.content)["id"]) + assert _upstream_response_id("responses", caller_id) == response_id, caller_id + assert _only_request(wire, "gemini", False) == _expected_body( + "responses", _user(_QUESTION), _REPLAYED_TURN, _user(_FOLLOW_UP) + ) + assert _spend_row(response_id)["status"] == "success" + + +_CACHE_NAME: Final = "cachedContents/synthetic-cache" +_CACHED_POLICY: Final = " ".join(f"policy clause {index} applies" for index in range(600)) +_EPHEMERAL: Final = {"type": "ephemeral"} + + +def _cached_reply(response_id: str) -> Reply: + return Reply( + body=json.dumps( + { + **_frame(response_id, ({"text": _ANSWER},), finished=True), + "usageMetadata": {**_USAGE, "promptTokenCount": 1300, "cachedContentTokenCount": 1290}, + } + ).encode() + ) + + +def _cached_chat_history() -> Sequence[Mapping[str, JsonValue]]: + return [ + {"role": "user", "content": [{"type": "text", "text": _CACHED_POLICY, "cache_control": _EPHEMERAL}]}, + { + "role": "assistant", + "content": _PRIOR_ANSWER, + "reasoning_content": _REASONING, + "thinking_blocks": [_thinking_block(_CLAUDE_SIGNATURE)], + "cache_control": _EPHEMERAL, + }, + {"role": "user", "content": [{"type": "text", "text": "Acknowledged.", "cache_control": _EPHEMERAL}]}, + {"role": "user", "content": _FOLLOW_UP}, + ] + + +def _cached_messages_history() -> Sequence[Mapping[str, JsonValue]]: + return [ + {"role": "user", "content": [{"type": "text", "text": _CACHED_POLICY, "cache_control": _EPHEMERAL}]}, + { + "role": "assistant", + "content": [ + _thinking_block(_CLAUDE_SIGNATURE), + {"type": "text", "text": _PRIOR_ANSWER, "cache_control": _EPHEMERAL}, + ], + }, + {"role": "user", "content": [{"type": "text", "text": "Acknowledged.", "cache_control": _EPHEMERAL}]}, + {"role": "user", "content": _FOLLOW_UP}, + ] + + +@pytest.mark.parametrize("endpoint", ("chat", "messages")) +def test_context_cached_thinking_turn_is_stored_without_its_signature(gateway: Gateway, endpoint: Endpoint) -> None: + response_id: Final = f"gemini-reply-{uuid.uuid4().hex}" + + def respond(request: Request) -> Reply: + if request.method == "GET": + return Reply(body=b"{}") + if request.target.endswith("cachedContents"): + return Reply(body=json.dumps({"name": _CACHE_NAME, "model": f"models/{_BACKEND}"}).encode()) + return _cached_reply(response_id) + + body: Final = ( + {"messages": _cached_chat_history()} + if endpoint == "chat" + else {"max_tokens": 64, "messages": _cached_messages_history()} + ) + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _register(gateway, scenario, "gemini", wire) + response: Final = gateway.request( + "POST", + _path(endpoint), + {"model": model, **body, **_CACHE_BUST}, + headers={"anthropic-version": "2023-06-01"}, + ) + assert response.status_code == 200, response.text + assert _JSON_OBJECT.validate_json(response.content)["id"] == response_id, response.text + assert _ANSWER in response.text, response.text + received: Final = wire.drain() + assert [(request.method, request.target.endswith("cachedContents")) for request in received] == [ + ("GET", True), + ("POST", True), + ("POST", False), + ], received + assert received[2].target == f"/models/{_BACKEND}:generateContent", received + assert all("thoughtSignature" not in request.body.decode() for request in received[1:]), received + stored: Final = _JSON_OBJECT.validate_json(received[1].body) + assert isinstance(stored["displayName"], str) and stored["displayName"], stored + assert stored == { + "contents": [_user(_CACHED_POLICY), _REPLAYED_TURN, _user("Acknowledged.")], + "model": f"models/{_BACKEND}", + "displayName": stored["displayName"], + "tools": None, + }, stored + assert _JSON_OBJECT.validate_json(received[2].body) == { + **_expected_body(endpoint, _user(_FOLLOW_UP)), + "cachedContent": _CACHE_NAME, + } + assert _spend_row(response_id)["status"] == "success" diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 0b37e033023..940bd71fdb4 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -1,4 +1,6 @@ import base64 +from collections.abc import Sequence +from itertools import chain from pathlib import Path from typing import Final @@ -17,7 +19,7 @@ from litellm.llms.vertex_ai.gemini.transformation import ( _get_highest_media_resolution, _extract_max_media_resolution_from_messages, ) -from litellm.types.llms.vertex_ai import BlobType +from litellm.types.llms.vertex_ai import BlobType, ContentType, PartType from litellm.types.utils import Message @@ -2769,3 +2771,58 @@ def test_convert_tool_response_with_url_image(monkeypatch: pytest.MonkeyPatch) - assert len(function_response["parts"]) == 1 inline_data: Final[BlobType] = function_response["parts"][0]["inline_data"] assert inline_data == {"data": base64.b64encode(WHITE_PNG).decode(), "mime_type": "image/png"} + + +CLAUDE_THINKING_SIGNATURE: Final = "EqQBCkYIBxgCKkCfQ2x0b3VkZS1zaWduYXR1cmUtbm90LW1pbnRlZC1ieS1nZW1pbmkSDJ3lXf5sD+QVqpFQmRoM" + + +def _parts_of(contents: Sequence[ContentType]) -> list[PartType]: + return list(chain.from_iterable(content["parts"] for content in contents)) + + +def test_thinking_block_signature_is_not_forwarded_to_gemini() -> None: + thinking: Final = "The user wants the capital of France." + messages: Final = [ + {"role": "user", "content": "Capital of France?"}, + { + "role": "assistant", + "content": "Paris.", + "reasoning_content": thinking, + "thinking_blocks": [{"type": "thinking", "thinking": thinking, "signature": CLAUDE_THINKING_SIGNATURE}], + }, + {"role": "user", "content": "And of Spain?"}, + ] + + parts: Final = _parts_of(_gemini_convert_messages_with_history(messages=messages, model="gemini-3.8-flash")) + + assert all("thoughtSignature" not in part for part in parts) + assert [part for part in parts if part.get("text") == thinking] == [{"thought": True, "text": thinking}] + + +def test_anthropic_messages_history_replays_to_gemini_without_claude_signature() -> None: + from litellm.litellm_core_utils.prompt_templates.factory import _get_dummy_thought_signature + from litellm.llms.anthropic.pass_through.adapters.transformation import LiteLLMAnthropicMessagesAdapter + + thinking: Final = "I should look the weather up before answering." + anthropic_messages: Final = [ + {"role": "user", "content": "Weather in Paris?"}, + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": thinking, "signature": CLAUDE_THINKING_SIGNATURE}, + {"type": "text", "text": "Let me check."}, + {"type": "tool_use", "id": "toolu_01A", "name": "get_weather", "input": {"city": "Paris"}}, + ], + }, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01A", "content": "Sunny, 22C"}]}, + ] + + chat_messages: Final = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai( + messages=anthropic_messages + ) + parts: Final = _parts_of(_gemini_convert_messages_with_history(messages=chat_messages, model="gemini-3.8-flash")) + + signatures: Final = [part["thoughtSignature"] for part in parts if "thoughtSignature" in part] + assert signatures == [_get_dummy_thought_signature()] + assert [part for part in parts if part.get("text") == thinking] == [{"thought": True, "text": thinking}] + assert next(part for part in parts if "function_call" in part)["thoughtSignature"] == _get_dummy_thought_signature()