Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
*.pid

# Local Codex inspection artifacts (don't ship)
codedb.snapshot
*.asar
app-asar-work/
*-bsLDOISN.js
Expand Down
24 changes: 22 additions & 2 deletions codex_shim/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
anthropic_to_response,
chat_completion_to_response,
chat_to_anthropic,
normalize_responses_usage,
responses_to_anthropic,
responses_to_chat,
)
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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


Expand Down
73 changes: 71 additions & 2 deletions codex_shim/translate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
44 changes: 36 additions & 8 deletions tests/test_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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},
}
)

Expand Down Expand Up @@ -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"

Expand All @@ -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
Expand Down Expand Up @@ -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()
Expand All @@ -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):
Expand All @@ -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},
}
)

Expand Down Expand Up @@ -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"]
Expand Down Expand Up @@ -723,4 +752,3 @@ async def test_switch_model_requires_slug(tmp_path, auth_missing):
assert resp.status == 400
finally:
await shim_client.close()

52 changes: 51 additions & 1 deletion tests/test_translate.py
Original file line number Diff line number Diff line change
@@ -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():
Expand Down Expand Up @@ -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,
},
}
Loading