fix(streaming): keep the served service_tier on streamed chunks and spend rows (#42870)

* fix(streaming): keep the provider's served service_tier on streamed chunks and spend rows

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(streaming): satisfy type-discipline and strict ruff budgets

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(streaming): stamp the served service_tier on every Responses bridge chunk

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(anthropic-adapter): expose streamed chunks so disconnects bill partial spend

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(service-tier): cover anthropic and responses served-tier billing paths

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(anthropic-adapter): return a chunks-exposing stream so disconnects bill partial spend

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(service-tier): bill disconnects through the router's anthropic stream wrapper

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* style: apply ruff format to the anthropic stream changes

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(coverage): ignore delegating properties the ast scan cannot see

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* style: keep the cast-ok reasons on the cast call line

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): cover served service_tier billing for streamed chat and messages, complete and disconnected

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(anthropic-cache): delegate chunks/messages/model through the messages stream cache writer

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(streaming): keep service_tier on OpenAI-compatible parsed chunks

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(streaming): parameterize delegated chunks and messages types

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(tests): follow the anthropic pass_through rename after merging main

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(anthropic): drain the logging worker between response cache tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(spend): cover azure, databricks, responses bridge and gemini served tiers in the stream billing integration test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(databricks): keep the served service_tier on streamed chunks and bill it

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(databricks): type the served service_tier chunk without a loose kwargs dict

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(cost): bill the served service_tier over the requested one

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(cost): drop explanatory comment from the tier resolution

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: kerry <kerry@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-09-29 12:54:17 -07:00 • committed by GitHub
parent 273489824a
commit e814532033
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
32 changed files with 1621 additions and 53 deletions

View file

@ -1334,6 +1334,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
):
super().__init__(streaming_response, sync_stream, json_mode)
self._chat_completion_id: str | None = None
self._served_service_tier: str | None = None
self._tool_call_index_map: dict[int, int] = {} # mutable-ok: per-stream accumulator state
def _handle_string_chunk(
@ -1598,6 +1599,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response_data.get("usage"))
provider_metadata: Final = _provider_metadata(response_data)
served_service_tier: Final = response_data.get("service_tier")
return ModelResponseStream(
choices=[
StreamingChoices(
@ -1611,6 +1613,11 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
],
usage=usage,
provider_specific_fields=dict(provider_metadata) or None, # mutable-ok: field is typed dict
**(
MappingProxyType({"service_tier": served_service_tier})
if isinstance(served_service_tier, str)
else MappingProxyType({})
),
)
else:
pass
@ -1639,12 +1646,28 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
ModelResponseStream: OpenAI-formatted streaming chunk
"""
verbose_logger.debug("Chat provider: transform_streaming_response called with chunk: %s", chunk)
return self._with_stream_scoped_id(
OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
chunk, tool_call_index_map=self._tool_call_index_map
self._remember_served_service_tier(chunk)
return self._with_served_service_tier(
self._with_stream_scoped_id(
OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(
chunk, tool_call_index_map=self._tool_call_index_map
)
)
)
def _remember_served_service_tier(self, chunk: dict[str, object]) -> None:
response_payload: Final = chunk.get("response")
if not isinstance(response_payload, dict):
return
served_tier: Final = response_payload.get("service_tier")
if isinstance(served_tier, str) and served_tier:
self._served_service_tier = served_tier
def _with_served_service_tier(self, chunk: "ModelResponseStream") -> "ModelResponseStream":
if self._served_service_tier is not None and chunk.model_dump().get("service_tier") is None:
setattr(chunk, "service_tier", self._served_service_tier) # noqa: B010 # pydantic extra, not a declared field
return chunk
def _with_stream_scoped_id(self, chunk: "ModelResponseStream") -> "ModelResponseStream":
if self._chat_completion_id is None:
self._chat_completion_id = chunk.id

View file

@ -26,6 +26,7 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import
TranscriptionUsageObjectTransformation,
)
from litellm.litellm_core_utils.llm_cost_calc.utils import (
_SERVICE_TIER_TO_COST_KEY_SUFFIX,
BilledTokenRates,
CostCalculatorUtils,
_generic_cost_per_character,
@ -704,7 +705,7 @@ def cost_per_token(
data_residency=data_residency,
)
elif custom_llm_provider == "databricks":
return databricks_cost_per_token(model=model, usage=usage_block)
return databricks_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
elif custom_llm_provider == "fireworks_ai":
return fireworks_ai_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "azure":
@ -969,6 +970,37 @@ def _normalize_service_tier(service_tier: object) -> str | None:
return service_tier
_BASE_PRICING_SERVICE_TIERS: Final[frozenset[str]] = frozenset({"default", "standard"})
def _resolve_billable_service_tier(requested: object, served: object) -> str | None:
"""Served tier wins when it names a priced tier or explicitly says base pricing; otherwise the request decides."""
served_lower: Final = served.lower() if isinstance(served, str) else None
if served_lower is not None and served_lower in _SERVICE_TIER_TO_COST_KEY_SUFFIX:
return served_lower
if served_lower in _BASE_PRICING_SERVICE_TIERS:
return None
return _normalize_service_tier(requested)
def _served_service_tier(completion_response: object, usage_object: Usage | None) -> str | None:
"""Find the tier the provider actually served: response, then usage, then Gemini trafficType."""
response_tier: Final = _extract_service_tier(completion_response)
if isinstance(response_tier, str):
return response_tier
usage_tier: Final = _extract_service_tier(usage_object)
if isinstance(usage_tier, str):
return usage_tier
hidden_params: Final = getattr(completion_response, "_hidden_params", None)
if hidden_params is None:
return None
provider_specific: Final = hidden_params.get("provider_specific_fields") or {}
raw_traffic_type: Final = provider_specific.get("traffic_type")
if not raw_traffic_type:
return None
return _map_traffic_type_to_service_tier(raw_traffic_type) or "default"
def _extract_service_tier(source: object) -> str | None:
"""Read a raw ``service_tier`` off a response body or usage object, dict or pydantic model alike."""
if isinstance(source, BaseModel):
@ -1388,23 +1420,14 @@ def completion_cost(
)
rerank_billed_units: RerankBilledUnits | None = None
# Extract service_tier from optional_params if not provided directly
if service_tier is None and optional_params is not None:
service_tier = optional_params.get("service_tier")
service_tier = _normalize_service_tier(service_tier)
# Extract service_tier from completion_response if not provided
if service_tier is None and completion_response is not None:
service_tier = _extract_service_tier(completion_response)
service_tier = _normalize_service_tier(service_tier)
# Extract service_tier from usage object if not provided
if service_tier is None and cost_per_token_usage_object is not None:
service_tier = _extract_service_tier(cost_per_token_usage_object)
service_tier = _normalize_service_tier(service_tier)
explicit_tier: Final = _normalize_service_tier(service_tier)
if explicit_tier is not None:
service_tier = explicit_tier
else:
service_tier = _resolve_billable_service_tier( # rebind-ok: resolved from request then response
requested=optional_params.get("service_tier") if optional_params is not None else None,
served=_served_service_tier(completion_response, cost_per_token_usage_object),
)
explicit_pricing: Final = custom_pricing is True or base_model is not None
selected_model: Final = _select_model_name_for_cost_calc(
@ -1494,15 +1517,6 @@ def completion_cost(
custom_llm_provider = hidden_params.get("custom_llm_provider", custom_llm_provider or None)
region_name = hidden_params.get("region_name", region_name)
# For Gemini/Vertex AI responses, trafficType is stored in
# provider_specific_fields. Map it to the service_tier used
# by the cost key lookup (_priority / _flex suffixes) so that
# ON_DEMAND_PRIORITY requests are billed at priority prices.
if service_tier is None:
provider_specific = hidden_params.get("provider_specific_fields") or {}
raw_traffic_type = provider_specific.get("traffic_type")
if raw_traffic_type:
service_tier = _map_traffic_type_to_service_tier(raw_traffic_type)
else:
if model is None:
raise ValueError(

View file

@ -1884,7 +1884,6 @@ class Logging(LiteLLMLoggingBaseClass):
"standard_built_in_tools_params": self.standard_built_in_tools_params,
"router_model_id": router_model_id,
"litellm_logging_obj": self,
"service_tier": (self.optional_params.get("service_tier") if self.optional_params else None),
"data_residency": (
self.litellm_params.get("data_residency")
if hasattr(self, "litellm_params") and self.litellm_params

View file

@ -109,6 +109,7 @@ class _BaseChunk(TypedDict, total=False):
created: ReadOnly[int]
model: ReadOnly[str]
system_fingerprint: ReadOnly[str | None]
service_tier: ReadOnly[str | None]
choices: ReadOnly[Required[Sequence[StreamingChoices]]]
_hidden_params: ReadOnly[_ChunkHiddenParams]
@ -369,6 +370,13 @@ class ChunkProcessor:
# Fall back to first chunk's model if no different model found
return first_chunk_model
@staticmethod
def _get_service_tier_from_chunks(chunks: Sequence["_BaseChunk"]) -> str | None:
return next(
(tier for chunk in reversed(chunks) if isinstance(tier := chunk.get("service_tier"), str) and tier),
None,
)
def build_base_response(self, chunks: Sequence["_BaseChunk"]) -> ModelResponse:
chunk = self.first_chunk
id: Final = ChunkProcessor._get_chunk_id(chunks)
@ -378,6 +386,7 @@ class ChunkProcessor:
# Get the actual model - for Azure Model Router, this finds the real model from later chunks
model: Final = ChunkProcessor._get_model_from_chunks(chunks, first_chunk_model)
system_fingerprint: Final = chunk.get("system_fingerprint", None)
service_tier: Final = ChunkProcessor._get_service_tier_from_chunks(chunks)
role: Final = ChunkProcessor._get_role_from_chunks(chunks)
finish_reason = "stop"
@ -399,6 +408,11 @@ class ChunkProcessor:
"created": created,
"model": model,
"system_fingerprint": system_fingerprint,
**(
MappingProxyType({"service_tier": service_tier})
if service_tier is not None
else MappingProxyType({})
),
"choices": [
{
"index": 0,

View file

@ -72,6 +72,12 @@ def _next_sync_or_exhausted(it: Any) -> object:
return _SYNC_ITER_EXHAUSTED
def _stamp_served_service_tier(response: ModelResponseStream, complete_streaming_response: ModelResponse) -> None:
served_tier: Final = complete_streaming_response.model_dump().get("service_tier")
if isinstance(served_tier, str) and served_tier:
setattr(response, "service_tier", served_tier) # noqa: B010 # pydantic extra, not a declared field
def is_async_iterable(obj: object) -> bool:
"""
Check if an object is an async iterable (can be used with 'async for').
@ -1876,6 +1882,7 @@ class CustomStreamWrapper:
"usage",
getattr(complete_streaming_response, "usage"),
)
_stamp_served_service_tier(response, complete_streaming_response)
try:
_cache_copy = complete_streaming_response.model_copy(deep=True)
_log_copy = complete_streaming_response.model_copy(deep=True)
@ -2127,6 +2134,7 @@ class CustomStreamWrapper:
"usage",
getattr(complete_streaming_response, "usage"),
)
_stamp_served_service_tier(response, complete_streaming_response)
try:
_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:

View file

@ -11,6 +11,7 @@ from typing import (
Final,
Literal,
Protocol,
cast,
get_args,
)
@ -35,6 +36,7 @@ from litellm.types.utils import AdapterCompletionStreamWrapper, Delta
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponseStream
@ -115,6 +117,18 @@ class _CombinedChunkSplitter:
self._async_iter: AsyncIterator[ModelResponseStream] | None = None
self._buffer: deque[ModelResponseStream] = deque()
@property
def chunks(self) -> "list[ModelResponseStream] | None":
return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream
"list[ModelResponseStream] | None", getattr(self._stream, "chunks", None)
)
@property
def messages(self) -> "list[AllMessageValues] | None":
return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream
"list[AllMessageValues] | None", getattr(self._stream, "messages", None)
)
@staticmethod
def _is_combined(chunk: "ModelResponseStream") -> bool:
"""True if ``chunk`` carries response content AND a finish_reason."""
@ -351,6 +365,18 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
text="",
)
@property
def chunks(self) -> "list[ModelResponseStream] | None":
return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream
"list[ModelResponseStream] | None", getattr(self.completion_stream, "chunks", None)
)
@property
def messages(self) -> "list[AllMessageValues] | None":
return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream
"list[AllMessageValues] | None", getattr(self.completion_stream, "messages", None)
)
def _merge_usage_into_held_stop_reason_chunk(self, chunk: Any) -> MessageBlockDelta:
"""Merge usage data from ``chunk`` into the held ``message_delta`` chunk.
@ -1173,3 +1199,37 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
return True
return False
class AnthropicSSEStream(AsyncIterator[bytes]):
"""
AsyncIterator[bytes] view of AnthropicStreamWrapper returned to callers of
translate_completion_output_params_streaming. Keeps the wrapper reachable so
the proxy's disconnect-time partial billing can read the inner chat stream's
collected chunks, messages, and model; a bare async generator would hide them.
"""
def __init__(self, anthropic_wrapper: AnthropicStreamWrapper) -> None:
self._anthropic_wrapper = anthropic_wrapper
self._byte_stream: Final[AsyncIterator[bytes]] = anthropic_wrapper.async_anthropic_sse_wrapper()
self._hidden_params: dict[
str, object
] = {} # mutable-ok: the proxy merges provider headers onto _hidden_params in place
@property
def chunks(self) -> "list[ModelResponseStream] | None":
return self._anthropic_wrapper.chunks
@property
def messages(self) -> "list[AllMessageValues] | None":
return self._anthropic_wrapper.messages
@property
def model(self) -> str:
return self._anthropic_wrapper.model
async def __anext__(self) -> bytes:
return await self._byte_stream.__anext__()
async def aclose(self) -> None:
await self._byte_stream.aclose()

View file

@ -201,7 +201,7 @@ from litellm.types.llms.openai import (
from litellm.types.utils import Choices, ModelResponse, StreamingChoices, Usage
from litellm.utils import supports_mid_conversation_system
from .streaming_iterator import AnthropicStreamWrapper
from .streaming_iterator import AnthropicSSEStream, AnthropicStreamWrapper
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
@ -341,7 +341,7 @@ class AnthropicAdapter:
)
# Return the SSE-wrapped version for proper event formatting.
if is_async:
return anthropic_wrapper.async_anthropic_sse_wrapper()
return AnthropicSSEStream(anthropic_wrapper)
return anthropic_wrapper.anthropic_sse_wrapper()

View file

@ -1,7 +1,7 @@
import re
from collections.abc import AsyncIterator, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from typing import TYPE_CHECKING, Final, cast
import litellm
from litellm._logging import verbose_logger
@ -17,6 +17,8 @@ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import (
if TYPE_CHECKING:
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponseStream
CACHED_STREAM_EVENTS_KEY: Final = "litellm_cached_anthropic_sse_events"
@ -51,6 +53,24 @@ class AnthropicMessagesStreamCacheWriter:
def has_buffered_provider_output(self) -> bool:
return getattr(self.stream, "has_buffered_provider_output", False) is True
@property
def chunks(self) -> "list[ModelResponseStream] | None":
return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream
"list[ModelResponseStream] | None", getattr(self.stream, "chunks", None)
)
@property
def messages(self) -> "list[AllMessageValues] | None":
return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream
"list[AllMessageValues] | None", getattr(self.stream, "messages", None)
)
@property
def model(self) -> str | None:
return cast( # cast-ok: model is a str on the inner stream
"str | None", getattr(self.stream, "model", None)
)
def __aiter__(self) -> "AnthropicMessagesStreamCacheWriter":
return self

View file

@ -778,6 +778,17 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator):
)
choice["delta"]["thinking_blocks"] = thinking_blocks
translated_choices.append(choice)
service_tier: Final = chunk.get("service_tier")
if isinstance(service_tier, str) and service_tier:
return ModelResponseStream(
id=chunk["id"],
object="chat.completion.chunk",
created=chunk["created"],
model=chunk["model"],
choices=translated_choices,
usage=chunk.get("usage"),
service_tier=service_tier,
)
return ModelResponseStream(
id=chunk["id"],
object="chat.completion.chunk",

View file

@ -30,7 +30,7 @@ def _registry_key(model: str) -> str:
)
def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
def cost_per_token(model: str, usage: Usage, service_tier: str | None = None) -> tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -45,4 +45,5 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
model=_registry_key(model),
usage=usage,
custom_llm_provider="databricks",
service_tier=service_tier,
)

View file

@ -890,6 +890,9 @@ class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator):
}
if "usage" in chunk and chunk["usage"] is not None:
kwargs["usage"] = chunk["usage"]
service_tier: Final = chunk.get("service_tier")
if isinstance(service_tier, str) and service_tier:
kwargs["service_tier"] = service_tier
return ModelResponseStream(**kwargs)
except Exception as e:
raise e

View file

@ -9296,9 +9296,10 @@ def _fast_serialize_simple_model_response_stream(
"object": getattr(chunk, "object", None),
"created": getattr(chunk, "created", None),
"model": model,
"service_tier": getattr(chunk, "service_tier", None),
"choices": [choice_dict],
}
for top_level_key in ("id", "object", "created"):
for top_level_key in ("id", "object", "created", "service_tier"):
if payload[top_level_key] is None:
payload.pop(top_level_key)
return orjson.dumps(payload)

View file

@ -638,6 +638,24 @@ class FallbackAwareAnthropicMessagesStream:
def has_buffered_provider_output(self) -> bool:
return getattr(self._source_iterator, "has_buffered_provider_output", False) is True
@property
def chunks(self) -> list[ModelResponseStream] | None:
return cast( # cast-ok: chunks is a list of ModelResponseStream on the inner stream
"list[ModelResponseStream] | None", getattr(self._source_iterator, "chunks", None)
)
@property
def messages(self) -> list[AllMessageValues] | None:
return cast( # cast-ok: messages is a list of AllMessageValues on the inner stream
"list[AllMessageValues] | None", getattr(self._source_iterator, "messages", None)
)
@property
def model(self) -> str | None:
return cast( # cast-ok: model is a str on the inner stream
"str | None", getattr(self._source_iterator, "model", None)
)
def adopt_fallback_source(self, fallback_response: object) -> None:
self._source_iterator = fallback_response
self.fallback_headers_adopted = True

View file

@ -81,6 +81,9 @@ ignored_function_names = [
"_merge_tools_from_deployment", # Tested indirectly via _update_kwargs_with_deployment (test files lack "router" in name)
"_invalidate_access_groups_cache", # Tested indirectly via set_model_list, upsert_model etc. (test files lack "router" in name)
"has_buffered_provider_output", # Property, so its reads in test_router.py are never an ast.Call
"chunks", # Property on FallbackAwareAnthropicMessagesStream, so its reads in tests are never an ast.Call
"messages", # Property on FallbackAwareAnthropicMessagesStream, so its reads in tests are never an ast.Call
"model", # Property on FallbackAwareAnthropicMessagesStream, so its reads in tests are never an ast.Call
"_request_header", # Tested through Claude Code session routing in test_router.py
"_claude_code_session_router_cache_key", # Tested through Claude Code session routing in test_router.py
"_delete_claude_code_session_router_binding", # Tested through Redis cleanup failure in test_router.py

View file

@ -13,6 +13,7 @@
- {id: llm.chat_completions.openai.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "gpt-4o vision; high usage"}
- {id: llm.chat_completions.openai.prompt_cache_5m.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: openai, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Prompt caching cost optimization"}
- {id: llm.chat_completions.openai.service_tier.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: service_tier, streaming: nonstream, assertions: [works], source: "OpenAI service_tier param", rationale: "OpenAI scale-tier request option is forwarded and echoed"}
- {id: llm.chat_completions.openai.service_tier.stream.echoes_served_tier, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: service_tier, streaming: stream, assertions: [works], source: "litellm_core_utils/streaming_handler.py", fail_before_fix: proven, rationale: "Every relayed stream chunk carries the service_tier OpenAI stamped on it, so a streaming caller can see which tier served the request"}
- {id: llm.chat_completions.openai.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "o-series reasoning; emerging"}
- {id: llm.chat_completions.openai.structured_output.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: openai, capability: structured_output, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "response_schema extraction"}
- {id: llm.chat_completions.anthropic.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route translated to Anthropic"}

View file

@ -61,6 +61,9 @@
- {id: quota_management.spend_tracking.stream_cache_read.bills_cache_read_rate, module: quota_management, tier: P1, behavior: spend_tracking, variant: stream_cache_read, assertions: [bills_cache_read_rate], exercised_on: [chat_completions], source: "litellm_core_utils/streaming_chunk_builder_utils.py", rationale: "A streamed call's reassembled usage keeps the cached-token detail so cache reads bill at the cache-read discount, not full input price (#34812)"}
- {id: quota_management.spend_tracking.messages_bridge.keeps_cache_tokens, module: quota_management, tier: P1, behavior: spend_tracking, variant: messages_bridge, assertions: [keeps_cache_tokens], exercised_on: [messages], source: "llms/anthropic/pass_through/responses_adapters/handler.py", rationale: "A /v1/messages request served by a Responses-only OpenAI model keeps its cache-read tokens and their discounted billing across the bridge (#34957)"}
- {id: quota_management.spend_tracking.service_tier.bills_tier_rates, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier, assertions: [bills_tier_rates], exercised_on: [chat_completions], source: "cost_calculator.py", rationale: "A priority service_tier call bills input, output, and reasoning at the deployment's *_priority rates and records the tier on the row (#35923, #35925)"}
- {id: quota_management.spend_tracking.service_tier_stream.records_served_tier, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier_stream, assertions: [records_served_tier], exercised_on: [chat_completions], source: "litellm_core_utils/streaming_chunk_builder_utils.py", fail_before_fix: proven, rationale: "A streamed call with no service_tier requested bills at the rates of the tier OpenAI stamps on its chunks and records that served tier on the row; the reassembled stream dropped the provider tier so the row recorded none and priced at the default rates"}
- {id: quota_management.spend_tracking.service_tier_stream.responses_records_served_tier, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier_stream, assertions: [records_served_tier], exercised_on: [responses], source: "responses/streaming_iterator.py", rationale: "A streamed /v1/responses call bills at the tier carried on the response.completed event's inner response and records that served tier on the spend row"}
- {id: quota_management.spend_tracking.service_tier_stream.messages_records_served_tier, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier_stream, assertions: [records_served_tier], exercised_on: [messages], source: "llms/anthropic/pass_through/adapters/streaming_iterator.py", rationale: "A streamed /v1/messages call on an OpenAI-backed deployment bills at the tier OpenAI served; the Anthropic wire format has no tier field, so the spend row is the only record of it"}
- {id: quota_management.spend_tracking.cost_headers.additive_components, module: quota_management, tier: P1, behavior: spend_tracking, variant: cost_headers, assertions: [additive_components], exercised_on: [chat_completions], source: "proxy/common_request_processing.py", rationale: "The x-litellm-response-cost-* component headers sum to the total, input covers only fresh tokens, and reasoning stays a subset of output (#36965)"}
- {id: quota_management.spend_tracking.passthrough_stream.injects_usage_cost, module: quota_management, tier: P1, behavior: spend_tracking, variant: passthrough_stream, assertions: [injects_usage_cost], exercised_on: [openai_passthrough], source: "proxy/pass_through_endpoints/streaming_handler.py", rationale: "With include_cost_in_streaming_usage on, the /openai passthrough's final streaming usage frame carries the proxy-computed cost (#36503). Uncovered: the flag is only settable in litellm_settings, and the shared e2e stack does not turn it on yet"}
- {id: quota_management.spend_tracking.websearch_interception.bills_under_request_session, module: quota_management, tier: P1, behavior: spend_tracking, variant: websearch_interception, assertions: [bills_under_request_session], exercised_on: [messages], source: "integrations/websearch_interception/handler.py", fail_before_fix: proven, rationale: "A web_search server tool the proxy intercepts into litellm.asearch writes its own asearch spend row, and that row carries the parent request's session_id so the session view counts the search and its cost next to the turn that triggered it (LIT-8063)"}

View file

@ -568,6 +568,17 @@ class AnthropicMessagesBody(BaseModel):
cache: dict[str, bool] | None = {"no-cache": True}
class ResponsesStreamBody(BaseModel):
"""POST /v1/responses body in the subset the spend tests stream with.
`input` stays a plain string: the tests only drive single-turn prompts."""
model: str
input: str
stream: bool = True
max_output_tokens: int | None = None
cache: dict[str, bool] | None = {"no-cache": True}
class CountTokensBody(BaseModel):
"""POST /v1/messages/count_tokens body: the /v1/messages shape minus
max_tokens (the endpoint only counts the prompt)."""

View file

@ -84,6 +84,7 @@ from models import (
OcrResponse,
RerankBody,
RerankResponse,
ResponsesStreamBody,
RouterCurrentValues,
RouterSettingsResponse,
SearchToolCreateBody,
@ -969,6 +970,9 @@ class ProxyClient:
def messages_stream(self, key: str, body: AnthropicMessagesBody) -> StreamingResponse:
return self.transport.stream("/v1/messages", headers=self.transport.bearer(key), json=body)
def responses_stream(self, key: str, body: ResponsesStreamBody) -> StreamingResponse:
return self.transport.stream("/v1/responses", headers=self.transport.bearer(key), json=body)
def embed(self, key: str, body: EmbedBody) -> Result[EmbedResponse]:
return self.transport.post(
"/embeddings",

View file

@ -13,27 +13,45 @@ priority processing and served the default tier, the test fails there instead of
producing a vacuous rate comparison. Reasoning is requested explicitly with
`reasoning_effort`, so the reasoning-rate assertion rests on a parameter the test
sets rather than on whatever the model happens to do by default.
The streaming cases pin the served-tier contract: OpenAI stamps the tier it actually
used on every stream chunk, and that echo is what the caller sees and what the bill
must be computed on. The request sets no service_tier, so the only place the tier
can come from is the provider's response. The spend row must record the served tier
and price input at that tier's rate, and every chunk the proxy relays must carry the
same service_tier the provider sent.
"""
import pytest
import json
import pytest
from cost_rows import (
approx_equal,
assert_fresh_tokens_billed_at,
assert_total_is_sum_of_components,
poll_cost_row,
poll_cost_row_where,
register_priced_model,
)
from e2e_config import unique_marker
from e2e_config import CHEAP_OPENAI_MODEL, unique_marker
from e2e_http import unwrap
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, LiteLLMParamsBody
from models import (
AnthropicMessagesBody,
ChatBody,
ChatMessage,
ChatStreamOptions,
LiteLLMParamsBody,
ResponsesStreamBody,
)
from pydantic import BaseModel
from spend_e2e_client import SpendClient
pytestmark = pytest.mark.e2e
BACKEND = "openai/gpt-5.6-luna"
OPENAI_API_KEY = "os.environ/OPENAI_API_KEY"
STREAM_BACKEND = f"openai/{CHEAP_OPENAI_MODEL}"
INPUT_RATE = 4e-05
OUTPUT_RATE = 8e-05
@ -42,6 +60,40 @@ PRIORITY_OUTPUT_RATE = 1.6e-04
REASONING_EFFORT = "high"
TIER_INPUT_RATES = {"default": INPUT_RATE, "priority": PRIORITY_INPUT_RATE}
class _StreamChunk(BaseModel):
id: str | None = None
service_tier: str | None = None
class _CompletedResponseObject(BaseModel):
id: str | None = None
service_tier: str | None = None
class _ResponsesStreamEvent(BaseModel):
type: str | None = None
response: _CompletedResponseObject | None = None
class _MessagesStreamEvent(BaseModel):
type: str | None = None
def _stream_chunks(events: list[str]) -> list[_StreamChunk]:
return [_StreamChunk.model_validate_json(event) for event in events if event.strip() != "[DONE]"]
def _served_tier(chunks: list[_StreamChunk]) -> str:
tiers = {chunk.service_tier for chunk in chunks if chunk.service_tier}
assert len(tiers) == 1, (
f"the relayed stream carried {tiers or 'no'} service tier(s) across {len(chunks)} chunks; OpenAI stamps "
"the served tier on every chat chunk, so exactly one tier must reach the caller"
)
return tiers.pop()
class TestServiceTierPricing:
@pytest.mark.covers("quota_management.spend_tracking.service_tier.bills_tier_rates")
@ -83,8 +135,7 @@ class TestServiceTierPricing:
)
)
assert chat.service_tier == "priority", (
f"OpenAI served tier {chat.service_tier!r} instead of priority; "
"tier billing was never exercised"
f"OpenAI served tier {chat.service_tier!r} instead of priority; tier billing was never exercised"
)
assert chat.id, f"chat response carried no id: {chat}"
@ -119,3 +170,150 @@ class TestServiceTierPricing:
)
assert_total_is_sum_of_components(row)
@pytest.mark.covers("quota_management.spend_tracking.service_tier_stream.records_served_tier")
def test_streamed_call_records_and_bills_the_served_tier(
self, client: SpendClient, resources: ResourceManager, scoped_key: str
) -> None:
model = register_priced_model(
client.proxy,
resources,
"tier-priced-stream",
LiteLLMParamsBody(
model=BACKEND,
api_key=OPENAI_API_KEY,
input_cost_per_token=INPUT_RATE,
output_cost_per_token=OUTPUT_RATE,
input_cost_per_token_priority=PRIORITY_INPUT_RATE,
output_cost_per_token_priority=PRIORITY_OUTPUT_RATE,
),
)
result = client.proxy.chat_stream(
scoped_key,
ChatBody(
model=model,
messages=[ChatMessage(role="user", content=f"{unique_marker()} reply with one word")],
max_completion_tokens=64,
stream=True,
),
)
assert result.ok and result.stream_events, (
f"streamed chat failed (status {result.status_code}): {result.body[:300]}"
)
chunks = _stream_chunks(result.stream_events)
served_tier = _served_tier(chunks)
assert served_tier in TIER_INPUT_RATES, f"no custom rate registered for served tier {served_tier!r}"
stream_id = chunks[0].id
assert stream_id, f"first stream chunk carried no id: {result.stream_events[0][:200]}"
row = poll_cost_row(client.proxy, stream_id)
assert row is not None, f"no spend row with a cost breakdown landed for {stream_id}"
assert row.breakdown.service_tier == served_tier, (
f"the provider served tier {served_tier!r} on every chunk but the bill records "
f"pricing basis {row.breakdown.service_tier!r}"
)
assert_fresh_tokens_billed_at(row, TIER_INPUT_RATES[served_tier])
assert_total_is_sum_of_components(row)
@pytest.mark.covers("llm.chat_completions.openai.service_tier.stream.echoes_served_tier")
def test_every_streamed_chunk_carries_the_served_tier(
self, client: SpendClient, resources: ResourceManager, scoped_key: str
) -> None:
model = register_priced_model(
client.proxy, resources, "tier-echo-stream", LiteLLMParamsBody(model=STREAM_BACKEND, api_key=OPENAI_API_KEY)
)
result = client.proxy.chat_stream(
scoped_key,
ChatBody(
model=model,
messages=[ChatMessage(role="user", content=f"{unique_marker()} reply with one word")],
max_completion_tokens=64,
stream=True,
stream_options=ChatStreamOptions(include_usage=True),
),
)
assert result.ok and result.stream_events, (
f"streamed chat failed (status {result.status_code}): {result.body[:300]}"
)
chunks = _stream_chunks(result.stream_events)
served_tier = _served_tier(chunks)
missing = [
json.loads(event) for event, chunk in zip(result.stream_events, chunks) if chunk.service_tier is None
]
assert not missing, (
f"{len(missing)} of {len(chunks)} relayed chunks dropped the provider's service_tier "
f"{served_tier!r}: {missing}"
)
@pytest.mark.covers("quota_management.spend_tracking.service_tier_stream.responses_records_served_tier")
def test_responses_stream_records_the_served_tier(
self, client: SpendClient, resources: ResourceManager, scoped_key: str
) -> None:
model = register_priced_model(
client.proxy,
resources,
"tier-responses-stream",
LiteLLMParamsBody(model=STREAM_BACKEND, api_key=OPENAI_API_KEY),
)
result = client.proxy.responses_stream(
scoped_key,
ResponsesStreamBody(model=model, input=f"{unique_marker()} reply with one word"),
)
assert result.ok and result.stream_events, (
f"streamed responses call failed (status {result.status_code}): {result.body[:300]}"
)
events = [_ResponsesStreamEvent.model_validate_json(event) for event in result.stream_events]
completed = next((event for event in reversed(events) if event.type == "response.completed"), None)
assert completed is not None and completed.response is not None, (
f"no response.completed event in the stream: {[e.type for e in events]}"
)
served_tier = completed.response.service_tier
assert served_tier, f"response.completed carried no service_tier: {completed.response}"
assert served_tier in TIER_INPUT_RATES, f"no custom rate registered for served tier {served_tier!r}"
row = poll_cost_row_where(client.proxy, scoped_key, lambda r: r.spend is not None and r.spend > 0)
assert row is not None, f"no spend row with a cost breakdown landed for the streamed responses call on {model}"
assert row.breakdown.service_tier == served_tier, (
f"response.completed served tier {served_tier!r} but the bill records "
f"pricing basis {row.breakdown.service_tier!r}"
)
@pytest.mark.covers("quota_management.spend_tracking.service_tier_stream.messages_records_served_tier")
def test_messages_stream_records_the_served_tier(
self, client: SpendClient, resources: ResourceManager, scoped_key: str
) -> None:
model = register_priced_model(
client.proxy,
resources,
"tier-messages-stream",
LiteLLMParamsBody(model=STREAM_BACKEND, api_key=OPENAI_API_KEY),
)
result = client.proxy.messages_stream(
scoped_key,
AnthropicMessagesBody(
model=model,
messages=[ChatMessage(role="user", content=f"{unique_marker()} reply with one word")],
max_tokens=64,
stream=True,
),
)
assert result.ok and result.stream_events, (
f"streamed messages call failed (status {result.status_code}): {result.body[:300]}"
)
events = [_MessagesStreamEvent.model_validate_json(event) for event in result.stream_events]
assert any(event.type == "message_delta" for event in events), (
f"the anthropic stream emitted no message_delta: {[e.type for e in events]}"
)
row = poll_cost_row_where(client.proxy, scoped_key, lambda r: r.spend is not None and r.spend > 0)
assert row is not None, f"no spend row with a cost breakdown landed for the streamed messages call on {model}"
served_tier = row.breakdown.service_tier
assert served_tier in TIER_INPUT_RATES and served_tier is not None, (
"the anthropic wire format carries no service_tier, so the bill is the only record of "
f"the tier OpenAI served; the row recorded pricing basis {served_tier!r}"
)

View file

@ -0,0 +1,660 @@
"""Served service_tier drives billing on streamed calls, complete and disconnected.
The scripted upstream answers OpenAI-compatible /chat/completions with SSE chunks
that carry service_tier "priority" and terminal usage. The deployment registers
distinct default and *_priority rates, so a bill computed on the wrong tier cannot
match the hand-computed expectation. /v1/messages deployments on hosted_vllm have
no anthropic-messages provider config, so they take the chat adapter: the
streamed response is an AnthropicStreamWrapper under AnthropicSSEStream, wrapped
by AnthropicMessagesStreamCacheWriter when litellm.cache is on and then by the
router's FallbackAwareAnthropicMessagesStream; each layer must delegate the
inner stream's chunks for disconnect billing to find them.
Azure streams run the same OpenAI chunk path against /openai/deployments, so the
served tier must reach the spend row there too (LIT-2850). Databricks streams go
through DatabricksChatResponseIterator.chunk_parser and the databricks branch of
cost_per_token (LIT-8121). The responses bridge relays Responses API SSE as chat
chunks, so the served tier remembered from response.created must land on both
the chunks and the row. Gemini reports capacity as usageMetadata.trafficType, which maps to
service_tier "flex" and the *_flex rates (LIT-6287, LIT-6292).
"""
import json
from collections.abc import Callable
from hashlib import sha256
from typing import Final
from uuid import uuid4
import pytest
from integration._support.client import Gateway, Scenario, eventually, object_value
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
PROMPT_TOKENS: Final = 30
COMPLETION_TOKENS: Final = 40
INPUT_RATE: Final = 0.001
OUTPUT_RATE: Final = 0.002
PRIORITY_INPUT_RATE: Final = 0.01
PRIORITY_OUTPUT_RATE: Final = 0.02
EXPECTED_FULL_SPEND: Final = PROMPT_TOKENS * PRIORITY_INPUT_RATE + COMPLETION_TOKENS * PRIORITY_OUTPUT_RATE
FLEX_INPUT_RATE: Final = 0.0005
FLEX_OUTPUT_RATE: Final = 0.001
EXPECTED_FLEX_SPEND: Final = PROMPT_TOKENS * FLEX_INPUT_RATE + COMPLETION_TOKENS * FLEX_OUTPUT_RATE
def _sse_frame(payload: dict[str, JsonValue]) -> bytes:
return f"data: {json.dumps(payload, separators=(',', ':'))}\n\n".encode()
def _chat_chunk(request_id: str, upstream_model: str, content: str, served_tier: str) -> dict[str, JsonValue]:
return {
"id": request_id,
"object": "chat.completion.chunk",
"created": 1,
"model": upstream_model,
"service_tier": served_tier,
"choices": [{"index": 0, "delta": {"role": "assistant", "content": content}, "finish_reason": None}],
}
def _respond_for(
request_id: str,
prompt: str,
*,
expected_target: str = "/v1/chat/completions",
pause: float = 0.4,
served_tier: str = "priority",
expected_requested_tier: str | None = None,
) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
if request.target == "/v1/models":
return Reply(
body=json.dumps({"object": "list", "data": [{"id": "gpt-4o-mini", "object": "model"}]}).encode()
)
assert request.target.startswith(expected_target), request.target
body: Final = json.loads(request.body)
assert body["messages"] == [{"role": "user", "content": prompt}], body
if expected_requested_tier is not None:
assert body.get("service_tier") == expected_requested_tier, body
upstream_model: Final = str(body["model"])
terminal: Final[dict[str, JsonValue]] = {
"id": request_id,
"object": "chat.completion.chunk",
"created": 1,
"model": upstream_model,
"service_tier": served_tier,
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
"usage": {
"prompt_tokens": PROMPT_TOKENS,
"completion_tokens": COMPLETION_TOKENS,
"total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS,
},
}
return Reply(
content_type="text/event-stream",
chunks=(
_sse_frame(_chat_chunk(request_id, upstream_model, "first", served_tier)),
_sse_frame(_chat_chunk(request_id, upstream_model, "second", served_tier)),
_sse_frame(_chat_chunk(request_id, upstream_model, "third", served_tier)),
_sse_frame(terminal),
b"data: [DONE]\n\n",
),
pause_between_chunks=pause,
)
return respond
def _tiered_model(
scenario: Scenario,
wire: Wire,
*,
litellm_model: str,
api_base: str | None = None,
**extra: JsonValue,
) -> str:
return scenario.model(
model=litellm_model,
api_base=api_base or f"{wire.url}/v1",
input_cost_per_token=INPUT_RATE,
output_cost_per_token=OUTPUT_RATE,
input_cost_per_token_priority=PRIORITY_INPUT_RATE,
output_cost_per_token_priority=PRIORITY_OUTPUT_RATE,
input_cost_per_token_flex=FLEX_INPUT_RATE,
output_cost_per_token_flex=FLEX_OUTPUT_RATE,
**extra,
)
def _events(lines: list[str]) -> list[dict[str, JsonValue]]:
return [
object_value(json.loads(line.removeprefix("data:")))
for line in lines
if line.startswith("data:") and line.removeprefix("data:").strip() != "[DONE]"
]
def _rows_for_key(key: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT request_id, status, prompt_tokens, completion_tokens, spend, metadata FROM "LiteLLM_SpendLogs" '
"WHERE api_key=%s",
(sha256(key.encode()).hexdigest(),),
)
def _single_spend_row(key: str) -> dict[str, JsonValue]:
rows: Final = eventually(lambda: _rows_for_key(key), lambda values: len(values) == 1, seconds=70)
return rows[0]
def _cost_breakdown(row: dict[str, JsonValue]) -> dict[str, JsonValue]:
metadata: Final = row["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
return object_value(parsed["cost_breakdown"])
@pytest.mark.timeout(120)
def test_completed_chat_stream_bills_the_served_tier(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
request_id: Final = f"chatcmpl-{uuid4().hex[:8]}"
with (
wire_server(_respond_for(request_id, prompt)) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(scenario, wire, litellm_model="openai/gpt-4o-mini")
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True},
key=key,
)
assert response.status_code == 200, response.text
chunks: Final = _events(list(response.iter_lines()))
assert len(chunks) == 4, chunks
tiers: Final = {chunk.get("service_tier") for chunk in chunks}
assert tiers == {"priority"}, f"every relayed chunk must carry the served tier: {tiers}"
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert row["request_id"] == request_id, row
assert row["prompt_tokens"] == PROMPT_TOKENS, row
assert row["completion_tokens"] == COMPLETION_TOKENS, row
assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row
breakdown: Final = _cost_breakdown(row)
assert breakdown["service_tier"] == "priority", breakdown
assert len(wire.drain()) == 1
@pytest.mark.timeout(120)
def test_disconnected_chat_stream_bills_partial_usage_at_the_served_tier(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
request_id: Final = f"chatcmpl-{uuid4().hex[:8]}"
with (
wire_server(_respond_for(request_id, prompt, pause=2.0)) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(scenario, wire, litellm_model="openai/gpt-4o-mini")
key: Final = scenario.key(models=[model])
with gateway.client.stream(
"POST",
"/v1/chat/completions",
json={
"model": model,
"messages": [{"role": "user", "content": prompt}],
"stream": True,
},
headers={"Authorization": f"Bearer {key}"},
) as response:
assert response.status_code == 200, response.read().decode()
first_event: Final = next(line for line in response.iter_lines() if line.startswith("data:"))
assert object_value(json.loads(first_event.removeprefix("data:")))["id"] == request_id
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert int(row["prompt_tokens"]) > 0, row
assert int(row["completion_tokens"]) == 1, row
assert float(str(row["spend"])) == pytest.approx(
int(row["prompt_tokens"]) * PRIORITY_INPUT_RATE + PRIORITY_OUTPUT_RATE
), row
breakdown: Final = _cost_breakdown(row)
assert breakdown["service_tier"] == "priority", breakdown
assert len(wire.drain()) == 1
@pytest.mark.timeout(120)
def test_completed_messages_stream_bills_the_served_tier(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
with (
wire_server(_respond_for(f"chatcmpl-{uuid4().hex[:8]}", prompt)) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(scenario, wire, litellm_model="hosted_vllm/gpt-4o-mini")
key: Final = scenario.key(models=[model])
with gateway.client.stream(
"POST",
"/v1/messages",
json={
"model": model,
"messages": [{"role": "user", "content": prompt}],
"max_tokens": COMPLETION_TOKENS,
"stream": True,
},
headers={"Authorization": f"Bearer {key}"},
) as response:
assert response.status_code == 200, response.read().decode()
events: Final = _events(list(response.iter_lines()))
assert events[0]["type"] == "message_start", events
assert any(event["type"] == "message_delta" for event in events), events
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert row["prompt_tokens"] == PROMPT_TOKENS, row
assert row["completion_tokens"] == COMPLETION_TOKENS, row
assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row
breakdown: Final = _cost_breakdown(row)
assert breakdown["service_tier"] == "priority", breakdown
assert len(wire.drain()) == 1
@pytest.mark.timeout(120)
def test_disconnected_messages_stream_bills_partial_usage_at_the_served_tier(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
with (
wire_server(_respond_for(f"chatcmpl-{uuid4().hex[:8]}", prompt, pause=2.0)) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(scenario, wire, litellm_model="hosted_vllm/gpt-4o-mini")
key: Final = scenario.key(models=[model])
with gateway.client.stream(
"POST",
"/v1/messages",
json={
"model": model,
"messages": [{"role": "user", "content": prompt}],
"max_tokens": COMPLETION_TOKENS,
"stream": True,
},
headers={"Authorization": f"Bearer {key}"},
) as response:
assert response.status_code == 200, response.read().decode()
first_event: Final = next(line for line in response.iter_lines() if line.startswith("data:"))
assert object_value(json.loads(first_event.removeprefix("data:")))["type"] == "message_start", first_event
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert float(str(row["spend"])) > 0, row
assert int(row["completion_tokens"]) < COMPLETION_TOKENS, row
breakdown: Final = _cost_breakdown(row)
assert breakdown["service_tier"] == "priority", breakdown
assert len(wire.drain()) == 1
def _responses_frame(event: str, payload: dict[str, JsonValue]) -> bytes:
return f"event: {event}\ndata: {json.dumps(payload, separators=(',', ':'))}\n\n".encode()
def _respond_responses_for(response_id: str, prompt: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.target == "/v1/responses", request.target
body: Final = json.loads(request.body)
assert prompt in json.dumps(body["input"]), body["input"]
assert body["stream"] is True, body
upstream_model: Final = str(body["model"])
text: Final = "firstsecondthird"
response_payload: Final[dict[str, JsonValue]] = {
"id": response_id,
"object": "response",
"model": upstream_model,
"status": "in_progress",
"service_tier": "priority",
"output": [],
}
message_item: Final[dict[str, JsonValue]] = {
"type": "message",
"id": "msg_1",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": text, "annotations": []}],
}
return Reply(
content_type="text/event-stream",
chunks=(
_responses_frame("response.created", {"type": "response.created", "response": response_payload}),
_responses_frame(
"response.output_item.added",
{
"type": "response.output_item.added",
"output_index": 0,
"item": {
"type": "message",
"id": "msg_1",
"status": "in_progress",
"role": "assistant",
"content": [],
},
},
),
_responses_frame(
"response.content_part.added",
{
"type": "response.content_part.added",
"item_id": "msg_1",
"output_index": 0,
"content_index": 0,
"part": {"type": "output_text", "text": ""},
},
),
*(
_responses_frame(
"response.output_text.delta",
{
"type": "response.output_text.delta",
"item_id": "msg_1",
"output_index": 0,
"content_index": 0,
"delta": delta,
},
)
for delta in ("first", "second", "third")
),
_responses_frame(
"response.output_text.done",
{
"type": "response.output_text.done",
"item_id": "msg_1",
"output_index": 0,
"content_index": 0,
"text": text,
},
),
_responses_frame(
"response.content_part.done",
{
"type": "response.content_part.done",
"item_id": "msg_1",
"output_index": 0,
"content_index": 0,
"part": {"type": "output_text", "text": text},
},
),
_responses_frame(
"response.output_item.done",
{"type": "response.output_item.done", "output_index": 0, "item": message_item},
),
_responses_frame(
"response.completed",
{
"type": "response.completed",
"response": {
**response_payload,
"status": "completed",
"output": [message_item],
"usage": {
"input_tokens": PROMPT_TOKENS,
"output_tokens": COMPLETION_TOKENS,
"total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS,
},
},
},
),
),
)
return respond
def _gemini_chunk(text: str) -> dict[str, JsonValue]:
return {"candidates": [{"index": 0, "content": {"role": "model", "parts": [{"text": text}]}}]}
def _respond_gemini_for(prompt: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.target.startswith("/models/gemini-2.5-flash:streamGenerateContent"), request.target
body: Final = json.loads(request.body)
assert prompt in json.dumps(body["contents"]), body["contents"]
terminal: Final[dict[str, JsonValue]] = {
"candidates": [{"index": 0, "content": {"role": "model", "parts": [{"text": ""}]}, "finishReason": "STOP"}],
"usageMetadata": {
"promptTokenCount": PROMPT_TOKENS,
"candidatesTokenCount": COMPLETION_TOKENS,
"totalTokenCount": PROMPT_TOKENS + COMPLETION_TOKENS,
"trafficType": "ON_DEMAND_FLEX",
},
}
return Reply(
content_type="text/event-stream",
chunks=(
_sse_frame(_gemini_chunk("first")),
_sse_frame(_gemini_chunk("second")),
_sse_frame(_gemini_chunk("third")),
_sse_frame(terminal),
),
)
return respond
@pytest.mark.timeout(120)
def test_azure_chat_stream_bills_the_served_tier(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
request_id: Final = f"chatcmpl-{uuid4().hex[:8]}"
with (
wire_server(
_respond_for(request_id, prompt, expected_target="/openai/deployments/gpt-4o-mini/chat/completions")
) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(
scenario,
wire,
litellm_model="azure/gpt-4o-mini",
api_base=wire.url,
api_version="2024-10-21",
)
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True},
key=key,
)
assert response.status_code == 200, response.text
chunks: Final = _events(list(response.iter_lines()))
assert len(chunks) == 4, chunks
tiers: Final = {chunk.get("service_tier") for chunk in chunks}
assert tiers == {"priority"}, f"every relayed chunk must carry the served tier: {tiers}"
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert row["prompt_tokens"] == PROMPT_TOKENS, row
assert row["completion_tokens"] == COMPLETION_TOKENS, row
assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row
breakdown: Final = _cost_breakdown(row)
assert breakdown["service_tier"] == "priority", breakdown
assert len(wire.drain()) == 1
@pytest.mark.timeout(120)
def test_databricks_chat_stream_bills_the_served_tier(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
request_id: Final = f"chatcmpl-{uuid4().hex[:8]}"
with (
wire_server(_respond_for(request_id, prompt, expected_target="/serving-endpoints/chat/completions")) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(
scenario,
wire,
litellm_model="databricks/dbrx-instruct",
api_base=f"{wire.url}/serving-endpoints",
)
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True},
key=key,
)
assert response.status_code == 200, response.text
chunks: Final = _events(list(response.iter_lines()))
assert len(chunks) == 4, chunks
tiers: Final = {chunk.get("service_tier") for chunk in chunks}
assert tiers == {"priority"}, f"every relayed chunk must carry the served tier: {tiers}"
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert row["prompt_tokens"] == PROMPT_TOKENS, row
assert row["completion_tokens"] == COMPLETION_TOKENS, row
assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row
breakdown: Final = _cost_breakdown(row)
assert breakdown["service_tier"] == "priority", breakdown
assert len(wire.drain()) == 1
@pytest.mark.timeout(120)
def test_responses_bridge_stream_bills_the_served_tier(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
with (
wire_server(_respond_responses_for(f"resp_{uuid4().hex[:8]}", prompt)) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(scenario, wire, litellm_model="openai/responses/gpt-4o-mini")
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True},
key=key,
)
assert response.status_code == 200, response.text
chunks: Final = _events(list(response.iter_lines()))
assert len(chunks) >= 4, chunks
tiers: Final = {chunk.get("service_tier") for chunk in chunks}
assert tiers == {"priority"}, f"every relayed chunk must carry the served tier: {tiers}"
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert row["prompt_tokens"] == PROMPT_TOKENS, row
assert row["completion_tokens"] == COMPLETION_TOKENS, row
assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row
breakdown: Final = _cost_breakdown(row)
assert breakdown["service_tier"] == "priority", breakdown
assert len(wire.drain()) == 1
@pytest.mark.timeout(120)
def test_gemini_chat_stream_bills_the_flex_tier(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
with (
wire_server(_respond_gemini_for(prompt)) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(
scenario,
wire,
litellm_model="gemini/gemini-2.5-flash",
api_base=wire.url,
)
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": prompt}], "stream": True},
key=key,
)
assert response.status_code == 200, response.text
chunks: Final = _events(list(response.iter_lines()))
assert len(chunks) >= 2, chunks
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert row["prompt_tokens"] == PROMPT_TOKENS, row
assert row["completion_tokens"] == COMPLETION_TOKENS, row
assert float(str(row["spend"])) == pytest.approx(EXPECTED_FLEX_SPEND), row
breakdown: Final = _cost_breakdown(row)
assert breakdown["service_tier"] == "flex", breakdown
assert len(wire.drain()) == 1
@pytest.mark.timeout(120)
def test_requested_priority_downgraded_to_default_bills_base_rates(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
request_id: Final = f"chatcmpl-{uuid4().hex[:8]}"
with (
wire_server(
_respond_for(request_id, prompt, served_tier="default", expected_requested_tier="priority")
) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(scenario, wire, litellm_model="openai/gpt-4o-mini")
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": prompt}],
"stream": True,
"service_tier": "priority",
},
key=key,
)
assert response.status_code == 200, response.text
chunks: Final = _events(list(response.iter_lines()))
assert len(chunks) == 4, chunks
tiers: Final = {chunk.get("service_tier") for chunk in chunks}
assert tiers == {"default"}, f"every relayed chunk must carry the served tier: {tiers}"
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert row["request_id"] == request_id, row
assert row["prompt_tokens"] == PROMPT_TOKENS, row
assert row["completion_tokens"] == COMPLETION_TOKENS, row
assert float(str(row["spend"])) == pytest.approx(
PROMPT_TOKENS * INPUT_RATE + COMPLETION_TOKENS * OUTPUT_RATE
), row
breakdown: Final = _cost_breakdown(row)
assert breakdown.get("service_tier") != "priority", breakdown
assert len(wire.drain()) == 1
@pytest.mark.timeout(120)
def test_requested_priority_with_auto_echo_bills_priority(gateway: Gateway) -> None:
prompt: Final = f"tier control {uuid4().hex[:8]}"
request_id: Final = f"chatcmpl-{uuid4().hex[:8]}"
with (
wire_server(_respond_for(request_id, prompt, served_tier="auto")) as wire,
gateway.scenario() as scenario,
):
model: Final = _tiered_model(scenario, wire, litellm_model="openai/gpt-4o-mini")
key: Final = scenario.key(models=[model])
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": prompt}],
"stream": True,
"service_tier": "priority",
},
key=key,
)
assert response.status_code == 200, response.text
chunks: Final = _events(list(response.iter_lines()))
assert len(chunks) == 4, chunks
row: Final = _single_spend_row(key)
assert row["status"] == "success", row
assert row["request_id"] == request_id, row
assert float(str(row["spend"])) == pytest.approx(EXPECTED_FULL_SPEND), row
breakdown: Final = _cost_breakdown(row)
assert breakdown["service_tier"] == "priority", breakdown
assert len(wire.drain()) == 1

View file

@ -2058,3 +2058,13 @@ async def test_queue_request_stream_is_untouched_while_keepalives_are_unconfigur
assert not any(chunk.startswith(b": ping") for chunk in chunks)
assert chunks[-1] == b"data: [DONE]\n\n"
def test_fast_serialize_simple_model_response_stream_keeps_served_service_tier():
chunk = _simple_chunk()
chunk.service_tier = "priority"
result = _fast_serialize_simple_model_response_stream(chunk)
assert result is not None
assert json.loads(result)["service_tier"] == "priority"

View file

@ -64,6 +64,7 @@ from litellm.proxy._types import ProxyErrorTypes, ProxyException
from litellm.proxy._types import UserAPIKeyAuth as ProxyUserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.router import Router
from litellm.router_utils.add_retry_fallback_headers import prepare_response_for_header_attachment
def test_attach_guardrail_information_copies_recorded_entries_onto_model_response():
@ -7337,6 +7338,61 @@ class TestStreamingClientDisconnectBilling:
assert standard_logging_object["total_tokens"] > 0
assert standard_logging_object["response_cost"] >= 0.002
@pytest.mark.asyncio
async def test_disconnect_bills_partial_spend_for_anthropic_adapter_stream(self):
"""
The proxy's cleanup gets the FallbackAwareAnthropicMessagesStream the
router returns for /v1/messages; its chunks/messages must delegate
through the translate_completion_output_params_streaming result to the
inner chat stream's collected chunks or a disconnect bills nothing.
"""
from litellm.llms.anthropic.pass_through.adapters.streaming_iterator import (
AnthropicSSEStream,
)
from litellm.llms.anthropic.pass_through.adapters.transformation import (
AnthropicAdapter,
)
from litellm.router import FallbackAwareAnthropicMessagesStream
async def _sse_frames() -> AsyncGenerator[bytes, None]:
yield b"event: message_start\n\n"
recorder = _RecordingSuccessLogger()
original_callbacks = litellm.callbacks
litellm.callbacks = [recorder]
try:
response = await self._start_partial_stream()
setattr(response.chunks[-1], "service_tier", "priority") # noqa: B010 # pydantic extra, not a declared field
source_iterator: Final = AnthropicAdapter().translate_completion_output_params_streaming(
response,
model=response.model or "gpt-4o-mini",
is_async=True,
litellm_logging_obj=response.logging_obj,
)
assert isinstance(source_iterator, AnthropicSSEStream)
streamed: Final = prepare_response_for_header_attachment(
FallbackAwareAnthropicMessagesStream(_sse_frames(), source_iterator)
)
billed: Final = await _bill_partial_streamed_spend_on_disconnect(
{"litellm_logging_obj": response.logging_obj},
streamed,
)
for _ in range(50):
if recorder.success_events:
break
await asyncio.sleep(0.1)
await asyncio.sleep(0.5)
finally:
litellm.callbacks = original_callbacks
assert billed is True
assert len(recorder.success_events) == 1
partial_response: Final = recorder.success_events[0]["response_obj"]
assert getattr(partial_response, "service_tier") == "priority"
assert partial_response.usage.total_tokens > 0
@pytest.mark.asyncio
async def test_completed_stream_does_not_double_bill_on_late_disconnect(self):
recorder = _RecordingSuccessLogger()

View file

@ -4352,3 +4352,39 @@ def test_map_optional_params_verbosity_merges_into_text():
verbosity_only_request,
)
assert verbosity_only_request["text"] == {"verbosity": "low"}
def test_response_completed_carries_the_served_service_tier():
from litellm.completion_extras.litellm_responses_transformation.transformation import (
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
result = iterator.chunk_parser(
{
"type": "response.completed",
"response": {"id": "resp_1", "status": "completed", "output": [], "service_tier": "default"},
}
)
assert result.model_dump()["service_tier"] == "default"
def test_every_bridged_chunk_after_response_created_carries_the_served_service_tier():
from litellm.completion_extras.litellm_responses_transformation.transformation import (
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
events = [
{"type": "response.created", "response": {"id": "resp_1", "status": "in_progress", "service_tier": "default"}},
{"type": "response.output_item.added", "output_index": 0, "item": {"type": "message"}},
{"type": "response.output_text.delta", "output_index": 0, "delta": "Hi"},
{"type": "response.output_item.done", "output_index": 0, "item": {"type": "message"}},
{"type": "response.completed", "response": {"id": "resp_1", "status": "completed", "output": []}},
]
relayed = [iterator.chunk_parser(event).model_dump().get("service_tier") for event in events]
assert relayed == ["default"] * len(events), relayed

View file

@ -8811,6 +8811,58 @@ async def test_async_failure_handler_delivers_failure_payload_to_custom_logger()
assert events.empty()
def test_responses_completed_event_bills_the_served_service_tier():
"""The served service_tier on response.completed's inner ResponsesAPIResponse
must reach the cost calculator, so a priority-served stream prices at the
priority rates instead of the default tier's."""
logging_obj: Final = LitellmLogging(
model="openai/gpt-5.1",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="aresponses",
start_time=time.time(),
litellm_call_id="resp-served-tier",
function_id="resp-served-tier",
)
logging_obj.update_environment_variables(
model="openai/gpt-5.1",
user="",
optional_params={},
litellm_params={},
custom_llm_provider="openai",
)
inner: Final = ResponsesAPIResponse(
id="resp-served-tier",
created_at=1,
object="response",
status="completed",
model="gpt-5.1",
output=[],
usage=ResponseAPIUsage(input_tokens=10, output_tokens=20, total_tokens=30),
service_tier="priority",
)
event: Final = ResponseCompletedEvent(type="response.completed", response=inner)
cost: Final = logging_obj._response_cost_calculator(result=event) # pyright: ignore[reportPrivateUsage] # parity with the suite's own direct calls
billed_response: Final = ModelResponse(
model="gpt-5.1",
usage=litellm.Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30),
)
tier_cost: Final = litellm.completion_cost(
completion_response=billed_response,
model="openai/gpt-5.1",
service_tier="priority",
)
default_cost: Final = litellm.completion_cost(
completion_response=billed_response,
model="openai/gpt-5.1",
)
assert cost == pytest.approx(tier_cost)
assert cost > default_cost
def _image_logging_obj() -> LitellmLogging:
logging_obj = LitellmLogging(
model="gpt-image-2",

View file

@ -1820,3 +1820,34 @@ def test_calculate_usage_keeps_a_reported_count_over_a_later_chunks_zero() -> No
)
assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (5, 17, 22)
def _tier_chunk(content: str, service_tier: str | None, finish_reason: str | None = None) -> ModelResponseStream:
return ModelResponseStream(
id="chatcmpl-tier",
created=1,
model="gpt-4.1-mini",
object="chat.completion.chunk",
choices=[StreamingChoices(finish_reason=finish_reason, index=0, delta=Delta(content=content, role=None))],
**({"service_tier": service_tier} if service_tier is not None else {}),
)
def test_stream_chunk_builder_records_the_last_service_tier_the_provider_stamped():
chunks = [
_tier_chunk("Hel", "auto"),
_tier_chunk("lo", None),
_tier_chunk("", "default", finish_reason="stop"),
]
response = stream_chunk_builder(chunks=chunks)
assert response is not None
assert response.model_dump()["service_tier"] == "default"
def test_stream_chunk_builder_omits_service_tier_when_no_chunk_carried_one():
response = stream_chunk_builder(chunks=[_tier_chunk("Hi", None, finish_reason="stop")])
assert response is not None
assert "service_tier" not in response.model_dump()

View file

@ -4983,3 +4983,50 @@ async def test_async_stream_without_usage_counts_tokens_off_the_event_loop():
assert chunks[-1].usage.prompt_tokens > 100_000
assert chunks[-1].usage.completion_tokens > 100_000
assert_loop_stayed_free(took, lags)
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.asyncio
async def test_openai_stream_relays_the_served_service_tier_on_every_chunk_including_usage(
logging_obj: Logging, sync_mode: bool
):
from litellm.utils import ModelResponseListIterator
def _chunk(content: str, finish_reason: str | None, usage: Usage | None, choices: bool = True):
return ModelResponseStream(
id="chatcmpl-tier",
created=1742056047,
model="gpt-4.1-mini",
choices=[StreamingChoices(finish_reason=finish_reason, index=0, delta=Delta(content=content))]
if choices
else [],
usage=usage,
service_tier="default",
)
logging_obj.update_environment_variables(
model="gpt-4.1-mini",
optional_params={"stream_options": {"include_usage": True}},
litellm_params={},
custom_llm_provider="openai",
)
wrapper = CustomStreamWrapper(
completion_stream=ModelResponseListIterator(
model_responses=[
_chunk("Hi", None, None),
_chunk("", "stop", None),
_chunk("", None, Usage(prompt_tokens=10, completion_tokens=1, total_tokens=11), choices=False),
]
),
model="gpt-4.1-mini",
custom_llm_provider="openai",
logging_obj=logging_obj,
stream_options={"include_usage": True},
)
relayed = (
[chunk.model_dump() for chunk in wrapper] if sync_mode else [chunk.model_dump() async for chunk in wrapper]
)
assert [chunk.get("service_tier") for chunk in relayed] == ["default"] * len(relayed), relayed
assert relayed[-1]["usage"]["total_tokens"] == 11

View file

@ -0,0 +1,87 @@
"""
Tests for AnthropicSSEStream, the object translate_completion_output_params_streaming
hands to the proxy for /v1/messages streaming. It must emit the same SSE bytes as
the wrapper's async_anthropic_sse_wrapper, propagate aclose into it, and expose the
wrapper's chunks/messages/model so disconnect-time partial billing can read them.
"""
from typing import Final
from unittest.mock import MagicMock
import pytest
from litellm.llms.anthropic.pass_through.adapters.streaming_iterator import (
AnthropicSSEStream,
AnthropicStreamWrapper,
)
from litellm.types.utils import Delta, StreamingChoices
def _make_chunk(delta: Delta, finish_reason: str | None = None) -> MagicMock:
chunk = MagicMock()
chunk.choices = [StreamingChoices(finish_reason=finish_reason, index=0, delta=delta, logprobs=None)]
chunk.usage = None
chunk._hidden_params = {}
return chunk
class _AsyncStream:
def __init__(self, items: list[MagicMock]):
self._it = iter(items)
self.chunks = list(items)
self.messages: list[dict] = [{"role": "user", "content": "hi"}]
def __aiter__(self):
return self
async def __anext__(self):
try:
return next(self._it)
except StopIteration:
raise StopAsyncIteration
def _streamed_events() -> AnthropicSSEStream:
upstream: Final = _AsyncStream(
[
_make_chunk(Delta(content="Once")),
_make_chunk(Delta(content=" upon"), finish_reason="stop"),
]
)
wrapper: Final = AnthropicStreamWrapper(completion_stream=upstream, model="gpt-4o-mini")
wrapper._message_id = "msg_test"
return AnthropicSSEStream(wrapper)
@pytest.mark.asyncio
async def test_sse_stream_yields_identical_bytes_to_the_wrappers_sse_wrapper():
upstream_a: Final = _AsyncStream(
[_make_chunk(Delta(content="Once")), _make_chunk(Delta(content=" upon"), finish_reason="stop")]
)
wrapper_a: Final = AnthropicStreamWrapper(completion_stream=upstream_a, model="gpt-4o-mini")
wrapper_a._message_id = "msg_test"
expected: Final = [event async for event in wrapper_a.async_anthropic_sse_wrapper()]
actual: Final = [event async for event in _streamed_events()]
assert actual == expected
@pytest.mark.asyncio
async def test_sse_stream_aclose_ends_the_wrapped_stream():
stream: Final = _streamed_events()
first: Final = await stream.__anext__()
assert first.startswith(b"event: message_start")
await stream.aclose()
with pytest.raises(StopAsyncIteration):
await stream.__anext__()
def test_sse_stream_exposes_chunks_messages_and_model():
stream: Final = _streamed_events()
assert stream.model == "gpt-4o-mini"
assert stream.messages == [{"role": "user", "content": "hi"}]
chunks: Final = stream.chunks
assert isinstance(chunks, list) and len(chunks) == 2

View file

@ -10,6 +10,7 @@ import litellm
from litellm._internal_context import in_post_response_phase
from litellm.caching.caching import Cache, LiteLLMCacheType
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.anthropic.pass_through.messages import handler
from litellm.llms.anthropic.pass_through.messages.response_cache import (
AnthropicMessagesStreamCacheWriter,
@ -63,6 +64,12 @@ async def _collect(stream: AsyncIterator[bytes]) -> List[bytes]:
return [chunk async for chunk in stream]
@pytest.fixture(autouse=True)
async def _drain_logging_worker():
yield
await GLOBAL_LOGGING_WORKER.flush()
@pytest.fixture
def local_cache():
previous_cache = litellm.cache
@ -282,6 +289,40 @@ class _HeldBackStream:
raise StopAsyncIteration
class _AttributedStream:
"""Stream stub carrying the billing attributes the disconnect helper reads."""
def __init__(self, chunks: list) -> None:
self.chunks = [object()]
self.messages = [{"role": "user", "content": "hi"}]
self.model = "gpt-4o-mini"
self._pending = list(chunks)
def __aiter__(self) -> "_AttributedStream":
return self
async def __anext__(self) -> bytes:
if not self._pending:
raise StopAsyncIteration
return self._pending.pop(0)
@pytest.mark.asyncio
async def test_cache_writer_exposes_inner_stream_billing_attributes(request_kwargs):
caching_handler = LLMCachingHandler(
original_function=handler.anthropic_messages,
request_kwargs=dict(request_kwargs),
start_time=datetime.datetime.now(),
)
inner = _AttributedStream(STREAM_EVENTS)
writer = AnthropicMessagesStreamCacheWriter(stream=inner, caching_handler=caching_handler)
assert writer.chunks is inner.chunks
assert writer.messages is inner.messages
assert writer.model == "gpt-4o-mini"
assert await _collect(writer) == STREAM_EVENTS
@pytest.mark.asyncio
async def test_stream_cache_write_runs_in_post_response_phase(request_kwargs, monkeypatch):
"""Every event, message_stop included, is already with the client when the stream write

View file

@ -883,3 +883,13 @@ def test_completion_merges_system_messages_when_one_has_empty_content(respx_mock
{"role": "system", "content": "You are terse."},
{"role": "user", "content": "Hello"},
]
def test_chunk_parser_relays_the_served_service_tier():
iterator = DatabricksChatResponseIterator(streaming_response=None, sync_stream=True)
with_tier: Final = iterator.chunk_parser({**_streaming_chunk(), "service_tier": "priority"})
assert with_tier.model_dump()["service_tier"] == "priority"
without_tier: Final = iterator.chunk_parser(_streaming_chunk())
assert getattr(without_tier, "service_tier", None) is None

View file

@ -156,8 +156,6 @@ def test_uncached_request_bills_every_prompt_token_at_the_input_rate(local_model
assert completion_cost == pytest.approx(200 * info["output_cost_per_token"])
@pytest.mark.parametrize("model", NEW_MODELS)
def test_new_models_carry_cache_pricing(local_model_cost_map: None, model: str) -> None:
info: Final = _model_info(model)
@ -232,3 +230,28 @@ def test_sonnet_5_ships_standard_rates_not_introductory(local_model_cost_map: No
for field in PRICE_FIELDS:
assert sonnet_5[field] == pytest.approx(sonnet_4_6[field]), field
def test_cost_per_token_bills_the_served_priority_tier(
local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch
) -> None:
rates: Final = {
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
"input_cost_per_token_priority": 0.01,
"output_cost_per_token_priority": 0.02,
"litellm_provider": "databricks",
"mode": "chat",
}
monkeypatch.setitem(litellm.model_cost, "databricks/dbrx-tiered-test", rates)
usage: Final = Usage(prompt_tokens=30, completion_tokens=40, total_tokens=70)
prompt_cost, completion_cost = cost_per_token(
model="databricks/dbrx-tiered-test", usage=usage, service_tier="priority"
)
assert prompt_cost == pytest.approx(30 * 0.01)
assert completion_cost == pytest.approx(40 * 0.02)
prompt_cost, completion_cost = cost_per_token(model="databricks/dbrx-tiered-test", usage=usage)
assert prompt_cost == pytest.approx(30 * 0.001)
assert completion_cost == pytest.approx(40 * 0.002)

View file

@ -248,6 +248,33 @@ class TestOpenAIChatCompletionStreamingHandler:
assert result.usage.completion_tokens == 350
assert result.usage.total_tokens == 14147
def test_chunk_parser_preserves_service_tier(self):
"""OpenAI-compatible upstreams serve a service_tier on every streamed
chunk; chunk_parser must keep it on the emitted ModelResponseStream so
disconnect billing and the reassembled response see the served tier."""
handler = OpenAIChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
tiered_chunk = {
"id": "gen-123",
"created": 1234567890,
"model": "openai/gpt-4o-mini",
"object": "chat.completion.chunk",
"choices": [
{
"index": 0,
"delta": {"role": "assistant", "content": ""},
"finish_reason": None,
}
],
"service_tier": "priority",
}
plain_chunk = {key: value for key, value in tiered_chunk.items() if key != "service_tier"}
assert handler.chunk_parser(tiered_chunk).model_dump().get("service_tier") == "priority"
assert handler.chunk_parser(plain_chunk).model_dump().get("service_tier") is None
def test_chunk_parser_raises_on_in_body_error_payload(self):
"""vLLM/sglang return HTTP 200 streams whose body carries the error,
e.g. data: {"error": {..., "code": 400}}. chunk_parser must surface it

View file

@ -1948,7 +1948,7 @@ def test_completion_cost_extracts_service_tier_from_usage(_local_model_cost_map)
def test_completion_cost_service_tier_priority(_local_model_cost_map):
"""Test that service_tier extraction follows priority: optional_params > completion_response > usage."""
"""Test that the served tier wins over the requested tier: response > usage > request."""
from litellm import completion_cost
# Test with gpt-5-nano which has flex pricing
@ -1965,7 +1965,7 @@ def test_completion_cost_service_tier_priority(_local_model_cost_map):
)
setattr(response, "service_tier", "priority")
# Test that optional_params takes priority over response and usage
# A request-level tier loses to the tier the response actually served
cost_from_params = completion_cost(
completion_response=response,
model=model,
@ -1973,20 +1973,18 @@ def test_completion_cost_service_tier_priority(_local_model_cost_map):
optional_params={"service_tier": "flex"},
)
# Test that response takes priority over usage when optional_params is not provided
completion_cost(
# Response takes priority over usage
cost_served_priority = completion_cost(
completion_response=response,
model=model,
custom_llm_provider="openai",
)
# Test that usage is used when neither optional_params nor response have service_tier
# Create a new response without service_tier attribute
# Create a new response without service_tier attribute so it falls back to usage
response_no_tier = ModelResponse(
usage=usage,
model=model,
)
# Don't set service_tier on response, so it will fall back to usage
cost_from_usage = completion_cost(
completion_response=response_no_tier,
@ -1994,12 +1992,13 @@ def test_completion_cost_service_tier_priority(_local_model_cost_map):
custom_llm_provider="openai",
)
# All should use flex pricing (from different sources)
assert cost_from_params > 0, "Cost from params should be greater than 0"
assert cost_from_usage > 0, "Cost from usage should be greater than 0"
# Costs should be similar (all using flex)
assert abs(cost_from_params - cost_from_usage) < 1e-6, "Costs from params and usage should be similar (both flex)"
# Requested flex is ignored once the response reports served priority
assert cost_from_params == pytest.approx(cost_served_priority), (
"request-level service_tier must defer to the served tier on the response"
)
def test_completion_cost_service_tier_for_bedrock(_local_model_cost_map):
@ -5468,3 +5467,100 @@ def test_completion_cost_is_zero_when_explicit_rates_are_zero(monkeypatch: pytes
)
assert cost == 0.0
@pytest.mark.parametrize(
("requested", "served", "expected"),
[
(None, "priority", "priority"),
("priority", "flex", "flex"),
("priority", "default", None),
("priority", "standard", None),
("priority", "auto", "priority"),
("priority", "scale", "priority"),
("priority", None, "priority"),
("auto", None, None),
(None, "Priority", "priority"),
("flex", "on_demand", "flex"),
],
)
def test_resolve_billable_service_tier(requested: object, served: object, expected: str | None) -> None:
from litellm.cost_calculator import _resolve_billable_service_tier
assert _resolve_billable_service_tier(requested=requested, served=served) == expected
def _served_tier_cost_model(monkeypatch: pytest.MonkeyPatch) -> str:
model: Final = "served-tier-cost-model"
monkeypatch.setitem(
litellm.model_cost,
model,
{
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
"input_cost_per_token_priority": 0.01,
"output_cost_per_token_priority": 0.02,
"litellm_provider": "openai",
"mode": "chat",
},
)
return model
def test_completion_cost_bills_base_when_served_default_overrides_requested_priority(
_local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch
) -> None:
model: Final = _served_tier_cost_model(monkeypatch)
response: Final = ModelResponse(
model=model,
usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150),
)
setattr(response, "service_tier", "default")
cost: Final = completion_cost(
completion_response=response,
model=model,
custom_llm_provider="openai",
optional_params={"service_tier": "priority"},
)
assert cost == pytest.approx(100 * 0.001 + 50 * 0.002)
def test_completion_cost_bills_priority_when_served_tier_overrides_missing_request(
_local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch
) -> None:
model: Final = _served_tier_cost_model(monkeypatch)
response: Final = ModelResponse(
model=model,
usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150),
)
setattr(response, "service_tier", "priority")
cost: Final = completion_cost(
completion_response=response,
model=model,
custom_llm_provider="openai",
)
assert cost == pytest.approx(100 * 0.01 + 50 * 0.02)
def test_completion_cost_bills_base_when_gemini_serves_on_demand(
_local_model_cost_map: None, monkeypatch: pytest.MonkeyPatch
) -> None:
model: Final = _served_tier_cost_model(monkeypatch)
response: Final = ModelResponse(
model=model,
usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150),
)
response._hidden_params["provider_specific_fields"] = {"traffic_type": "ON_DEMAND"}
cost: Final = completion_cost(
completion_response=response,
model=model,
custom_llm_provider="openai",
optional_params={"service_tier": "priority"},
)
assert cost == pytest.approx(100 * 0.001 + 50 * 0.002)