refactor(passthrough): inject the streaming prompt-token counter instead of a default-image flag

This commit is contained in:
mateo-berri 2026-09-05 02:09:26 -07:00
parent df6fb9e5d9
commit 2f981d14e4
7 changed files with 142 additions and 18 deletions

View file

@ -1,6 +1,6 @@
import base64
import time
from collections.abc import Iterator, Mapping, Sequence
from collections.abc import Callable, Iterator, Mapping, Sequence
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, TypeAlias, TypedDict, Union, cast
@ -994,7 +994,7 @@ class ChunkProcessor:
completion_output: str,
messages: Sequence | None = None,
reasoning_tokens: int | None = None,
use_default_image_token_count: bool = False,
count_prompt_tokens: Callable[[], int] | None = None,
) -> Usage:
"""
Calculate usage for the given chunks.
@ -1019,8 +1019,8 @@ class ChunkProcessor:
cost: Final[float | None] = calculated_usage_per_chunk["cost"]
try:
returned_usage.prompt_tokens = prompt_tokens or token_counter(
model=model, messages=messages, use_default_image_token_count=use_default_image_token_count
returned_usage.prompt_tokens = prompt_tokens or (
count_prompt_tokens() if count_prompt_tokens else token_counter(model=model, messages=messages)
)
except Exception: # don't allow this failing to block a complete streaming response from being returned
print_verbose("token_counter failed, assuming prompt tokens is 0")

View file

@ -19,7 +19,7 @@ import random
import sys
import time
import traceback
from collections.abc import AsyncIterator, Coroutine, Iterable, Mapping, Sequence
from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Mapping, Sequence
from concurrent import futures
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from copy import deepcopy
@ -8674,7 +8674,7 @@ def stream_chunk_builder(
start_time=None,
end_time=None,
logging_obj: Optional["Logging"] = None,
use_default_image_token_count: bool = False,
count_prompt_tokens: Callable[[], int] | None = None,
) -> ModelResponse | TextCompletionResponse | None:
try:
if chunks is None:
@ -8748,7 +8748,7 @@ def stream_chunk_builder(
completion_output=completion_output,
messages=messages,
reasoning_tokens=0,
use_default_image_token_count=use_default_image_token_count,
count_prompt_tokens=count_prompt_tokens,
)
setattr(response, "usage", usage)
@ -8926,7 +8926,7 @@ def stream_chunk_builder(
completion_output=completion_output,
messages=messages,
reasoning_tokens=reasoning_tokens,
use_default_image_token_count=use_default_image_token_count,
count_prompt_tokens=count_prompt_tokens,
)
setattr(response, "usage", usage)

View file

@ -13,6 +13,7 @@ import httpx
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_IMAGE_TOKEN_COUNT
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import (
get_standard_logging_object_payload,
@ -97,6 +98,45 @@ def _is_openai_compatible_url(url_route: str | None) -> bool:
return False
def _is_remote_high_detail_image(part: object) -> bool:
if not isinstance(part, Mapping) or part.get("type") != "image_url":
return False
image_url: Final = part.get("image_url")
if not isinstance(image_url, Mapping):
return False
url: Final = image_url.get("url")
return isinstance(url, str) and url.startswith(("http://", "https://")) and image_url.get("detail") == "high"
def _content_parts(message: Mapping[str, object]) -> Sequence[object]:
content: Final = message.get("content")
return content if isinstance(content, list) else ()
def _without_remote_high_detail_images(message: Mapping[str, object]) -> Mapping[str, object]:
if not isinstance(message.get("content"), list):
return message
kept_parts: Final = [ # mutable-ok: token_counter reads message content only when it is a list
part for part in _content_parts(message) if not _is_remote_high_detail_image(part)
]
return {**message, "content": kept_parts} # mutable-ok: token_counter rejects any message that is not a dict
def count_relayed_prompt_tokens(model: str, messages: Sequence[Mapping[str, object]] | None) -> int:
if messages is None:
return 0
remote_high_detail_images: Final = sum(
1 for message in messages for part in _content_parts(message) if _is_remote_high_detail_image(part)
)
local_messages: Final = [ # mutable-ok: token_counter takes a list of messages
_without_remote_high_detail_images(message) for message in messages
]
return (
litellm.token_counter(model=model, messages=local_messages)
+ DEFAULT_IMAGE_TOKEN_COUNT * remote_high_detail_images
)
class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
"""
OpenAI-specific passthrough logging handler that provides cost tracking for /chat/completions endpoints.
@ -561,7 +601,9 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
# Build complete response from chunks
complete_streaming_response: Final = litellm.stream_chunk_builder(
chunks=all_openai_chunks, messages=messages, use_default_image_token_count=True
chunks=all_openai_chunks,
messages=messages,
count_prompt_tokens=lambda: count_relayed_prompt_tokens(model, messages),
)
return complete_streaming_response

View file

@ -57,7 +57,7 @@
"limit": 3
},
"BLE001": {
"limit": 2910
"limit": 2904
},
"C401": {
"limit": 8
@ -201,7 +201,7 @@
"limit": 310
},
"SIM103": {
"limit": 113
"limit": 110
},
"SIM113": {
"limit": 3

View file

@ -2,8 +2,10 @@ import json
from unittest.mock import MagicMock
import httpx
import pytest
import litellm
from litellm.constants import DEFAULT_IMAGE_TOKEN_COUNT
from litellm.llms.azure.passthrough.transformation import AzurePassthroughConfig
@ -159,7 +161,7 @@ def test_azure_passthrough_streaming_chunks_count_remote_image_prompt_tokens_wit
"role": "user",
"content": [
{"type": "text", "text": "Describe this"},
{"type": "image_url", "image_url": {"url": "http://127.0.0.1:9/doc.png"}},
{"type": "image_url", "image_url": {"url": "http://127.0.0.1:9/doc.png", "detail": "high"}},
],
}
]
@ -174,10 +176,10 @@ def test_azure_passthrough_streaming_chunks_count_remote_image_prompt_tokens_wit
endpoint="openai/deployments/gpt-4.1-mini/chat/completions",
)
text_only_messages = [{"role": "user", "content": [{"type": "text", "text": "Describe this"}]}]
assert isinstance(response, ModelResponse)
assert response.usage.prompt_tokens > 0
assert response.usage.prompt_tokens == litellm.token_counter(
model="gpt-4.1-mini", messages=messages, use_default_image_token_count=True
assert response.usage.prompt_tokens == (
litellm.token_counter(model="gpt-4.1-mini", messages=text_only_messages) + DEFAULT_IMAGE_TOKEN_COUNT
)
@ -218,3 +220,37 @@ def test_azure_passthrough_url_prefers_the_deployments_api_version():
)
assert url.params["api-version"] == "2024-10-21"
def test_azure_passthrough_url_strips_the_leading_router_model_segment():
url, _ = AzurePassthroughConfig().get_complete_url(
api_base="https://my-resource.openai.azure.com",
api_key="key",
model="gpt-4.1-mini",
endpoint="gpt-4.1-mini/openai/deployments/gpt-4.1-mini/chat/completions",
request_query_params={"api-version": "2024-10-21"},
litellm_params={},
)
assert str(url) == "https://my-resource.openai.azure.com/openai/deployments/gpt-4.1-mini/chat/completions?api-version=2024-10-21"
def test_azure_passthrough_url_rewrites_the_model_group_only_as_a_whole_segment():
url, _ = AzurePassthroughConfig().get_complete_url(
api_base="https://my-resource.openai.azure.com",
api_key="key",
model="gpt-4.1-mini",
endpoint="gpt/openai/deployments/gpt-4o/chat/completions",
request_query_params={"api-version": "2024-10-21"},
litellm_params={"litellm_metadata": {"model_group": "gpt"}},
)
assert str(url) == "https://my-resource.openai.azure.com/openai/deployments/gpt-4o/chat/completions?api-version=2024-10-21"
@pytest.mark.parametrize(
"request_data, expected",
[({"stream": True}, True), ({"stream": 1}, True), ({"stream": False}, False), ({}, False)],
)
def test_azure_passthrough_is_streaming_request_reads_the_stream_flag(request_data, expected):
assert AzurePassthroughConfig().is_streaming_request(endpoint="openai/deployments/x/chat/completions", request_data=request_data) is expected

View file

@ -8,9 +8,11 @@ import pytest
import litellm
from litellm.constants import DEFAULT_IMAGE_TOKEN_COUNT
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler import (
OpenAIPassthroughLoggingHandler,
count_relayed_prompt_tokens,
)
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
@ -2037,3 +2039,47 @@ class TestOpenAIPassthroughEmbeddingsSpendLog:
if __name__ == "__main__":
pytest.main([__file__])
ONE_PIXEL_PNG_DATA_URL = (
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNkYAAAAAYAAjCB0C8AAAAASUVORK5CYII="
)
UNREACHABLE_IMAGE_URL = "http://127.0.0.1:9/doc.png"
TEXT_ONLY_MESSAGES = [{"role": "user", "content": [{"type": "text", "text": "Describe this"}]}]
def _image_messages(url: str, detail: str) -> list[dict]:
return [
{
"role": "user",
"content": [
{"type": "text", "text": "Describe this"},
{"type": "image_url", "image_url": {"url": url, "detail": detail}},
],
}
]
def test_count_relayed_prompt_tokens_counts_a_data_url_image_exactly():
messages = _image_messages(ONE_PIXEL_PNG_DATA_URL, "high")
assert count_relayed_prompt_tokens("gpt-4.1-mini", messages) == litellm.token_counter(
model="gpt-4.1-mini", messages=messages
)
def test_count_relayed_prompt_tokens_keeps_a_low_detail_remote_image_at_the_base_count():
messages = _image_messages(UNREACHABLE_IMAGE_URL, "low")
assert count_relayed_prompt_tokens("gpt-4.1-mini", messages) == litellm.token_counter(
model="gpt-4.1-mini", messages=messages
)
assert count_relayed_prompt_tokens("gpt-4.1-mini", messages) < DEFAULT_IMAGE_TOKEN_COUNT
def test_count_relayed_prompt_tokens_estimates_only_the_remote_high_detail_image():
messages = _image_messages(UNREACHABLE_IMAGE_URL, "high")
assert count_relayed_prompt_tokens("gpt-4.1-mini", messages) == (
litellm.token_counter(model="gpt-4.1-mini", messages=TEXT_ONLY_MESSAGES) + DEFAULT_IMAGE_TOKEN_COUNT
)

View file

@ -1,6 +1,6 @@
{
"LIT001": {
"limit": 22166
"limit": 22156
},
"LIT002": {
"limit": 26745
@ -27,10 +27,10 @@
"limit": 0
},
"LIT010": {
"limit": 16446
"limit": 16434
},
"LIT011": {
"limit": 5506
"limit": 5504
},
"LIT012": {
"limit": 4486