From f58983b9eb0b2ad1e624be0064b6b4a63ef4e1a2 Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Tue, 8 Sep 2026 20:55:58 +0200 Subject: [PATCH] fix: account for Responses input and terminal usage snapshots --- routstr/upstream/base.py | 4 +- routstr/upstream/count_tokens.py | 39 +++++++ tests/unit/test_count_tokens_local.py | 36 +++++++ tests/unit/test_x_cashu_missing_usage.py | 132 ++++++++++++++++++++++- 4 files changed, 207 insertions(+), 4 deletions(-) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 2ea06e0b..d41b2664 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -3892,7 +3892,7 @@ class BaseUpstreamProvider: continue usage_estimator.observe(data_json) - if usage_data is None: + if not usage_data: usage_data = _estimated_usage(usage_estimator, model) if usage_data: logger.warning( @@ -4935,7 +4935,7 @@ class BaseUpstreamProvider: elif not model and payload.get("model"): model = payload["model"] - if usage_data is None: + if not usage_data: usage_data = _estimated_usage(usage_estimator, model) logger.warning( "No usage in streaming Responses API response, billing from local token estimate", diff --git a/routstr/upstream/count_tokens.py b/routstr/upstream/count_tokens.py index 8ebeb6d9..51ffbfa4 100644 --- a/routstr/upstream/count_tokens.py +++ b/routstr/upstream/count_tokens.py @@ -57,6 +57,37 @@ def _count_with_litellm( if not isinstance(messages, list): messages = [] + if "input" in body: + response_input = body["input"] + if isinstance(response_input, str): + messages = [{"role": "user", "content": response_input}] + elif isinstance(response_input, list): + messages = [] + for item in response_input: + if not isinstance(item, dict) or "role" not in item: + raise ValueError( + "Responses input requires fallback token estimation" + ) + content = item.get("content", "") + if isinstance(content, list): + parts = [] + for part in content: + if not isinstance(part, dict) or part.get("type") not in ( + "input_text", + "output_text", + "text", + ): + raise ValueError( + "Non-text Responses input requires fallback token estimation" + ) + parts.append({"type": "text", "text": part.get("text", "")}) + content = parts + messages.append({"role": item["role"], "content": content}) + else: + raise ValueError("Unsupported Responses input") + if body.get("instructions"): + messages.insert(0, {"role": "system", "content": body["instructions"]}) + prompt_token_ids = 0 if include_legacy_prompt: prompt = body.get("prompt") @@ -177,6 +208,14 @@ class MissingUsageEstimator: def observe(self, response_data: object) -> None: if isinstance(response_data, dict): event_type = response_data.get("type") + if event_type in ("response.completed", "response.incomplete"): + response = response_data.get("response") + if isinstance(response, dict) and isinstance( + response.get("output"), list + ): + # Terminal output is a snapshot, not another text delta. + self._output_parts = _generated_text(response["output"]) + return if isinstance(event_type, str) and event_type.endswith(".done"): # Responses API ``*.done`` events repeat text already streamed # via ``*.delta`` events; counting both would double-bill. diff --git a/tests/unit/test_count_tokens_local.py b/tests/unit/test_count_tokens_local.py index 85a8ec72..23bf379c 100644 --- a/tests/unit/test_count_tokens_local.py +++ b/tests/unit/test_count_tokens_local.py @@ -272,3 +272,39 @@ def test_uses_forwarded_model_id_when_present() -> None: assert captured["model"] == "claude-3-5-sonnet-20241022" assert _read_payload(response)["input_tokens"] == 7 + + +def test_responses_instructions_are_counted_as_system_text() -> None: + body = {"model": "gpt-4o", "input": "Hi", "instructions": "Be concise."} + with patch.object( + count_tokens_module.litellm, "token_counter", return_value=12 + ) as counter: + usage = MissingUsageEstimator(_body(body), None).response_data()["usage"] + + assert usage["input_tokens"] == 12 + counter.assert_called_once_with( + model="gpt-4o", + messages=[ + {"role": "system", "content": "Be concise."}, + {"role": "user", "content": "Hi"}, + ], + tools=None, + ) + + +def test_responses_tool_results_use_fallback_instead_of_empty_messages() -> None: + body = { + "model": "gpt-4o", + "input": [ + { + "type": "function_call_output", + "call_id": "call_1", + "output": "result " * 100, + } + ], + } + with patch.object(count_tokens_module.litellm, "token_counter") as counter: + usage = MissingUsageEstimator(_body(body), None).response_data()["usage"] + + counter.assert_not_called() + assert usage["input_tokens"] > 100 diff --git a/tests/unit/test_x_cashu_missing_usage.py b/tests/unit/test_x_cashu_missing_usage.py index 19f5ae34..ceaabade 100644 --- a/tests/unit/test_x_cashu_missing_usage.py +++ b/tests/unit/test_x_cashu_missing_usage.py @@ -40,6 +40,7 @@ async def _settle( *, responses_api: bool = False, request_body: bytes | None = REQUEST_BODY, + unit: str = "msat", ) -> tuple[Any, AsyncMock, AsyncMock]: provider = BaseUpstreamProvider(base_url="http://test", api_key="test-key") get_cost = AsyncMock(side_effect=provider.get_x_cashu_cost) @@ -56,7 +57,7 @@ async def _settle( result = await handler( response=response, amount=10_000, - unit="msat", + unit=unit, max_cost_for_model=9_000, mint=None, request_body=request_body, @@ -113,7 +114,7 @@ async def test_non_streaming_chat_without_usage_bills_from_estimate() -> None: @pytest.mark.asyncio async def test_streaming_responses_without_usage_bills_from_estimate() -> None: - events = [ + events: list[dict[str, Any]] = [ {"type": "response.created", "response": {"model": "gpt-5-mini"}}, {"type": "response.output_text.delta", "delta": "Why did the chicken"}, {"type": "response.output_text.done", "text": "Why did the chicken"}, @@ -139,3 +140,130 @@ async def test_non_streaming_responses_without_usage_bills_from_estimate() -> No usage = _billed_usage(get_cost) assert usage is not None assert usage["output_tokens"] > 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("responses_api", [False, True]) +async def test_empty_streaming_usage_uses_estimate(responses_api: bool) -> None: + payload: dict[str, Any] = {"model": "gpt-4o", "usage": {}} + if responses_api: + payload["output"] = [{"content": [{"type": "output_text", "text": "Hello"}]}] + else: + payload["choices"] = [{"delta": {"content": "Hello"}}] + + _, get_cost, _ = await _settle(_sse([payload]), responses_api=responses_api) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage.get("estimated") is True + assert usage["output_tokens"] > 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("terminal_type", ["response.completed", "response.incomplete"]) +async def test_responses_terminal_output_is_not_billed_twice( + terminal_type: str, +) -> None: + payload = { + "model": "gpt-4o", + "output": [{"content": [{"type": "output_text", "text": "Hello world"}]}], + } + events: list[dict[str, Any]] = [ + {"type": "response.output_text.delta", "delta": "Hello world"}, + {"type": terminal_type, "response": payload}, + ] + _, streaming_cost, _ = await _settle(_sse(events), responses_api=True) + _, json_cost, _ = await _settle(_json(payload), responses_api=True) + + stream_usage = _billed_usage(streaming_cost) + json_usage = _billed_usage(json_cost) + assert stream_usage is not None and json_usage is not None + for field in ("input_tokens", "output_tokens", "total_tokens"): + assert stream_usage[field] == json_usage[field] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("input_shape", ["string", "message", "content_blocks"]) +async def test_responses_estimate_includes_input( + stream: bool, input_shape: str +) -> None: + prompt = "Explain how payment reservations work. " * 100 + response_input: Any = prompt + if input_shape == "message": + response_input = [{"role": "user", "content": prompt}] + elif input_shape == "content_blocks": + response_input = [ + {"role": "user", "content": [{"type": "input_text", "text": prompt}]} + ] + request_body = json.dumps( + { + "model": "gpt-4o", + "instructions": "Answer concisely.", + "input": response_input, + } + ).encode() + payload = { + "model": "gpt-4o", + "output": [{"content": [{"type": "output_text", "text": "Hello"}]}], + } + response = _sse([payload]) if stream else _json(payload) + _, get_cost, _ = await _settle( + response, responses_api=True, request_body=request_body + ) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage["input_tokens"] > 100 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("responses_api", [False, True]) +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("tokens", [0, 17]) +async def test_reported_usage_takes_precedence( + responses_api: bool, stream: bool, tokens: int +) -> None: + payload: dict[str, Any] = { + "model": "gpt-4o", + "usage": {"input_tokens": tokens, "output_tokens": tokens}, + } + if responses_api: + payload["output"] = [{"content": [{"type": "output_text", "text": "Hello"}]}] + else: + payload["choices"] = [{"message": {"content": "Hello"}}] + + response = _sse([payload]) if stream else _json(payload) + _, get_cost, _ = await _settle(response, responses_api=responses_api) + + usage = _billed_usage(get_cost) + assert usage is not None + assert usage["input_tokens"] == tokens + assert usage["output_tokens"] == tokens + assert "estimated" not in usage + + +@pytest.mark.asyncio +@pytest.mark.parametrize("responses_api", [False, True]) +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("unit", ["sat", "msat"]) +async def test_pricing_error_refunds_full_prepayment( + responses_api: bool, stream: bool, unit: str +) -> None: + payload = { + "model": "unpriced-model", + "usage": {"input_tokens": 10, "output_tokens": 5}, + } + response = _sse([payload]) if stream else _json(payload) + with patch( + "routstr.payment.cost_calculation._get_pricing_rates", + side_effect=ValueError("No pricing for model"), + ): + result, _, send_refund = await _settle( + response, responses_api=responses_api, unit=unit + ) + + send_refund.assert_awaited_once_with(10_000, unit, None, request_id=None) + assert result.status_code == 200 + assert result.headers["X-Cashu"] == "cashuBrefund" + assert result.headers["X-Routstr-Cost-Msats"] == "0"