From 094be588e50c165729156d557f2a68963a14bc21 Mon Sep 17 00:00:00 2001 From: Tharusha Nirmal Date: Thu, 28 May 2026 13:29:07 +0530 Subject: [PATCH] fix: normalize byok responses usage --- .gitignore | 1 + codex_shim/server.py | 24 ++++++++++++-- codex_shim/translate.py | 73 +++++++++++++++++++++++++++++++++++++++-- tests/test_server.py | 44 ++++++++++++++++++++----- tests/test_translate.py | 52 ++++++++++++++++++++++++++++- 5 files changed, 181 insertions(+), 13 deletions(-) diff --git a/.gitignore b/.gitignore index 7f3c6e14..710873e3 100644 --- a/.gitignore +++ b/.gitignore @@ -4,6 +4,7 @@ *.pid # Local Codex inspection artifacts (don't ship) +codedb.snapshot *.asar app-asar-work/ *-bsLDOISN.js diff --git a/codex_shim/server.py b/codex_shim/server.py index e5113ed8..ff12957a 100644 --- a/codex_shim/server.py +++ b/codex_shim/server.py @@ -29,6 +29,7 @@ anthropic_to_response, chat_completion_to_response, chat_to_anthropic, + normalize_responses_usage, responses_to_anthropic, responses_to_chat, ) @@ -640,7 +641,7 @@ async def finish(self, response: web.StreamResponse) -> None: async def write_chat_delta(self, response: web.StreamResponse, chunk: dict[str, Any]) -> None: usage = chunk.get("usage") if isinstance(usage, dict): - self.usage = usage + self.usage = normalize_responses_usage(usage) choice = (chunk.get("choices") or [{}])[0] delta = choice.get("delta") or {} reasoning = delta.get("reasoning_content") or delta.get("reasoning") @@ -699,6 +700,11 @@ async def _chat_tool_delta(self, response: web.StreamResponse, call: dict[str, A # ------------------------------------------------------------------ async def write_anthropic_delta(self, response: web.StreamResponse, event: dict[str, Any]) -> None: event_type = event.get("type") + if event_type == "message_start": + message = event.get("message") or {} + usage = message.get("usage") + if isinstance(usage, dict): + self.usage = normalize_responses_usage(usage) if event_type == "content_block_start": block = event.get("content_block") or {} idx = int(event.get("index", 0)) @@ -769,7 +775,19 @@ async def write_anthropic_delta(self, response: web.StreamResponse, event: dict[ elif event_type == "message_delta": usage = event.get("usage") if isinstance(usage, dict): - self.usage = usage + if self.usage is None or any( + key in usage for key in ("input_tokens", "prompt_tokens", "cache_read_input_tokens", "cache_creation_input_tokens") + ): + normalized = normalize_responses_usage(usage) + if normalized is not None: + self.usage = normalized if self.usage is None else {**self.usage, **normalized} + output_tokens = usage.get("output_tokens") + if isinstance(output_tokens, int) and not isinstance(output_tokens, bool): + if self.usage is None: + self.usage = normalize_responses_usage(usage) + else: + self.usage["output_tokens"] = output_tokens + self.usage["total_tokens"] = int(self.usage.get("input_tokens") or 0) + output_tokens elif event_type == "content_block_stop": idx = int(event.get("index", 0)) tool_state = self.tool_calls.get(("anthropic", idx)) @@ -1058,6 +1076,8 @@ def _response(self, status: str, *, final: bool = False) -> dict[str, Any]: } if self.usage is not None: payload["usage"] = self.usage + elif final: + payload["usage"] = {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0} return payload diff --git a/codex_shim/translate.py b/codex_shim/translate.py index 65a7daa0..cda169bd 100644 --- a/codex_shim/translate.py +++ b/codex_shim/translate.py @@ -263,12 +263,81 @@ def chat_completion_to_response(payload: dict[str, Any], requested_model: str) - "status": "completed", "model": requested_model, "output": output, - "usage": payload.get("usage"), + "usage": normalize_responses_usage(payload.get("usage")), } def anthropic_to_response(payload: dict[str, Any], requested_model: str) -> dict[str, Any]: - return chat_completion_to_response(anthropic_to_chat_response(payload, requested_model), requested_model) + response = chat_completion_to_response(anthropic_to_chat_response(payload, requested_model), requested_model) + response["usage"] = normalize_responses_usage(payload.get("usage")) + return response + + +def normalize_responses_usage(usage: Any) -> dict[str, Any] | None: + if not isinstance(usage, dict): + return None + + input_tokens = _int_token(usage.get("input_tokens")) + if input_tokens is None: + input_tokens = _int_token(usage.get("prompt_tokens")) + + output_tokens = _int_token(usage.get("output_tokens")) + if output_tokens is None: + output_tokens = _int_token(usage.get("completion_tokens")) + + total_tokens = _int_token(usage.get("total_tokens")) + if total_tokens is None and input_tokens is not None and output_tokens is not None: + total_tokens = input_tokens + output_tokens + + if input_tokens is None: + input_tokens = max(total_tokens - output_tokens, 0) if total_tokens is not None and output_tokens is not None else 0 + if output_tokens is None: + output_tokens = max(total_tokens - input_tokens, 0) if total_tokens is not None else 0 + if total_tokens is None: + total_tokens = input_tokens + output_tokens + + normalized: dict[str, Any] = { + "input_tokens": input_tokens, + "output_tokens": output_tokens, + "total_tokens": total_tokens, + } + + input_details: dict[str, Any] = {} + if isinstance(usage.get("input_tokens_details"), dict): + input_details.update(usage["input_tokens_details"]) + if isinstance(usage.get("prompt_tokens_details"), dict): + input_details.update(usage["prompt_tokens_details"]) + + cache_read = _int_token(usage.get("cache_read_input_tokens")) + if cache_read is not None: + input_details.setdefault("cached_tokens", cache_read) + input_details.setdefault("cache_read_input_tokens", cache_read) + cache_created = _int_token(usage.get("cache_creation_input_tokens")) + if cache_created is not None: + input_details.setdefault("cache_creation_input_tokens", cache_created) + + if input_details: + normalized["input_tokens_details"] = input_details + + output_details: dict[str, Any] = {} + if isinstance(usage.get("output_tokens_details"), dict): + output_details.update(usage["output_tokens_details"]) + if isinstance(usage.get("completion_tokens_details"), dict): + output_details.update(usage["completion_tokens_details"]) + if output_details: + normalized["output_tokens_details"] = output_details + + return normalized + + +def _int_token(value: Any) -> int | None: + if isinstance(value, bool): + return None + if isinstance(value, int): + return value + if isinstance(value, float) and value.is_integer(): + return int(value) + return None def strip_think(text: str) -> str: diff --git a/tests/test_server.py b/tests/test_server.py index 798233da..9d743765 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -188,7 +188,7 @@ async def chat(request): { "id": "chatcmpl_fake", "choices": [{"message": {"role": "assistant", "content": "hello"}}], - "usage": {"total_tokens": 3}, + "usage": {"prompt_tokens": 2, "completion_tokens": 1, "total_tokens": 3}, } ) @@ -220,6 +220,7 @@ async def chat(request): assert resp.status == 200 payload = await resp.json() assert payload["output"][0]["content"][0]["text"] == "hello" + assert payload["usage"] == {"input_tokens": 2, "output_tokens": 1, "total_tokens": 3} assert captured["body"]["model"] == "real-openai" assert captured["headers"]["Authorization"] == "Bearer secret" @@ -244,7 +245,9 @@ async def chat(request): response = web.StreamResponse(headers={"Content-Type": "text/event-stream"}) await response.prepare(request) await response.write(b'data: {"choices":[{"delta":{"content":"hello"}}]}\n\n') - await response.write(b'data: {"choices":[{"delta":{}}],"usage":{"prompt_tokens":4,"completion_tokens":2,"total_tokens":6}}\n\n') + await response.write( + b'data: {"choices":[{"delta":{}}],"usage":{"prompt_tokens":4,"completion_tokens":2,"total_tokens":6,"prompt_tokens_details":{"cached_tokens":3}}}\n\n' + ) await response.write(b"data: [DONE]\n\n") await response.write_eof() return response @@ -276,7 +279,12 @@ async def chat(request): assert resp.status == 200 events = _sse_events(await resp.text()) completed = [event for event in events if event.get("type") == "response.completed"][-1] - assert completed["response"]["usage"] == {"prompt_tokens": 4, "completion_tokens": 2, "total_tokens": 6} + assert completed["response"]["usage"] == { + "input_tokens": 4, + "output_tokens": 2, + "total_tokens": 6, + "input_tokens_details": {"cached_tokens": 3}, + } await shim_client.close() await upstream_client.close() @@ -292,19 +300,40 @@ async def write(self, data: bytes): downstream = FakeResponse() state = ResponsesStreamState("claude-real") + await state.write_anthropic_delta( + downstream, + { + "type": "message_start", + "message": { + "usage": { + "input_tokens": 5, + "cache_read_input_tokens": 4, + "output_tokens": 1, + } + }, + }, + ) await state.write_anthropic_delta( downstream, { "type": "message_delta", "delta": {"stop_reason": "end_turn"}, - "usage": {"input_tokens": 5, "output_tokens": 3}, + "usage": {"output_tokens": 3}, }, ) await state.finish(downstream) events = _sse_events(b"".join(downstream.chunks).decode()) completed = [event for event in events if event.get("type") == "response.completed"][-1] - assert completed["response"]["usage"] == {"input_tokens": 5, "output_tokens": 3} + assert completed["response"]["usage"] == { + "input_tokens": 5, + "output_tokens": 3, + "total_tokens": 8, + "input_tokens_details": { + "cached_tokens": 4, + "cache_read_input_tokens": 4, + }, + } async def test_responses_compact_routes_to_openai_chat_and_returns_compacted_window(tmp_path): @@ -316,7 +345,7 @@ async def chat(request): { "id": "chatcmpl_compact", "choices": [{"message": {"role": "assistant", "content": "Task: keep implementing compact support."}}], - "usage": {"total_tokens": 11}, + "usage": {"prompt_tokens": 9, "completion_tokens": 2, "total_tokens": 11}, } ) @@ -361,7 +390,7 @@ async def chat(request): assert payload["status"] == "completed" assert payload["model"] == "real-openai" assert payload["output"][0]["content"][0]["text"] == "Task: keep implementing compact support." - assert payload["usage"] == {"total_tokens": 11} + assert payload["usage"] == {"input_tokens": 9, "output_tokens": 2, "total_tokens": 11} assert captured["body"]["model"] == "real-openai" assert captured["body"]["stream"] is False assert "service_tier" not in captured["body"] @@ -723,4 +752,3 @@ async def test_switch_model_requires_slug(tmp_path, auth_missing): assert resp.status == 400 finally: await shim_client.close() - diff --git a/tests/test_translate.py b/tests/test_translate.py index 6e9aae92..8d4fd0a8 100644 --- a/tests/test_translate.py +++ b/tests/test_translate.py @@ -1,6 +1,6 @@ from __future__ import annotations -from codex_shim.translate import chat_completion_to_response, responses_to_anthropic, responses_to_chat +from codex_shim.translate import anthropic_to_response, chat_completion_to_response, responses_to_anthropic, responses_to_chat def test_responses_to_chat_text_input(): @@ -228,3 +228,53 @@ def test_chat_completion_to_response_strips_think(): out = chat_completion_to_response(payload, "slug") assert out["model"] == "slug" assert out["output"][0]["content"][0]["text"] == "Hello" + + +def test_chat_completion_to_response_normalizes_cached_usage(): + payload = { + "id": "chatcmpl_1", + "choices": [{"message": {"role": "assistant", "content": "Hello"}}], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 2, + "total_tokens": 12, + "prompt_tokens_details": {"cached_tokens": 8}, + "completion_tokens_details": {"reasoning_tokens": 1}, + }, + } + + out = chat_completion_to_response(payload, "slug") + + assert out["usage"] == { + "input_tokens": 10, + "output_tokens": 2, + "total_tokens": 12, + "input_tokens_details": {"cached_tokens": 8}, + "output_tokens_details": {"reasoning_tokens": 1}, + } + + +def test_anthropic_to_response_normalizes_cache_usage(): + payload = { + "id": "msg_1", + "content": [{"type": "text", "text": "Hello"}], + "usage": { + "input_tokens": 10, + "cache_read_input_tokens": 8, + "cache_creation_input_tokens": 2, + "output_tokens": 3, + }, + } + + out = anthropic_to_response(payload, "slug") + + assert out["usage"] == { + "input_tokens": 10, + "output_tokens": 3, + "total_tokens": 13, + "input_tokens_details": { + "cached_tokens": 8, + "cache_read_input_tokens": 8, + "cache_creation_input_tokens": 2, + }, + }