mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(anthropic): count a leading system run through count_tokens' system parameter (#45463)
* fix(anthropic): count a leading system run through count_tokens' system parameter Anthropic's count_tokens rejects role "system" at the head of messages, so /v1/responses/input_tokens with instructions, and /v1/messages/count_tokens with a system-role message, fell back to the local tokenizer. The shared Anthropic count_tokens transformation now lifts the leading run of system messages into the top-level system parameter, the way the chat path sends it, after any system the caller set. Anthropic direct, Azure AI Anthropic, and Bedrock Mantle share that transformation. * test(integration): cover the count_tokens leading-system lift across Anthropic, Azure AI and Bedrock Mantle --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
72ab736863
commit
d34ad35281
7 changed files with 1254 additions and 34 deletions
|
|
@ -31,7 +31,7 @@ Anthropic wire shape is built later by ``anthropic_messages_pt``.
|
|||
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from itertools import chain, groupby
|
||||
from typing import Final, Literal, TypeAlias
|
||||
from typing import Final, Literal, TypeAlias, TypeVar
|
||||
|
||||
from litellm.types.llms.anthropic import AnthropicMessagesSystemMessageParam, AnthropicSystemMessageContent
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -55,6 +55,7 @@ _RENDERED_ASSISTANT_PART_TYPES: Final = frozenset({"text", "server_tool_use"})
|
|||
_THINKING_BLOCK_TYPES: Final = frozenset({"thinking", "redacted_thinking"})
|
||||
|
||||
_MessageKind: TypeAlias = Literal["system", "tool", "user", "other"]
|
||||
_Message: Final = TypeVar("_Message")
|
||||
_TextPart: TypeAlias = tuple[str, ChatCompletionCachedContent | None]
|
||||
|
||||
|
||||
|
|
@ -96,8 +97,8 @@ def _kind(message: object) -> _MessageKind:
|
|||
|
||||
|
||||
def split_leading_system_run(
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> tuple[tuple[AllMessageValues, ...], tuple[AllMessageValues, ...]]:
|
||||
messages: Sequence[_Message],
|
||||
) -> tuple[tuple[_Message, ...], tuple[_Message, ...]]:
|
||||
"""Split ``messages`` into the leading run of system messages and everything after it."""
|
||||
leading_count: Final = next(
|
||||
(index for index, message in enumerate(messages) if not is_system_message(message)),
|
||||
|
|
@ -160,9 +161,16 @@ def _anthropic_text_block(part: _TextPart) -> AnthropicSystemMessageContent:
|
|||
return cached
|
||||
|
||||
|
||||
def anthropic_system_blocks(run: Sequence[object]) -> tuple[AnthropicSystemMessageContent, ...]:
|
||||
"""The top-level ``system`` blocks for a run of system messages: every non-empty text part in order,
|
||||
each keeping its ``cache_control``, which is the shape the chat path sends for the leading run."""
|
||||
parts: Final = chain.from_iterable(_text_parts(message) for message in run)
|
||||
return tuple(_anthropic_text_block(part) for part in parts)
|
||||
|
||||
|
||||
def anthropic_system_messages(message: object) -> tuple[AnthropicMessagesSystemMessageParam, ...]:
|
||||
"""The Anthropic wire message for a system message, or nothing when it carries no text."""
|
||||
blocks: Final = tuple(_anthropic_text_block(part) for part in _text_parts(message))
|
||||
blocks: Final = anthropic_system_blocks((message,))
|
||||
if not blocks:
|
||||
return ()
|
||||
wire: Final[AnthropicMessagesSystemMessageParam] = {
|
||||
|
|
|
|||
|
|
@ -11,11 +11,16 @@ from typing import Final
|
|||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
|
||||
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import (
|
||||
anthropic_system_blocks,
|
||||
split_leading_system_run,
|
||||
)
|
||||
from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers
|
||||
from litellm.llms.anthropic.wif import resolve_anthropic_base
|
||||
from litellm.types.llms.openai import ChatCompletionImageObject
|
||||
|
||||
_COUNT_REQUEST: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_SYSTEM_BLOCKS: Final = TypeAdapter(list[JsonValue])
|
||||
_IMAGE_BLOCK: Final = TypeAdapter(ChatCompletionImageObject)
|
||||
COUNT_TOKEN_OPTION_NAMES: Final = ("thinking", "tool_choice", "output_config")
|
||||
|
||||
|
|
@ -48,6 +53,27 @@ def _count_content(content: JsonValue) -> JsonValue:
|
|||
return [_count_block(block) for block in content] if isinstance(content, list) else content
|
||||
|
||||
|
||||
def _lift_leading_system(
|
||||
messages: Sequence[Mapping[str, JsonValue]], system: JsonValue
|
||||
) -> tuple[tuple[Mapping[str, JsonValue], ...], JsonValue]:
|
||||
"""Move the leading run of system-role messages into the top-level ``system`` parameter.
|
||||
|
||||
count_tokens only takes the initial system prompt there and answers 400 on ``role: "system"``
|
||||
at the head of ``messages``; the chat path sends the same run as ``system``. A caller's own
|
||||
``system`` keeps its place ahead of the lifted blocks, and a ``system`` that is neither text
|
||||
nor a block list is left as sent, messages included, for the provider to judge.
|
||||
"""
|
||||
leading, conversation = split_leading_system_run(messages)
|
||||
if not leading or not (system is None or isinstance(system, (str, list))):
|
||||
return tuple(messages), system
|
||||
lifted: Final = _SYSTEM_BLOCKS.validate_python(list(anthropic_system_blocks(leading)))
|
||||
if isinstance(system, list):
|
||||
return conversation, [*system, *lifted]
|
||||
if isinstance(system, str) and system:
|
||||
return conversation, [{"type": "text", "text": system}, *lifted]
|
||||
return conversation, lifted or system
|
||||
|
||||
|
||||
class AnthropicCountTokensConfig:
|
||||
"""
|
||||
Configuration and transformation logic for Anthropic CountTokens API.
|
||||
|
|
@ -85,16 +111,24 @@ class AnthropicCountTokensConfig:
|
|||
"""
|
||||
Transform request to Anthropic CountTokens format.
|
||||
|
||||
Includes optional system and tools fields for accurate token counting.
|
||||
Includes optional system and tools fields for accurate token counting; a leading run of
|
||||
system-role messages is counted through ``system``, the only place count_tokens accepts it.
|
||||
"""
|
||||
options: Final[Mapping[str, JsonValue]] = optional_params or MappingProxyType({})
|
||||
counted_messages, counted_system = _lift_leading_system(messages, system)
|
||||
return _COUNT_REQUEST.validate_python(
|
||||
MappingProxyType(
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{**message, "content": _count_content(message["content"])} for message in messages],
|
||||
"messages": [
|
||||
{**message, "content": _count_content(message["content"])} for message in counted_messages
|
||||
],
|
||||
**MappingProxyType(
|
||||
{key: value for key, value in (("system", system), ("tools", tools)) if value is not None}
|
||||
{
|
||||
key: value
|
||||
for key, value in (("system", counted_system), ("tools", tools))
|
||||
if value is not None
|
||||
}
|
||||
),
|
||||
**MappingProxyType(
|
||||
{key: value for key, value in options.items() if key in COUNT_TOKEN_OPTION_NAMES}
|
||||
|
|
|
|||
159
tests/integration/providers/_count_tokens_system_lift.py
Normal file
159
tests/integration/providers/_count_tokens_system_lift.py
Normal file
|
|
@ -0,0 +1,159 @@
|
|||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
import openai
|
||||
from integration._support.client import Gateway
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
TOKEN_COUNTING_BETA: Final = "token-counting-2024-11-01"
|
||||
ANTHROPIC_VERSION: Final = "2023-06-01"
|
||||
REJECTION: Final = 'messages.0: Unexpected role "system". The Messages API accepts a top-level `system` parameter'
|
||||
|
||||
USER_TEXT: Final = "Count this message"
|
||||
USER: Final[dict[str, JsonValue]] = {"role": "user", "content": USER_TEXT}
|
||||
ASSISTANT: Final[dict[str, JsonValue]] = {"role": "assistant", "content": "One."}
|
||||
FOLLOW_UP: Final[dict[str, JsonValue]] = {"role": "user", "content": "Again"}
|
||||
INSTRUCTION: Final = "You are a terse assistant"
|
||||
REMINDER: Final = "Answer in one sentence"
|
||||
LEADING: Final[dict[str, JsonValue]] = {"role": "system", "content": INSTRUCTION}
|
||||
LIFTED: Final[dict[str, JsonValue]] = {"type": "text", "text": INSTRUCTION}
|
||||
MID_SYSTEM: Final[dict[str, JsonValue]] = {"role": "system", "content": REMINDER}
|
||||
EPHEMERAL: Final[dict[str, JsonValue]] = {"type": "ephemeral"}
|
||||
ONE_HOUR: Final[dict[str, JsonValue]] = {"type": "ephemeral", "ttl": "1h"}
|
||||
IMAGE_PART: Final[dict[str, JsonValue]] = {
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
|
||||
},
|
||||
}
|
||||
FIVE_KB: Final = "Answer in one sentence. " * 214
|
||||
CALLER_SYSTEM: Final = "Prefer metric units"
|
||||
CALLER_BLOCKS: Final[list[JsonValue]] = [
|
||||
{"type": "text", "text": CALLER_SYSTEM},
|
||||
{"type": "text", "text": "Never guess", "cache_control": EPHEMERAL},
|
||||
]
|
||||
TOOLS: Final[list[JsonValue]] = [
|
||||
{
|
||||
"name": "get_weather",
|
||||
"description": "Look up the current weather for a city",
|
||||
"input_schema": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string", "description": "City to look up"}},
|
||||
"required": ["city"],
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LiftCase:
|
||||
messages: tuple[dict[str, JsonValue], ...]
|
||||
system: JsonValue | None
|
||||
lifted: list[JsonValue] | None
|
||||
|
||||
|
||||
LIFT_CASES: Final[Mapping[str, LiftCase]] = MappingProxyType(
|
||||
{
|
||||
"string": LiftCase((LEADING, USER), None, [LIFTED]),
|
||||
"run_with_cache_control": LiftCase(
|
||||
(
|
||||
{"role": "system", "content": INSTRUCTION, "cache_control": ONE_HOUR},
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL},
|
||||
{"type": "text", "text": ""},
|
||||
IMAGE_PART,
|
||||
],
|
||||
},
|
||||
USER,
|
||||
),
|
||||
None,
|
||||
[
|
||||
{"type": "text", "text": INSTRUCTION, "cache_control": ONE_HOUR},
|
||||
{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL},
|
||||
],
|
||||
),
|
||||
"caller_system_string_first": LiftCase(
|
||||
(LEADING, USER), CALLER_SYSTEM, [{"type": "text", "text": CALLER_SYSTEM}, LIFTED]
|
||||
),
|
||||
"caller_system_blocks_first": LiftCase((LEADING, USER), CALLER_BLOCKS, [*CALLER_BLOCKS, LIFTED]),
|
||||
"caller_empty_system_dropped": LiftCase((LEADING, USER), "", [LIFTED]),
|
||||
"empty_content_dropped": LiftCase(({"role": "system", "content": ""}, USER), None, None),
|
||||
"image_only_content_dropped": LiftCase(({"role": "system", "content": [IMAGE_PART]}, USER), None, None),
|
||||
"integer_content_dropped": LiftCase(({"role": "system", "content": 5}, USER), None, None),
|
||||
"integer_text_part_dropped": LiftCase(
|
||||
({"role": "system", "content": [{"type": "text", "text": 7}]}, USER), None, None
|
||||
),
|
||||
"five_kb": LiftCase(({"role": "system", "content": FIVE_KB}, USER), None, [{"type": "text", "text": FIVE_KB}]),
|
||||
}
|
||||
)
|
||||
STRING_CASE: Final = LIFT_CASES["string"]
|
||||
|
||||
|
||||
def count_request(model: str, case: LiftCase) -> dict[str, JsonValue]:
|
||||
return {"model": model, "messages": list(case.messages), **({} if case.system is None else {"system": case.system})}
|
||||
|
||||
|
||||
def expected_count_body(model: str, case: LiftCase) -> dict[str, JsonValue]:
|
||||
conversation: Final[list[JsonValue]] = [message for message in case.messages if message["role"] != "system"]
|
||||
return {"model": model, "messages": conversation, **({} if case.lifted is None else {"system": case.lifted})}
|
||||
|
||||
|
||||
def _opening_message(message: JsonValue) -> bool:
|
||||
return isinstance(message, dict) and message.get("role") in ("user", "assistant")
|
||||
|
||||
|
||||
def _conversation_message(message: JsonValue) -> bool:
|
||||
return isinstance(message, dict) and message.get("role") in ("user", "assistant", "system")
|
||||
|
||||
|
||||
def _anthropic_tool(tool: JsonValue) -> bool:
|
||||
return isinstance(tool, dict) and isinstance(tool.get("name"), str) and isinstance(tool.get("input_schema"), dict)
|
||||
|
||||
|
||||
def accepts_count_body(body: Mapping[str, JsonValue]) -> bool:
|
||||
# Anthropic, Azure AI Foundry and Bedrock Mantle count_tokens verdicts observed live on 2026-10-09: an empty
|
||||
# messages list and a system role at messages[0] answer 400 invalid_request_error, a later system role counts
|
||||
messages: Final = body.get("messages")
|
||||
tools: Final = body.get("tools", [])
|
||||
return (
|
||||
isinstance(messages, list)
|
||||
and len(messages) > 0
|
||||
and _opening_message(messages[0])
|
||||
and all(map(_conversation_message, messages))
|
||||
and isinstance(body.get("system", ""), (str, list))
|
||||
and isinstance(tools, list)
|
||||
and all(map(_anthropic_tool, tools))
|
||||
)
|
||||
|
||||
|
||||
def anthropic_client(gateway: Gateway) -> anthropic.Anthropic:
|
||||
return anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic:
|
||||
return anthropic.AsyncAnthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def openai_client(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),
|
||||
)
|
||||
|
||||
|
||||
def async_openai_client(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),
|
||||
)
|
||||
|
|
@ -0,0 +1,572 @@
|
|||
import asyncio
|
||||
import json
|
||||
import socket
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import AbstractContextManager, ExitStack
|
||||
from dataclasses import dataclass
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final, cast
|
||||
|
||||
import httpcore
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
from anthropic.types import MessageParam
|
||||
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from integration.providers._count_tokens_system_lift import (
|
||||
ANTHROPIC_VERSION,
|
||||
ASSISTANT,
|
||||
EPHEMERAL,
|
||||
FOLLOW_UP,
|
||||
IMAGE_PART,
|
||||
INSTRUCTION,
|
||||
JSON_OBJECT,
|
||||
LEADING,
|
||||
LIFT_CASES,
|
||||
LIFTED,
|
||||
MID_SYSTEM,
|
||||
REJECTION,
|
||||
REMINDER,
|
||||
STRING_CASE,
|
||||
TOKEN_COUNTING_BETA,
|
||||
TOOLS,
|
||||
USER,
|
||||
USER_TEXT,
|
||||
LiftCase,
|
||||
accepts_count_body,
|
||||
anthropic_client,
|
||||
async_anthropic_client,
|
||||
async_openai_client,
|
||||
count_request,
|
||||
expected_count_body,
|
||||
openai_client,
|
||||
)
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_MODEL: Final = "claude-opus-5-5"
|
||||
_API_KEY: Final = "synthetic-count-tokens-key"
|
||||
_COUNT: Final = 3131
|
||||
_PROVIDER_TOKENIZERS: Final = frozenset({"anthropic_api", "azure_ai_anthropic_api"})
|
||||
_SDK_MESSAGES: Final = cast(
|
||||
list[MessageParam], [LEADING, USER]
|
||||
) # cast-ok: the SDK types reject the role the proxy lifts
|
||||
_STRING_BODY: Final = expected_count_body(_MODEL, STRING_CASE)
|
||||
_PROXY_MODULE: Final = "integration._support.proxy"
|
||||
_PROBES_PER_ROUND: Final = 8
|
||||
_CLIENT_ADDRESS: Final = TypeAdapter(tuple[str, int])
|
||||
_ARGUMENTS: Final = TypeAdapter(tuple[str, ...])
|
||||
_NAME: Final = TypeAdapter(str)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Provider:
|
||||
prefix: str
|
||||
target: str
|
||||
tokenizer: str
|
||||
azure: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Deployment:
|
||||
provider: _Provider
|
||||
port: int
|
||||
model: str
|
||||
|
||||
|
||||
_ANTHROPIC: Final = _Provider("anthropic", "/v1/messages/count_tokens", "anthropic_api", False)
|
||||
_AZURE: Final = _Provider("azure_ai", "/anthropic/v1/messages/count_tokens", "azure_ai_anthropic_api", True)
|
||||
_PROVIDERS: Final = MappingProxyType({"anthropic": _ANTHROPIC, "azure_ai": _AZURE})
|
||||
|
||||
|
||||
@pytest.fixture(params=_PROVIDERS.keys())
|
||||
def provider(request: pytest.FixtureRequest) -> _Provider:
|
||||
name: Final[object] = request.param # pyright: ignore[reportAny] # pytest types the fixture param as Any
|
||||
return _PROVIDERS[_NAME.validate_python(name)]
|
||||
|
||||
|
||||
def _json_reply(status: int, payload: Mapping[str, JsonValue]) -> Reply:
|
||||
return Reply(status=status, body=json.dumps(payload).encode())
|
||||
|
||||
|
||||
def _rejected(status: int) -> Reply:
|
||||
return _json_reply(status, {"type": "error", "error": {"type": "invalid_request_error", "message": REJECTION}})
|
||||
|
||||
|
||||
def _counted(request: Request) -> Reply:
|
||||
accepted: Final = accepts_count_body(JSON_OBJECT.validate_json(request.body))
|
||||
return _json_reply(200, {"input_tokens": _COUNT}) if accepted else _rejected(400)
|
||||
|
||||
|
||||
def _rejecting(status: int) -> Callable[[Request], Reply]:
|
||||
def count(_request: Request) -> Reply:
|
||||
return _rejected(status)
|
||||
|
||||
return count
|
||||
|
||||
|
||||
def _holding(held: SimpleQueue[str], release: threading.Event, seconds: float) -> Callable[[Request], Reply]:
|
||||
def hold(request: Request) -> Reply:
|
||||
held.put(request.target)
|
||||
assert release.wait(timeout=seconds), "Held count was never released"
|
||||
return _counted(request)
|
||||
|
||||
return hold
|
||||
|
||||
|
||||
def _peer(provider: _Provider, count: Callable[[Request], Reply] = _counted) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.target == provider.target:
|
||||
return count(request)
|
||||
return _json_reply(404, {"error": f"unscripted target {request.target}"})
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def _listening(deployment: _Deployment, count: Callable[[Request], Reply] = _counted) -> AbstractContextManager[Wire]:
|
||||
return wire_server(_peer(deployment.provider, count), port=deployment.port)
|
||||
|
||||
|
||||
def _reserved_port() -> int:
|
||||
with socket.socket() as reserve:
|
||||
reserve.bind(("127.0.0.1", 0))
|
||||
return _CLIENT_ADDRESS.validate_python(reserve.getsockname())[1]
|
||||
|
||||
|
||||
def _cmdline(process: psutil.Process) -> tuple[str, ...]:
|
||||
try:
|
||||
return _ARGUMENTS.validate_python(process.cmdline())
|
||||
except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess):
|
||||
return ()
|
||||
|
||||
|
||||
def _serves(cmdline: tuple[str, ...], proxy_port: int) -> bool:
|
||||
if _PROXY_MODULE not in cmdline or "--port" not in cmdline:
|
||||
return False
|
||||
return cmdline[cmdline.index("--port") + 1] == str(proxy_port)
|
||||
|
||||
|
||||
def _listens(process: psutil.Process, proxy_port: int) -> bool:
|
||||
return any(
|
||||
connection.status == psutil.CONN_LISTEN and connection.laddr and connection.laddr.port == proxy_port
|
||||
for connection in process.net_connections(kind="tcp")
|
||||
)
|
||||
|
||||
|
||||
def _proxy_workers(proxy_port: int) -> frozenset[int]:
|
||||
(master,) = tuple(process for process in psutil.process_iter() if _serves(_cmdline(process), proxy_port))
|
||||
spawned: Final = frozenset(child.pid for child in master.children() if _listens(child, proxy_port))
|
||||
return spawned or frozenset({master.pid})
|
||||
|
||||
|
||||
def _holder(workers: frozenset[int], client_port: int) -> int | None:
|
||||
def holds(pid: int) -> bool:
|
||||
return any(
|
||||
connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == client_port
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
)
|
||||
|
||||
return next((pid for pid in sorted(workers) if holds(pid)), None)
|
||||
|
||||
|
||||
def _client_port(response: httpx.Response) -> int:
|
||||
stream: Final[object] = response.extensions["network_stream"] # pyright: ignore[reportAny] # httpx types extensions as Any
|
||||
assert isinstance(stream, httpcore.NetworkStream), stream
|
||||
return _CLIENT_ADDRESS.validate_python(stream.get_extra_info("client_addr"))[1]
|
||||
|
||||
|
||||
def _probe(gateway: Gateway, body: Mapping[str, JsonValue], workers: frozenset[int]) -> tuple[int | None, JsonValue]:
|
||||
with httpx.Client(base_url=str(gateway.client.base_url), timeout=30, trust_env=False) as client:
|
||||
response: Final = client.post(
|
||||
"/v1/messages/count_tokens", json=dict(body), headers={"Authorization": f"Bearer {gateway.key}"}
|
||||
)
|
||||
return _holder(workers, _client_port(response)), JSON_OBJECT.validate_json(response.content).get("input_tokens")
|
||||
|
||||
|
||||
def _round(
|
||||
gateway: Gateway, body: Mapping[str, JsonValue], workers: frozenset[int]
|
||||
) -> tuple[tuple[int | None, JsonValue], ...]:
|
||||
def probe(_index: int) -> tuple[int | None, JsonValue]:
|
||||
return _probe(gateway, body, workers)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=_PROBES_PER_ROUND) as pool:
|
||||
return tuple(pool.map(probe, range(_PROBES_PER_ROUND)))
|
||||
|
||||
|
||||
def _settled_on_every_worker(gateway: Gateway, body: Mapping[str, JsonValue]) -> None:
|
||||
proxy_port: Final = gateway.client.base_url.port
|
||||
assert proxy_port is not None, gateway.client.base_url
|
||||
workers: Final = _proxy_workers(proxy_port)
|
||||
eventually(
|
||||
lambda: _round(gateway, body, workers),
|
||||
lambda observed: (
|
||||
frozenset(pid for pid, _ in observed) == workers and all(count == _COUNT for _, count in observed)
|
||||
),
|
||||
seconds=60,
|
||||
)
|
||||
|
||||
|
||||
def _deploy(gateway: Gateway, scenario: Scenario, provider: _Provider) -> _Deployment:
|
||||
port: Final = _reserved_port()
|
||||
model: Final = scenario.model(
|
||||
model_info=None, model=f"{provider.prefix}/{_MODEL}", api_base=f"http://127.0.0.1:{port}", api_key=_API_KEY
|
||||
)
|
||||
with wire_server(_peer(provider), port=port):
|
||||
_settled_on_every_worker(gateway, {"model": model, "messages": [USER]})
|
||||
return _Deployment(provider, port, model)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def deployments() -> Iterator[Mapping[str, _Deployment]]:
|
||||
with gateway_from_environment() as gateway, gateway.scenario() as scenario:
|
||||
yield MappingProxyType({chosen.prefix: _deploy(gateway, scenario, chosen) for chosen in (_ANTHROPIC, _AZURE)})
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def deployment(provider: _Provider, deployments: Mapping[str, _Deployment]) -> _Deployment:
|
||||
return deployments[provider.prefix]
|
||||
|
||||
|
||||
def _count(gateway: Gateway, body: Mapping[str, JsonValue], key: str | None = None) -> httpx.Response:
|
||||
return gateway.request("POST", "/v1/messages/count_tokens", body, key=key)
|
||||
|
||||
|
||||
def _payload(response: httpx.Response) -> dict[str, JsonValue]:
|
||||
assert response.status_code == 200, response.text
|
||||
return JSON_OBJECT.validate_json(response.content)
|
||||
|
||||
|
||||
def _local_count(gateway: Gateway, body: Mapping[str, JsonValue]) -> int:
|
||||
response: Final = gateway.request("POST", "/utils/token_counter", body, params={"call_endpoint": "false"})
|
||||
payload: Final = _payload(response)
|
||||
total: Final = payload["total_tokens"]
|
||||
assert payload["tokenizer_type"] not in _PROVIDER_TOKENIZERS, response.text
|
||||
assert isinstance(total, int) and total > 0 and total != _COUNT, response.text
|
||||
return total
|
||||
|
||||
|
||||
def _bodies(wire: Wire, provider: _Provider) -> tuple[dict[str, JsonValue], ...]:
|
||||
received: Final = wire.drain()
|
||||
for request in received:
|
||||
assert (request.method, request.target) == ("POST", provider.target), request.target
|
||||
assert request.headers["anthropic-version"] == ANTHROPIC_VERSION, request.headers
|
||||
assert TOKEN_COUNTING_BETA in request.headers["anthropic-beta"], request.headers
|
||||
assert request.headers["content-type"] == "application/json", request.headers
|
||||
assert request.headers["x-api-key"] == _API_KEY, request.headers
|
||||
assert (request.headers.get("api-key") == _API_KEY) is provider.azure, request.headers
|
||||
return tuple(JSON_OBJECT.validate_json(request.body) for request in received)
|
||||
|
||||
|
||||
def _clients(stack: ExitStack, base_url: str, count: int) -> tuple[httpx.Client, ...]:
|
||||
return tuple(
|
||||
stack.enter_context(httpx.Client(base_url=base_url, timeout=30, trust_env=False)) for _ in range(count)
|
||||
)
|
||||
|
||||
|
||||
def _counted_on(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue]:
|
||||
response: Final = client.post(
|
||||
"/v1/messages/count_tokens", json=dict(body), headers={"Authorization": f"Bearer {key}"}
|
||||
)
|
||||
return response.status_code, JSON_OBJECT.validate_json(response.content).get("input_tokens")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("case", LIFT_CASES.values(), ids=LIFT_CASES.keys())
|
||||
def test_messages_count_tokens_lifts_the_leading_system_run(
|
||||
deployment: _Deployment, gateway: Gateway, case: LiftCase
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
response: Final = _count(gateway, count_request(deployment.model, case))
|
||||
assert _payload(response) == {"input_tokens": _COUNT}, response.text
|
||||
assert _bodies(peer, deployment.provider) == (expected_count_body(_MODEL, case),)
|
||||
|
||||
|
||||
def test_anthropic_sdk_count_tokens_lifts_the_leading_system_message(deployment: _Deployment, gateway: Gateway) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
counted: Final = anthropic_client(gateway).messages.count_tokens(model=deployment.model, messages=_SDK_MESSAGES)
|
||||
assert counted.input_tokens == _COUNT, counted
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
|
||||
|
||||
|
||||
def test_async_anthropic_sdk_count_tokens_lifts_the_leading_system_message(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
counted: Final = asyncio.run(
|
||||
async_anthropic_client(gateway).messages.count_tokens(model=deployment.model, messages=_SDK_MESSAGES)
|
||||
)
|
||||
assert counted.input_tokens == _COUNT, counted
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
|
||||
|
||||
|
||||
def test_utils_token_counter_call_endpoint_counts_a_leading_system_through_the_provider(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/utils/token_counter",
|
||||
count_request(deployment.model, STRING_CASE),
|
||||
params={"call_endpoint": "true"},
|
||||
)
|
||||
payload: Final = _payload(response)
|
||||
expected_tokenizer: Final = deployment.provider.tokenizer
|
||||
assert (payload["total_tokens"], payload["tokenizer_type"]) == (_COUNT, expected_tokenizer), response.text
|
||||
assert payload["original_response"] == {"input_tokens": _COUNT}, response.text
|
||||
assert (payload["request_model"], payload["model_used"]) == (deployment.model, _MODEL), response.text
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
|
||||
|
||||
|
||||
def test_utils_token_counter_local_mode_never_calls_the_peer_for_a_leading_system(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
assert _local_count(gateway, count_request(deployment.model, STRING_CASE)) > 0
|
||||
assert peer.drain() == ()
|
||||
|
||||
|
||||
def test_responses_input_tokens_lifts_instructions(deployment: _Deployment, gateway: Gateway) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/responses/input_tokens",
|
||||
{"model": deployment.model, "input": USER_TEXT, "instructions": INSTRUCTION},
|
||||
)
|
||||
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _COUNT}, response.text
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
|
||||
|
||||
|
||||
def test_responses_input_tokens_lifts_instructions_ahead_of_a_leading_system_item(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/responses/input_tokens",
|
||||
{"model": deployment.model, "input": [MID_SYSTEM, USER], "instructions": INSTRUCTION},
|
||||
)
|
||||
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _COUNT}, response.text
|
||||
assert _bodies(peer, deployment.provider) == (
|
||||
{"model": _MODEL, "messages": [USER], "system": [LIFTED, {"type": "text", "text": REMINDER}]},
|
||||
)
|
||||
|
||||
|
||||
def test_openai_sdk_input_tokens_lifts_instructions(deployment: _Deployment, gateway: Gateway) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
counted: Final = openai_client(gateway).responses.input_tokens.count(
|
||||
model=deployment.model, input=USER_TEXT, instructions=INSTRUCTION
|
||||
)
|
||||
assert counted.input_tokens == _COUNT, counted
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
|
||||
|
||||
|
||||
def test_async_openai_sdk_input_tokens_lifts_instructions(deployment: _Deployment, gateway: Gateway) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
counted: Final = asyncio.run(
|
||||
async_openai_client(gateway).responses.input_tokens.count(
|
||||
model=deployment.model, input=USER_TEXT, instructions=INSTRUCTION
|
||||
)
|
||||
)
|
||||
assert counted.input_tokens == _COUNT, counted
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
|
||||
|
||||
|
||||
def test_messages_count_tokens_forwards_tools_beside_the_lifted_system(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
response: Final = _count(gateway, {**count_request(deployment.model, STRING_CASE), "tools": TOOLS})
|
||||
assert _payload(response) == {"input_tokens": _COUNT}, response.text
|
||||
assert _bodies(peer, deployment.provider) == ({**_STRING_BODY, "tools": TOOLS},)
|
||||
|
||||
|
||||
def test_messages_count_tokens_keeps_a_mid_conversation_system_in_place(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
body: Final[dict[str, JsonValue]] = {
|
||||
"model": deployment.model,
|
||||
"messages": [LEADING, USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP],
|
||||
}
|
||||
response: Final = _count(gateway, body)
|
||||
assert _payload(response) == {"input_tokens": _COUNT}, response.text
|
||||
assert _bodies(peer, deployment.provider) == (
|
||||
{"model": _MODEL, "messages": [USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP], "system": [LIFTED]},
|
||||
)
|
||||
|
||||
|
||||
def test_messages_count_tokens_falls_back_locally_when_every_message_is_system(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
body: Final[dict[str, JsonValue]] = {"model": deployment.model, "messages": [LEADING]}
|
||||
local: Final = _local_count(gateway, body)
|
||||
response: Final = _count(gateway, body)
|
||||
assert _payload(response) == {"input_tokens": local}, response.text
|
||||
assert _bodies(peer, deployment.provider) == ({"model": _MODEL, "messages": [], "system": [LIFTED]},)
|
||||
|
||||
|
||||
def test_messages_count_tokens_leaves_the_request_untouched_for_a_non_text_system(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
local: Final = _local_count(gateway, count_request(deployment.model, STRING_CASE))
|
||||
response: Final = _count(gateway, {**count_request(deployment.model, STRING_CASE), "system": 5})
|
||||
assert _payload(response) == {"input_tokens": local}, response.text
|
||||
assert _bodies(peer, deployment.provider) == ({"model": _MODEL, "messages": [LEADING, USER], "system": 5},)
|
||||
|
||||
|
||||
def test_messages_count_tokens_repeated_request_lifts_each_time(deployment: _Deployment, gateway: Gateway) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
answers: Final = tuple(
|
||||
_payload(_count(gateway, count_request(deployment.model, STRING_CASE))) for _ in range(2)
|
||||
)
|
||||
assert answers == ({"input_tokens": _COUNT},) * 2
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,) * 2
|
||||
|
||||
|
||||
def test_messages_count_tokens_answers_a_leading_system_without_content_before_any_peer_call(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
body: Final[dict[str, JsonValue]] = {"model": deployment.model, "messages": [{"role": "system"}, USER]}
|
||||
local: Final = _local_count(gateway, body)
|
||||
response: Final = _count(gateway, body)
|
||||
assert _payload(response) == {"input_tokens": local}, response.text
|
||||
assert peer.drain() == ()
|
||||
follow_up: Final = _count(gateway, count_request(deployment.model, STRING_CASE))
|
||||
assert _payload(follow_up) == {"input_tokens": _COUNT}, follow_up.text
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
|
||||
|
||||
|
||||
def test_messages_count_tokens_duplicate_messages_key_lifts_the_last_value(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
first: Final = json.dumps([USER])
|
||||
last: Final = json.dumps([LEADING, USER])
|
||||
response: Final = gateway.client.post(
|
||||
"/v1/messages/count_tokens",
|
||||
content=f'{{"model": "{deployment.model}", "messages": {first}, "messages": {last}}}',
|
||||
headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"},
|
||||
)
|
||||
assert _payload(response) == {"input_tokens": _COUNT}, response.text
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", [400, 403, 404, 500, 503])
|
||||
def test_messages_count_tokens_falls_back_locally_when_the_peer_rejects_the_lifted_body(
|
||||
deployment: _Deployment, gateway: Gateway, status: int
|
||||
) -> None:
|
||||
with _listening(deployment, _rejecting(status)) as peer:
|
||||
body: Final = count_request(deployment.model, STRING_CASE)
|
||||
local: Final = _local_count(gateway, body)
|
||||
response: Final = _count(gateway, body)
|
||||
assert _payload(response) == {"input_tokens": local}, response.text
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
|
||||
|
||||
|
||||
def test_messages_count_tokens_unauthenticated_request_never_reaches_the_peer(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with _listening(deployment) as peer:
|
||||
response: Final = _count(gateway, count_request(deployment.model, STRING_CASE), key="sk-not-a-key")
|
||||
assert response.status_code == 401, response.text
|
||||
assert peer.drain() == ()
|
||||
|
||||
|
||||
def test_peer_outage_between_concurrent_waves_falls_back_then_recovers(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
with ExitStack() as stack:
|
||||
clients: Final = _clients(stack, str(gateway.client.base_url), 8)
|
||||
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
|
||||
body: Final = count_request(deployment.model, STRING_CASE)
|
||||
local: Final = _local_count(gateway, body)
|
||||
|
||||
def count(client: httpx.Client) -> tuple[int, JsonValue]:
|
||||
return _counted_on(client, gateway.key, body)
|
||||
|
||||
with _listening(deployment) as peer:
|
||||
assert tuple(pool.map(count, clients)) == ((200, _COUNT),) * len(clients)
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,) * len(clients)
|
||||
assert tuple(pool.map(count, clients)) == ((200, local),) * len(clients)
|
||||
with _listening(deployment) as revived:
|
||||
assert tuple(pool.map(count, clients)) == ((200, _COUNT),) * len(clients)
|
||||
assert _bodies(revived, deployment.provider) == (_STRING_BODY,) * len(clients)
|
||||
|
||||
|
||||
def test_slow_peer_holds_concurrent_lifted_counts_without_stalling_the_proxy(
|
||||
deployment: _Deployment, gateway: Gateway
|
||||
) -> None:
|
||||
held: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
release: Final = threading.Event()
|
||||
with ExitStack() as stack:
|
||||
clients: Final = _clients(stack, str(gateway.client.base_url), 6)
|
||||
peer: Final = stack.enter_context(_listening(deployment, _holding(held, release, 20)))
|
||||
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
|
||||
stack.callback(release.set)
|
||||
body: Final = count_request(deployment.model, STRING_CASE)
|
||||
futures: Final = tuple(pool.submit(_counted_on, client, gateway.key, body) for client in clients)
|
||||
eventually(held.qsize, lambda size: size == len(clients), seconds=30)
|
||||
assert gateway.request("GET", "/health/liveliness").status_code == 200
|
||||
assert _local_count(gateway, body) > 0
|
||||
assert not any(future.done() for future in futures)
|
||||
release.set()
|
||||
assert tuple(future.result(timeout=30) for future in futures) == ((200, _COUNT),) * len(clients)
|
||||
assert _bodies(peer, deployment.provider) == (_STRING_BODY,) * len(clients)
|
||||
|
||||
|
||||
def test_chat_completions_on_the_same_deployment_keeps_a_mid_conversation_system_in_place(
|
||||
deployments: Mapping[str, _Deployment], gateway: Gateway
|
||||
) -> None:
|
||||
identity: Final = f"msg_{uuid.uuid4().hex}"
|
||||
anthropic_deployment: Final = deployments[_ANTHROPIC.prefix]
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert (request.method, request.target) == ("POST", "/v1/messages"), request.target
|
||||
return _json_reply(
|
||||
200,
|
||||
{
|
||||
"id": identity,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": _MODEL,
|
||||
"content": [{"type": "text", "text": "done"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": None,
|
||||
"usage": {"input_tokens": 12, "output_tokens": 3},
|
||||
},
|
||||
)
|
||||
|
||||
with wire_server(respond, port=anthropic_deployment.port) as peer:
|
||||
reminder: Final[dict[str, JsonValue]] = {
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL},
|
||||
{"type": "text", "text": ""},
|
||||
IMAGE_PART,
|
||||
],
|
||||
}
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": anthropic_deployment.model,
|
||||
"max_tokens": 16,
|
||||
"messages": [LEADING, USER, reminder, ASSISTANT, FOLLOW_UP],
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
(sent,) = peer.drain()
|
||||
body: Final = JSON_OBJECT.validate_json(sent.body)
|
||||
assert body["system"] == [LIFTED], body
|
||||
assert body["messages"] == [
|
||||
{"role": "user", "content": [{"type": "text", "text": USER_TEXT}]},
|
||||
{"role": "system", "content": [{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL}]},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "One."}]},
|
||||
{"role": "user", "content": [{"type": "text", "text": "Again"}]},
|
||||
], body
|
||||
|
|
@ -11,7 +11,7 @@ from contextlib import ExitStack
|
|||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from typing import Final, cast
|
||||
|
||||
import anthropic
|
||||
import httpx
|
||||
|
|
@ -24,6 +24,25 @@ from integration._support.bedrock_runtime_peer import respond as runtime_generat
|
|||
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value
|
||||
from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from integration.providers._count_tokens_system_lift import (
|
||||
ASSISTANT,
|
||||
FOLLOW_UP,
|
||||
INSTRUCTION,
|
||||
LEADING,
|
||||
LIFT_CASES,
|
||||
LIFTED,
|
||||
MID_SYSTEM,
|
||||
REMINDER,
|
||||
STRING_CASE,
|
||||
USER,
|
||||
USER_TEXT,
|
||||
LiftCase,
|
||||
accepts_count_body,
|
||||
async_openai_client,
|
||||
count_request,
|
||||
expected_count_body,
|
||||
openai_client,
|
||||
)
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
pytestmark = pytest.mark.timeout(2 * graceful_stop_seconds() + 120)
|
||||
|
|
@ -81,6 +100,10 @@ def _mantle_body(**fields: JsonValue) -> dict[str, JsonValue]:
|
|||
|
||||
|
||||
_MANTLE_BARE: Final = _mantle_body()
|
||||
_MANTLE_LIFTED: Final = _mantle_body(system=[LIFTED])
|
||||
_LIFT_SDK_MESSAGES: Final = cast(
|
||||
list[MessageParam], [LEADING, USER]
|
||||
) # cast-ok: the SDK types reject the role the proxy lifts
|
||||
_MANTLE_FULL: Final = _mantle_body(system=_SYSTEM, tools=_TOOLS)
|
||||
|
||||
|
||||
|
|
@ -138,25 +161,8 @@ def _rejecting(status: int) -> Callable[[Request], Reply]:
|
|||
return count
|
||||
|
||||
|
||||
def _anthropic_message(message: JsonValue) -> bool:
|
||||
return isinstance(message, dict) and message.get("role") in ("user", "assistant")
|
||||
|
||||
|
||||
def _anthropic_tool(tool: JsonValue) -> bool:
|
||||
return isinstance(tool, dict) and isinstance(tool.get("name"), str) and isinstance(tool.get("input_schema"), dict)
|
||||
|
||||
|
||||
def _strict(request: Request) -> Reply:
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
messages: Final = body.get("messages")
|
||||
tools: Final = body.get("tools", [])
|
||||
accepted: Final = (
|
||||
isinstance(messages, list)
|
||||
and all(map(_anthropic_message, messages))
|
||||
and isinstance(body.get("system", ""), (str, list))
|
||||
and isinstance(tools, list)
|
||||
and all(map(_anthropic_tool, tools))
|
||||
)
|
||||
accepted: Final = accepts_count_body(_JSON_OBJECT.validate_json(request.body))
|
||||
return _mantle_counted(request) if accepted else _rejected(400)
|
||||
|
||||
|
||||
|
|
@ -577,7 +583,7 @@ def test_bedrock_passthrough_count_tokens_still_answers_the_runtime_rejection(
|
|||
assert mantle.drain() == ()
|
||||
|
||||
|
||||
def test_responses_input_tokens_with_instructions_still_counts_locally(
|
||||
def test_responses_input_tokens_with_instructions_counts_through_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
|
|
@ -592,14 +598,322 @@ def test_responses_input_tokens_with_instructions_still_counts_locally(
|
|||
"/v1/responses/input_tokens",
|
||||
{"model": model, "input": "Count this message", "instructions": "Be terse"},
|
||||
)
|
||||
payload: Final = _payload(response)
|
||||
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _MANTLE_COUNT}, response.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
(sent,) = _mantle_bodies(mantle)
|
||||
messages: Final = sent["messages"]
|
||||
assert isinstance(messages, list) and messages[0] == {"role": "system", "content": "Be terse"}, sent
|
||||
assert "system" not in sent, sent
|
||||
local: Final = _local_count(gateway, {"model": model, "messages": messages})
|
||||
assert payload == {"object": "response.input_tokens", "input_tokens": local}, response.text
|
||||
assert _mantle_bodies(mantle) == (_mantle_body(system=[{"type": "text", "text": "Be terse"}]),)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("case", LIFT_CASES.values(), ids=LIFT_CASES.keys())
|
||||
def test_messages_count_tokens_lifts_the_leading_system_run_for_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int, case: LiftCase
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
response: Final = _count(gateway, count_request(model, case))
|
||||
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (expected_count_body(_OPUS_BASE, case),)
|
||||
|
||||
|
||||
def test_anthropic_sdk_count_tokens_lifts_the_leading_system_through_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
counted: Final = _anthropic_client(gateway).messages.count_tokens(model=model, messages=_LIFT_SDK_MESSAGES)
|
||||
assert counted.input_tokens == _MANTLE_COUNT, counted
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
|
||||
|
||||
|
||||
def test_async_anthropic_sdk_count_tokens_lifts_the_leading_system_through_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
counted: Final = asyncio.run(
|
||||
_async_anthropic_client(gateway).messages.count_tokens(model=model, messages=_LIFT_SDK_MESSAGES)
|
||||
)
|
||||
assert counted.input_tokens == _MANTLE_COUNT, counted
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
|
||||
|
||||
|
||||
def test_utils_token_counter_call_endpoint_counts_a_leading_system_through_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
response: Final = gateway.request(
|
||||
"POST", "/utils/token_counter", count_request(model, STRING_CASE), params={"call_endpoint": "true"}
|
||||
)
|
||||
payload: Final = _payload(response)
|
||||
assert (payload["total_tokens"], payload["tokenizer_type"]) == (_MANTLE_COUNT, "bedrock_mantle_api")
|
||||
assert payload["original_response"] == {"input_tokens": _MANTLE_COUNT}, response.text
|
||||
assert (payload["request_model"], payload["model_used"]) == (model, _OPUS), response.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
|
||||
|
||||
|
||||
def test_openai_sdk_input_tokens_lifts_instructions_through_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
counted: Final = openai_client(gateway).responses.input_tokens.count(
|
||||
model=model, input=USER_TEXT, instructions=INSTRUCTION
|
||||
)
|
||||
assert counted.input_tokens == _MANTLE_COUNT, counted
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
|
||||
|
||||
|
||||
def test_async_openai_sdk_input_tokens_lifts_instructions_through_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
counted: Final = asyncio.run(
|
||||
async_openai_client(gateway).responses.input_tokens.count(
|
||||
model=model, input=USER_TEXT, instructions=INSTRUCTION
|
||||
)
|
||||
)
|
||||
assert counted.input_tokens == _MANTLE_COUNT, counted
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
|
||||
|
||||
|
||||
def test_responses_input_tokens_lifts_instructions_ahead_of_a_leading_system_item_through_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/responses/input_tokens",
|
||||
{"model": model, "input": [MID_SYSTEM, USER], "instructions": INSTRUCTION},
|
||||
)
|
||||
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _MANTLE_COUNT}, response.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_mantle_body(system=[LIFTED, {"type": "text", "text": REMINDER}]),)
|
||||
|
||||
|
||||
def test_messages_count_tokens_forwards_tools_beside_the_lifted_system_to_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
response: Final = _count(gateway, {**count_request(model, STRING_CASE), "tools": _TOOLS})
|
||||
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_mantle_body(system=[LIFTED], tools=_TOOLS),)
|
||||
|
||||
|
||||
def test_messages_count_tokens_keeps_a_mid_conversation_system_in_place_for_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
body: Final[dict[str, JsonValue]] = {
|
||||
"model": model,
|
||||
"messages": [LEADING, USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP],
|
||||
}
|
||||
response: Final = _count(gateway, body)
|
||||
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (
|
||||
_mantle_body(messages=[USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP], system=[LIFTED]),
|
||||
)
|
||||
|
||||
|
||||
def test_messages_count_tokens_falls_back_locally_when_every_message_is_system_for_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
body: Final[dict[str, JsonValue]] = {"model": model, "messages": [LEADING]}
|
||||
local: Final = _local_count(gateway, body)
|
||||
response: Final = _count(gateway, body)
|
||||
assert _payload(response) == {"input_tokens": local}, response.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_mantle_body(messages=[], system=[LIFTED]),)
|
||||
|
||||
|
||||
def test_messages_count_tokens_leaves_a_leading_system_in_place_beside_a_non_text_system_for_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
local: Final = _local_count(gateway, count_request(model, STRING_CASE))
|
||||
response: Final = _count(gateway, {**count_request(model, STRING_CASE), "system": 5})
|
||||
assert _payload(response) == {"input_tokens": local}, response.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_mantle_body(messages=[LEADING, USER], system=5),)
|
||||
|
||||
|
||||
def test_messages_count_tokens_repeated_request_lifts_the_leading_system_each_time_for_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
answers: Final = tuple(_payload(_count(gateway, count_request(model, STRING_CASE))) for _ in range(2))
|
||||
assert answers == ({"input_tokens": _MANTLE_COUNT},) * 2
|
||||
assert _runtime_count_targets(runtime) == (f"/model/{_OPUS_BASE}/count-tokens",) * 2
|
||||
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,) * 2
|
||||
|
||||
|
||||
def test_messages_count_tokens_answers_a_leading_system_without_content_before_any_mantle_call(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
body: Final[dict[str, JsonValue]] = {"model": model, "messages": [{"role": "system"}, USER]}
|
||||
local: Final = _local_count(gateway, body)
|
||||
response: Final = _count(gateway, body)
|
||||
assert _payload(response) == {"input_tokens": local}, response.text
|
||||
assert mantle.drain() == ()
|
||||
follow_up: Final = _count(gateway, count_request(model, STRING_CASE))
|
||||
assert _payload(follow_up) == {"input_tokens": _MANTLE_COUNT}, follow_up.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
|
||||
|
||||
|
||||
def test_messages_count_tokens_duplicate_messages_key_lifts_the_last_value_for_mantle(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with (
|
||||
wire_server(_runtime) as runtime,
|
||||
wire_server(_mantle(_strict), port=mantle_port) as mantle,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
first: Final = json.dumps([USER])
|
||||
last: Final = json.dumps([LEADING, USER])
|
||||
response: Final = gateway.client.post(
|
||||
"/v1/messages/count_tokens",
|
||||
content=f'{{"model": "{model}", "messages": {first}, "messages": {last}}}',
|
||||
headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"},
|
||||
)
|
||||
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
|
||||
|
||||
|
||||
def test_disabled_token_counter_counts_a_leading_system_through_mantle(
|
||||
gateway: Gateway, mantle_port: int, tmp_path: Path
|
||||
) -> None:
|
||||
with ExitStack() as stack:
|
||||
runtime: Final = stack.enter_context(wire_server(_runtime))
|
||||
config: Final = _owned_config(
|
||||
tmp_path / "disabled-token-counter-lift.yaml", runtime.url, {"disable_token_counter": True}
|
||||
)
|
||||
owned: Final = stack.enter_context(
|
||||
owned_proxy_process(
|
||||
gateway,
|
||||
tmp_path,
|
||||
_mantle_environment(mantle_port),
|
||||
config=config,
|
||||
workers=2,
|
||||
remove_environment=_INHERITED_BEARER,
|
||||
)
|
||||
)
|
||||
with wire_server(_mantle(_strict), port=mantle_port) as strict:
|
||||
counted: Final = _count(owned.gateway, count_request(_OWNED_OPUS, STRING_CASE))
|
||||
assert _payload(counted) == {"input_tokens": _MANTLE_COUNT}, counted.text
|
||||
assert _mantle_bodies(strict) == (_MANTLE_LIFTED,)
|
||||
assert len(_runtime_count_targets(runtime)) == 1
|
||||
|
||||
|
||||
def test_mantle_outage_between_concurrent_lifted_waves_falls_back_then_recovers(
|
||||
counting_proxy: OwnedProxy, mantle_port: int
|
||||
) -> None:
|
||||
gateway: Final = counting_proxy.gateway
|
||||
with ExitStack() as stack:
|
||||
clients: Final = _clients(stack, str(gateway.client.base_url), 8)
|
||||
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
|
||||
runtime: Final = stack.enter_context(wire_server(_runtime))
|
||||
scenario: Final = stack.enter_context(gateway.scenario())
|
||||
model: Final = _deployment(scenario, runtime.url)
|
||||
body: Final = count_request(model, STRING_CASE)
|
||||
local: Final = _local_count(gateway, body)
|
||||
|
||||
def count(client: httpx.Client) -> tuple[int, JsonValue]:
|
||||
return _counted_on(client, gateway.key, body)
|
||||
|
||||
with wire_server(_mantle(_strict), port=mantle_port) as mantle:
|
||||
assert tuple(pool.map(count, clients)) == ((200, _MANTLE_COUNT),) * len(clients)
|
||||
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,) * len(clients)
|
||||
assert tuple(pool.map(count, clients)) == ((200, local),) * len(clients)
|
||||
with wire_server(_mantle(_strict), port=mantle_port) as revived:
|
||||
assert tuple(pool.map(count, clients)) == ((200, _MANTLE_COUNT),) * len(clients)
|
||||
assert _mantle_bodies(revived) == (_MANTLE_LIFTED,) * len(clients)
|
||||
assert len(_runtime_count_targets(runtime)) == 3 * len(clients)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", [400, 403, 404, 500, 503])
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import litellm
|
|||
from litellm.litellm_core_utils.prompt_templates.common_utils import encrypted_reasoning_signature
|
||||
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import (
|
||||
CONVERTED_SYSTEM_NOTE,
|
||||
anthropic_system_blocks,
|
||||
place_mid_conversation_system,
|
||||
split_leading_system_run,
|
||||
)
|
||||
|
|
@ -390,3 +391,23 @@ def test_flagged_placement_keeps_a_system_before_an_assistant_turn_whose_empty_t
|
|||
)
|
||||
|
||||
assert _roles(placed) == ["user", "system", "assistant", "user"]
|
||||
|
||||
|
||||
def test_anthropic_system_blocks_keeps_text_parts_with_their_cache_control_and_drops_the_rest():
|
||||
run = [
|
||||
{"role": "system", "content": "one", "cache_control": {"type": "ephemeral", "ttl": "1h"}},
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": "two", "cache_control": {"type": "ephemeral"}},
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "image_url", "image_url": {"url": "data:image/png;base64,aW1hZ2U="}},
|
||||
],
|
||||
},
|
||||
{"role": "system", "content": ""},
|
||||
]
|
||||
|
||||
assert anthropic_system_blocks(run) == (
|
||||
{"type": "text", "text": "one", "cache_control": {"type": "ephemeral", "ttl": "1h"}},
|
||||
{"type": "text", "text": "two", "cache_control": {"type": "ephemeral"}},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -329,3 +329,115 @@ async def test_remote_image_fetch_keeps_counting_handler_event_loop_responsive(
|
|||
]}],
|
||||
}
|
||||
assert messages == original
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"config_type", (AnthropicCountTokensConfig, AzureAIAnthropicCountTokensConfig)
|
||||
)
|
||||
def test_count_lifts_the_leading_system_run_into_system(
|
||||
config_type: type[AnthropicCountTokensConfig],
|
||||
) -> None:
|
||||
"""A Responses ``instructions`` arrives as a leading system-role message. count_tokens answers 400
|
||||
on that role at the head of ``messages`` and only takes the initial prompt in ``system``, so the
|
||||
leading run moves there with its cache_control, empty text dropped, and a later reminder stays."""
|
||||
cache_control: Final[dict[str, JsonValue]] = {"type": "ephemeral"}
|
||||
messages: Final[list[dict[str, JsonValue]]] = [
|
||||
{"role": "system", "content": "Be terse", "cache_control": cache_control},
|
||||
{"role": "system", "content": [{"type": "text", "text": "Answer in French"}, {"type": "text", "text": ""}]},
|
||||
{"role": "user", "content": "Hello, how are you?"},
|
||||
{"role": "system", "content": "later reminder"},
|
||||
{"role": "assistant", "content": "Bonjour."},
|
||||
]
|
||||
original: Final = deepcopy(messages)
|
||||
result: Final = config_type().transform_request_to_count_tokens(model="claude-opus-5-5", messages=messages)
|
||||
|
||||
assert result == {
|
||||
"model": "claude-opus-5-5",
|
||||
"system": [
|
||||
{"type": "text", "text": "Be terse", "cache_control": cache_control},
|
||||
{"type": "text", "text": "Answer in French"},
|
||||
],
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, how are you?"},
|
||||
{"role": "system", "content": "later reminder"},
|
||||
{"role": "assistant", "content": "Bonjour."},
|
||||
],
|
||||
}
|
||||
assert messages == original
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("system", "expected_system"),
|
||||
(
|
||||
(None, [{"type": "text", "text": "Be terse"}]),
|
||||
("", [{"type": "text", "text": "Be terse"}]),
|
||||
("Answer in French", [{"type": "text", "text": "Answer in French"}, {"type": "text", "text": "Be terse"}]),
|
||||
(
|
||||
[{"type": "text", "text": "Answer in French", "cache_control": {"type": "ephemeral"}}],
|
||||
[
|
||||
{"type": "text", "text": "Answer in French", "cache_control": {"type": "ephemeral"}},
|
||||
{"type": "text", "text": "Be terse"},
|
||||
],
|
||||
),
|
||||
),
|
||||
ids=["absent", "empty", "string", "blocks"],
|
||||
)
|
||||
def test_count_keeps_the_callers_system_ahead_of_the_lifted_run(
|
||||
system: JsonValue, expected_system: list[dict[str, JsonValue]]
|
||||
) -> None:
|
||||
result: Final = AnthropicCountTokensConfig().transform_request_to_count_tokens(
|
||||
model="claude-opus-5-5",
|
||||
messages=[{"role": "system", "content": "Be terse"}, {"role": "user", "content": "hi"}],
|
||||
system=system,
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"model": "claude-opus-5-5",
|
||||
"system": expected_system,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
|
||||
|
||||
def test_count_leaves_a_non_text_system_and_its_messages_as_sent() -> None:
|
||||
"""A malformed ``system`` is the provider's to reject, so nothing is rearranged around it."""
|
||||
messages: Final[list[dict[str, JsonValue]]] = [
|
||||
{"role": "system", "content": "Be terse"},
|
||||
{"role": "user", "content": "hi"},
|
||||
]
|
||||
result: Final = AnthropicCountTokensConfig().transform_request_to_count_tokens(
|
||||
model="claude-opus-5-5", messages=messages, system=5
|
||||
)
|
||||
|
||||
assert result == {"model": "claude-opus-5-5", "system": 5, "messages": messages}
|
||||
|
||||
|
||||
def test_count_drops_a_leading_system_message_without_text() -> None:
|
||||
result: Final = AnthropicCountTokensConfig().transform_request_to_count_tokens(
|
||||
model="claude-opus-5-5",
|
||||
messages=[{"role": "system", "content": ""}, {"role": "user", "content": "hi"}],
|
||||
)
|
||||
|
||||
assert result == {"model": "claude-opus-5-5", "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handler_sends_the_leading_system_run_as_system_not_as_a_message(httpx_transport_clients):
|
||||
"""The wire body is what the provider judges: ``system`` carries the prompt and no message has
|
||||
``role: "system"``, so a Responses ``instructions`` is counted by Anthropic instead of 400ing."""
|
||||
with respx.mock:
|
||||
route = respx.post("https://gateway.example/v1/messages/count_tokens").mock(
|
||||
return_value=httpx.Response(200, json={"input_tokens": 21})
|
||||
)
|
||||
result = await AnthropicCountTokensHandler().handle_count_tokens_request(
|
||||
model="claude-opus-5-5",
|
||||
messages=[{"role": "system", "content": "Be terse"}, {"role": "user", "content": "Hello, how are you?"}],
|
||||
auth_header={"x-api-key": "sk-ant-api03-test-key"},
|
||||
api_base="https://gateway.example",
|
||||
)
|
||||
|
||||
assert result == {"input_tokens": 21}
|
||||
assert TypeAdapter(dict[str, JsonValue]).validate_json(route.calls.last.request.content) == {
|
||||
"model": "claude-opus-5-5",
|
||||
"system": [{"type": "text", "text": "Be terse"}],
|
||||
"messages": [{"role": "user", "content": "Hello, how are you?"}],
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue