mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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>
This commit is contained in:
parent
7cdb3d3760
commit
bc2df7cdcd
4 changed files with 1392 additions and 23 deletions
|
|
@ -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:
|
||||
|
|
|
|||
344
tests/integration/providers/test_gemini_thinking_replay_chaos.py
Normal file
344
tests/integration/providers/test_gemini_thinking_replay_chaos.py
Normal file
|
|
@ -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)))
|
||||
990
tests/integration/providers/test_gemini_thinking_replay_wire.py
Normal file
990
tests/integration/providers/test_gemini_thinking_replay_wire.py
Normal file
|
|
@ -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"
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue