mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
273489824a
commit
e814532033
32 changed files with 1621 additions and 53 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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)"}
|
||||
|
|
|
|||
|
|
@ -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)."""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
660
tests/integration/spend/test_service_tier_stream_billing.py
Normal file
660
tests/integration/spend/test_service_tier_stream_billing.py
Normal 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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue