mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
support /message/count_tokens endpoint
This commit is contained in:
@@ -36,6 +36,7 @@ class CostDataError(BaseModel):
|
|||||||
async def calculate_cost(
|
async def calculate_cost(
|
||||||
response_data: dict, max_cost: int, session: AsyncSession
|
response_data: dict, max_cost: int, session: AsyncSession
|
||||||
) -> CostData | MaxCostData | CostDataError:
|
) -> CostData | MaxCostData | CostDataError:
|
||||||
|
print(response_data)
|
||||||
"""Calculate the cost of an API request based on token usage.
|
"""Calculate the cost of an API request based on token usage.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ from ..payment.models import (
|
|||||||
from ..payment.price import sats_usd_price
|
from ..payment.price import sats_usd_price
|
||||||
from ..wallet import recieve_token, send_token
|
from ..wallet import recieve_token, send_token
|
||||||
from . import messages_dispatch
|
from . import messages_dispatch
|
||||||
|
from .count_tokens import count_tokens_locally
|
||||||
from .litellm_routing import detect_litellm_prefix
|
from .litellm_routing import detect_litellm_prefix
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
logger = get_logger(__name__)
|
||||||
@@ -2159,6 +2160,12 @@ class BaseUpstreamProvider:
|
|||||||
"""
|
"""
|
||||||
path = self.normalize_request_path(path, model_obj)
|
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 (
|
if (
|
||||||
path.endswith("messages")
|
path.endswith("messages")
|
||||||
and not path.endswith("count_tokens")
|
and not path.endswith("count_tokens")
|
||||||
@@ -3357,6 +3364,12 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
request_body = await request.body()
|
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 (
|
if (
|
||||||
path.endswith("messages")
|
path.endswith("messages")
|
||||||
and not path.endswith("count_tokens")
|
and not path.endswith("count_tokens")
|
||||||
|
|||||||
@@ -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",
|
||||||
|
)
|
||||||
@@ -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
|
||||||
@@ -908,7 +908,7 @@ async def test_forward_request_skips_litellm_when_provider_supports_messages() -
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@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()
|
provider = _make_provider()
|
||||||
key = _make_key()
|
key = _make_key()
|
||||||
model = _make_model()
|
model = _make_model()
|
||||||
@@ -923,19 +923,25 @@ async def test_forward_request_skips_litellm_for_count_tokens() -> None:
|
|||||||
with patch.object(
|
with patch.object(
|
||||||
provider,
|
provider,
|
||||||
"prepare_request_body",
|
"prepare_request_body",
|
||||||
side_effect=RuntimeError("stop here"),
|
side_effect=AssertionError("upstream should not be called"),
|
||||||
):
|
):
|
||||||
with pytest.raises(RuntimeError, match="stop here"):
|
response = await provider.forward_request(
|
||||||
await provider.forward_request(
|
request=request,
|
||||||
request=request,
|
path="messages/count_tokens",
|
||||||
path="messages/count_tokens",
|
headers={},
|
||||||
headers={},
|
request_body=_anthropic_request_body(),
|
||||||
request_body=_anthropic_request_body(),
|
key=key,
|
||||||
key=key,
|
max_cost_for_model=10_000,
|
||||||
max_cost_for_model=10_000,
|
session=session,
|
||||||
session=session,
|
model_obj=model,
|
||||||
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
|
@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()
|
provider = _make_provider()
|
||||||
model = _make_model()
|
model = _make_model()
|
||||||
request = _make_request()
|
request = _make_request()
|
||||||
@@ -1024,18 +1030,24 @@ async def test_forward_x_cashu_request_skips_litellm_for_count_tokens() -> None:
|
|||||||
with patch.object(
|
with patch.object(
|
||||||
provider,
|
provider,
|
||||||
"prepare_request_body",
|
"prepare_request_body",
|
||||||
side_effect=RuntimeError("stop here"),
|
side_effect=AssertionError("upstream should not be called"),
|
||||||
):
|
):
|
||||||
with pytest.raises(RuntimeError, match="stop here"):
|
response = await provider.forward_x_cashu_request(
|
||||||
await provider.forward_x_cashu_request(
|
request=request,
|
||||||
request=request,
|
path="v1/messages/count_tokens",
|
||||||
path="v1/messages/count_tokens",
|
headers={},
|
||||||
headers={},
|
amount=5_000,
|
||||||
amount=5_000,
|
unit="sat",
|
||||||
unit="sat",
|
max_cost_for_model=10_000,
|
||||||
max_cost_for_model=10_000,
|
model_obj=model,
|
||||||
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
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user