From 164ed775c83542bb013a158463f0bb0cce8f6fee Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Sat, 9 May 2026 14:43:03 +0200 Subject: [PATCH] support /message/count_tokens endpoint --- routstr/payment/cost_calculation.py | 1 + routstr/upstream/base.py | 13 ++ routstr/upstream/count_tokens.py | 108 +++++++++++ tests/unit/test_count_tokens_local.py | 177 +++++++++++++++++++ tests/unit/test_messages_litellm_dispatch.py | 62 ++++--- 5 files changed, 336 insertions(+), 25 deletions(-) create mode 100644 routstr/upstream/count_tokens.py create mode 100644 tests/unit/test_count_tokens_local.py diff --git a/routstr/payment/cost_calculation.py b/routstr/payment/cost_calculation.py index 2dbf9a93..4fb61f49 100644 --- a/routstr/payment/cost_calculation.py +++ b/routstr/payment/cost_calculation.py @@ -36,6 +36,7 @@ class CostDataError(BaseModel): async def calculate_cost( response_data: dict, max_cost: int, session: AsyncSession ) -> CostData | MaxCostData | CostDataError: + print(response_data) """Calculate the cost of an API request based on token usage. Args: diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 644bb25b..2f3a6221 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -42,6 +42,7 @@ from ..payment.models import ( from ..payment.price import sats_usd_price from ..wallet import recieve_token, send_token from . import messages_dispatch +from .count_tokens import count_tokens_locally from .litellm_routing import detect_litellm_prefix logger = get_logger(__name__) @@ -2159,6 +2160,12 @@ class BaseUpstreamProvider: """ path = self.normalize_request_path(path, model_obj) + if ( + path.endswith("messages/count_tokens") + and not self.supports_anthropic_messages + ): + return count_tokens_locally(request_body, model_obj) + if ( path.endswith("messages") and not path.endswith("count_tokens") @@ -3357,6 +3364,12 @@ class BaseUpstreamProvider: request_body = await request.body() + if ( + path.endswith("messages/count_tokens") + and not self.supports_anthropic_messages + ): + return count_tokens_locally(request_body, model_obj) + if ( path.endswith("messages") and not path.endswith("count_tokens") diff --git a/routstr/upstream/count_tokens.py b/routstr/upstream/count_tokens.py new file mode 100644 index 00000000..d114561c --- /dev/null +++ b/routstr/upstream/count_tokens.py @@ -0,0 +1,108 @@ +"""Local handling of Anthropic ``/v1/messages/count_tokens`` for upstreams +that do not natively expose the endpoint. + +Most non-Anthropic upstreams (OpenAI-compat, Gemini OpenAI-compat, +OpenRouter chat-completions, generic providers) return 400/404 when asked +to ``POST /messages/count_tokens``. Claude Code and other Anthropic SDK +clients call this endpoint before each turn to size context windows and +trigger compaction, so a failure breaks the whole chat. + +We answer locally. ``litellm.token_counter`` understands the Anthropic +message shape and the per-model tokenizers, so we prefer it. If it raises +(unknown model, encoding lookup failure, ...), we fall back to the +project's own ``estimate_tokens`` heuristic, which is always defined and +never raises. +""" + +from __future__ import annotations + +import json +from typing import Any + +import litellm +from fastapi.responses import Response + +from ..core import get_logger +from ..payment.helpers import estimate_tokens +from ..payment.models import Model + +logger = get_logger(__name__) + + +def _parse_request_body(request_body: bytes | None) -> dict[str, Any]: + if not request_body: + return {} + try: + parsed = json.loads(request_body) + except (ValueError, TypeError): + return {} + return parsed if isinstance(parsed, dict) else {} + + +def _count_with_litellm(model: str, body: dict[str, Any]) -> int: + messages = body.get("messages") + if not isinstance(messages, list): + messages = [] + + system = body.get("system") + if isinstance(system, str) and system: + messages = [{"role": "system", "content": system}, *messages] + elif isinstance(system, list): + text = "".join( + block.get("text", "") + for block in system + if isinstance(block, dict) and block.get("type") == "text" + ) + if text: + messages = [{"role": "system", "content": text}, *messages] + + tools = body.get("tools") if isinstance(body.get("tools"), list) else None + + return int( + litellm.token_counter( + model=model, + messages=messages, + tools=tools, + ) + ) + + +def count_tokens_locally( + request_body: bytes | None, + model_obj: Model | None, +) -> Response: + """Return an Anthropic-compatible count_tokens response without + touching the upstream. Always returns 200; never raises.""" + body = _parse_request_body(request_body) + + model_name = "" + if model_obj is not None: + model_name = model_obj.forwarded_model_id or model_obj.id or "" + if not model_name: + body_model = body.get("model") + if isinstance(body_model, str): + model_name = body_model + + input_tokens: int + try: + input_tokens = _count_with_litellm(model_name, body) + except Exception as exc: + messages = body.get("messages") + fallback_messages = messages if isinstance(messages, list) else [] + input_tokens = estimate_tokens(fallback_messages) + logger.debug( + "litellm token_counter failed; using local estimator", + extra={ + "model": model_name, + "error": str(exc), + "error_type": type(exc).__name__, + "estimated_tokens": input_tokens, + }, + ) + + payload = {"input_tokens": max(0, int(input_tokens))} + return Response( + content=json.dumps(payload).encode(), + status_code=200, + media_type="application/json", + ) diff --git a/tests/unit/test_count_tokens_local.py b/tests/unit/test_count_tokens_local.py new file mode 100644 index 00000000..6435948b --- /dev/null +++ b/tests/unit/test_count_tokens_local.py @@ -0,0 +1,177 @@ +"""Unit tests for the local count_tokens shim. + +The shim runs whenever an upstream that does not support Anthropic's +``/v1/messages`` endpoint is asked for a token count. It must always +return a 200 JSON ``{"input_tokens": N}`` response and must never raise. +""" + +from __future__ import annotations + +import json +from typing import Any +from unittest.mock import patch + +from routstr.payment.models import Architecture, Model, Pricing +from routstr.upstream import count_tokens as count_tokens_module +from routstr.upstream.count_tokens import count_tokens_locally + + +def _make_model(model_id: str = "anthropic/claude-3-5-sonnet") -> Model: + pricing = Pricing(prompt=0.000003, completion=0.000015) + architecture = Architecture( + modality="text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="cl100k_base", + instruct_type=None, + ) + return Model( + id=model_id, + name=model_id, + created=0, + description="", + context_length=200_000, + architecture=architecture, + pricing=pricing, + ) + + +def _body(payload: dict[str, Any]) -> bytes: + return json.dumps(payload).encode() + + +def _read_payload(response: Any) -> dict[str, Any]: + body = response.body if isinstance(response.body, bytes) else bytes(response.body) + return json.loads(body.decode()) + + +def test_returns_input_tokens_for_simple_messages() -> None: + model = _make_model() + request_body = _body( + { + "model": model.id, + "messages": [{"role": "user", "content": "hello world"}], + } + ) + + response = count_tokens_locally(request_body, model) + + assert response.status_code == 200 + assert response.media_type == "application/json" + payload = _read_payload(response) + assert "input_tokens" in payload + assert isinstance(payload["input_tokens"], int) + assert payload["input_tokens"] >= 0 + + +def test_falls_back_to_estimator_when_litellm_raises() -> None: + model = _make_model() + request_body = _body( + { + "model": model.id, + "messages": [{"role": "user", "content": "this is a longer message"}], + } + ) + + with patch.object( + count_tokens_module, + "_count_with_litellm", + side_effect=RuntimeError("boom"), + ): + response = count_tokens_locally(request_body, model) + + assert response.status_code == 200 + payload = _read_payload(response) + assert payload["input_tokens"] >= 1 + + +def test_handles_missing_model_object() -> None: + request_body = _body( + { + "model": "anthropic/claude-3-5-sonnet", + "messages": [{"role": "user", "content": "hi"}], + } + ) + + response = count_tokens_locally(request_body, None) + + assert response.status_code == 200 + payload = _read_payload(response) + assert payload["input_tokens"] >= 0 + + +def test_handles_empty_request_body() -> None: + response = count_tokens_locally(b"", _make_model()) + + assert response.status_code == 200 + payload = _read_payload(response) + assert payload["input_tokens"] >= 0 + + +def test_handles_malformed_json() -> None: + response = count_tokens_locally(b"not-json", _make_model()) + + assert response.status_code == 200 + payload = _read_payload(response) + assert payload["input_tokens"] >= 0 + + +def test_includes_system_prompt_in_count() -> None: + model = _make_model() + short = _body( + { + "model": model.id, + "messages": [{"role": "user", "content": "hi"}], + } + ) + with_system = _body( + { + "model": model.id, + "system": "You are a helpful assistant with a long preamble " * 10, + "messages": [{"role": "user", "content": "hi"}], + } + ) + + short_count = _read_payload(count_tokens_locally(short, model))["input_tokens"] + long_count = _read_payload(count_tokens_locally(with_system, model))["input_tokens"] + + assert long_count > short_count + + +def test_supports_anthropic_system_block_list() -> None: + model = _make_model() + request_body = _body( + { + "model": model.id, + "system": [{"type": "text", "text": "be terse" * 50}], + "messages": [{"role": "user", "content": "ok"}], + } + ) + + response = count_tokens_locally(request_body, model) + + payload = _read_payload(response) + assert payload["input_tokens"] > 0 + + +def test_uses_forwarded_model_id_when_present() -> None: + model = _make_model("anthropic/claude-3-5-sonnet") + model.forwarded_model_id = "claude-3-5-sonnet-20241022" + request_body = _body( + { + "model": "ignored", + "messages": [{"role": "user", "content": "hi"}], + } + ) + + captured: dict[str, Any] = {} + + def _capture(model_name: str, body: dict[str, Any]) -> int: + captured["model"] = model_name + return 7 + + with patch.object(count_tokens_module, "_count_with_litellm", side_effect=_capture): + response = count_tokens_locally(request_body, model) + + assert captured["model"] == "claude-3-5-sonnet-20241022" + assert _read_payload(response)["input_tokens"] == 7 diff --git a/tests/unit/test_messages_litellm_dispatch.py b/tests/unit/test_messages_litellm_dispatch.py index 20482bb1..b9c0a767 100644 --- a/tests/unit/test_messages_litellm_dispatch.py +++ b/tests/unit/test_messages_litellm_dispatch.py @@ -908,7 +908,7 @@ async def test_forward_request_skips_litellm_when_provider_supports_messages() - @pytest.mark.asyncio -async def test_forward_request_skips_litellm_for_count_tokens() -> None: +async def test_forward_request_handles_count_tokens_locally() -> None: provider = _make_provider() key = _make_key() model = _make_model() @@ -923,19 +923,25 @@ async def test_forward_request_skips_litellm_for_count_tokens() -> None: with patch.object( provider, "prepare_request_body", - side_effect=RuntimeError("stop here"), + side_effect=AssertionError("upstream should not be called"), ): - with pytest.raises(RuntimeError, match="stop here"): - await provider.forward_request( - request=request, - path="messages/count_tokens", - headers={}, - request_body=_anthropic_request_body(), - key=key, - max_cost_for_model=10_000, - session=session, - model_obj=model, - ) + response = await provider.forward_request( + request=request, + path="messages/count_tokens", + headers={}, + request_body=_anthropic_request_body(), + key=key, + max_cost_for_model=10_000, + session=session, + model_obj=model, + ) + + assert response.status_code == 200 + body = response.body if isinstance(response.body, bytes) else bytes(response.body) + payload = json.loads(body.decode()) + assert "input_tokens" in payload + assert isinstance(payload["input_tokens"], int) + assert payload["input_tokens"] >= 0 # --------------------------------------------------------------------------- @@ -1010,7 +1016,7 @@ async def test_forward_x_cashu_request_skips_litellm_when_native_messages() -> N @pytest.mark.asyncio -async def test_forward_x_cashu_request_skips_litellm_for_count_tokens() -> None: +async def test_forward_x_cashu_request_handles_count_tokens_locally() -> None: provider = _make_provider() model = _make_model() request = _make_request() @@ -1024,18 +1030,24 @@ async def test_forward_x_cashu_request_skips_litellm_for_count_tokens() -> None: with patch.object( provider, "prepare_request_body", - side_effect=RuntimeError("stop here"), + side_effect=AssertionError("upstream should not be called"), ): - with pytest.raises(RuntimeError, match="stop here"): - await provider.forward_x_cashu_request( - request=request, - path="v1/messages/count_tokens", - headers={}, - amount=5_000, - unit="sat", - max_cost_for_model=10_000, - model_obj=model, - ) + response = await provider.forward_x_cashu_request( + request=request, + path="v1/messages/count_tokens", + headers={}, + amount=5_000, + unit="sat", + max_cost_for_model=10_000, + model_obj=model, + mint="https://mint", + payment_token_hash="h", + ) + + assert response.status_code == 200 + body = response.body if isinstance(response.body, bytes) else bytes(response.body) + payload = json.loads(body.decode()) + assert "input_tokens" in payload # ---------------------------------------------------------------------------