fix(count_tokens): keep assistant turns on the provider counting API

Assistant list content was forwarded to /v1/responses/input_tokens as chat
`text` blocks, which the Responses API rejects (it accepts only output_text
and refusal inside an assistant turn). The 400 sent the whole request to the
local tokenizer, so any conversation with an assistant turn silently lost
provider-exact counting, including the image counting added in 73ab647b1c.

Assistant content now collapses to the plain string the Responses API counts
identically, and image parts are kept to user turns where they are legal.
This commit is contained in:
mateo-berri 2026-08-31 12:59:40 -07:00
parent 73ab647b1c
commit fe90c6f6fc
2 changed files with 77 additions and 5 deletions

View file

@ -23,6 +23,8 @@ class ResponsesInputImagePart(TypedDict):
ResponsesInputPart = ResponsesInputTextPart | ResponsesInputImagePart
ResponsesContentRole = Literal["user", "assistant"]
def _chat_image_block_to_responses_part(image_url: object) -> ResponsesInputImagePart | None:
url: Final = image_url.get("url") if isinstance(image_url, Mapping) else image_url
@ -37,7 +39,7 @@ def _chat_image_block_to_responses_part(image_url: object) -> ResponsesInputImag
return part
def _chat_block_to_responses_part(block: object) -> ResponsesInputPart | None:
def _chat_block_to_responses_part(block: object, role: ResponsesContentRole) -> ResponsesInputPart | None:
if isinstance(block, str):
bare: Final[ResponsesInputTextPart] = {"type": "input_text", "text": block}
return bare
@ -51,7 +53,7 @@ def _chat_block_to_responses_part(block: object) -> ResponsesInputPart | None:
"text": text_value if isinstance(text_value, str) else "",
}
return text
case "image_url":
case "image_url" if role == "user":
return _chat_image_block_to_responses_part(block.get("image_url"))
case _:
return None
@ -59,10 +61,15 @@ def _chat_block_to_responses_part(block: object) -> ResponsesInputPart | None:
def chat_content_blocks_to_responses_content(
content: Sequence[object],
role: ResponsesContentRole,
) -> str | tuple[ResponsesInputPart, ...]:
"""Text-only content collapses to a joined string, so text-only counts stay unchanged."""
"""Text-only content collapses to a joined string, which every role accepts and counts identically.
Only a user turn may carry an image part: the Responses API rejects any part but
output_text and refusal inside an assistant turn.
"""
parts: Final = tuple(
part for part in (_chat_block_to_responses_part(block) for block in content) if part is not None
part for part in (_chat_block_to_responses_part(block, role) for block in content) if part is not None
)
if any(part["type"] != "input_text" for part in parts):
return parts
@ -182,11 +189,13 @@ class OpenAICountTokensConfig:
instructions_parts.append("\n".join(text_parts))
elif role == "user":
if isinstance(content, list):
content = chat_content_blocks_to_responses_content(content)
content = chat_content_blocks_to_responses_content(content, "user")
input_items.append({"role": "user", "content": content})
elif role == "assistant":
# Map tool_calls to Responses API function_call items
tool_calls = msg.get("tool_calls")
if isinstance(content, list):
content = chat_content_blocks_to_responses_content(content, "assistant")
if content:
input_items.append({"role": "assistant", "content": content})
if tool_calls:

View file

@ -260,6 +260,69 @@ def test_messages_to_responses_input_drops_unmappable_blocks():
)
def test_messages_to_responses_input_assistant_blocks_collapse_to_a_string():
"""An assistant turn must never forward chat `text` blocks.
The Responses API only accepts output_text and refusal inside an assistant turn, so
forwarding them 400s the whole request and silently drops the count back to the local
tokenizer, which is exactly what defeats the image fix above.
"""
messages = [
{"role": "user", "content": [{"type": "text", "text": "What is the capital of France?"}]},
{"role": "assistant", "content": [{"type": "text", "text": "Paris."}]},
]
input_items, _ = OpenAICountTokensConfig.messages_to_responses_input(messages)
assert input_items == [
{"role": "user", "content": "What is the capital of France?"},
{"role": "assistant", "content": "Paris."},
]
def test_messages_to_responses_input_assistant_image_block_is_dropped():
"""An image part is illegal inside an assistant turn, so it must not reach the provider."""
messages = [
{
"role": "assistant",
"content": [
{"type": "text", "text": "Here it is"},
{"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}},
],
}
]
input_items, _ = OpenAICountTokensConfig.messages_to_responses_input(messages)
assert input_items == [{"role": "assistant", "content": "Here it is"}]
def test_messages_to_responses_input_keeps_user_image_alongside_an_assistant_turn():
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "What is in this image?"},
{"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}},
],
},
{"role": "assistant", "content": [{"type": "text", "text": "A cat."}]},
]
input_items, _ = OpenAICountTokensConfig.messages_to_responses_input(messages)
assert input_items == [
{
"role": "user",
"content": (
{"type": "input_text", "text": "What is in this image?"},
{"type": "input_image", "image_url": "https://example.com/cat.png", "detail": "auto"},
),
},
{"role": "assistant", "content": "A cat."},
]
def test_validate_request_valid():
"""Test that valid requests pass validation."""
config = OpenAICountTokensConfig()