mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
Merge remote-tracking branch 'origin/main' into litellm_fix_mcp_jwt_oauth_persistence
This commit is contained in:
commit
ece1de73b4
137 changed files with 6047 additions and 5642 deletions
|
|
@ -96,6 +96,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/langfuse/",
|
||||
"/vllm/",
|
||||
"/mistral/",
|
||||
"/nvidia_nim/",
|
||||
"/groq/",
|
||||
"/voyage/",
|
||||
"/cursor/",
|
||||
|
|
|
|||
|
|
@ -40,4 +40,4 @@ if not logger.handlers:
|
|||
logging.Formatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s")
|
||||
)
|
||||
logger.addHandler(handler)
|
||||
logger.setLevel(logging.INFO)
|
||||
logger.setLevel(os.getenv("LITELLM_LOG", "INFO").upper())
|
||||
|
|
|
|||
|
|
@ -401,10 +401,14 @@ def _parse_json_logs_env(value: str | None) -> bool:
|
|||
return (value or "").lower() == "true"
|
||||
|
||||
|
||||
def resolve_log_level(log_level: str) -> int:
|
||||
return getattr(logging, log_level.upper())
|
||||
|
||||
|
||||
json_logs: Final = _parse_json_logs_env(os.getenv("JSON_LOGS"))
|
||||
# Create a handler for the logger (you may need to adapt this based on your needs)
|
||||
log_level: Final = os.getenv("LITELLM_LOG", "DEBUG")
|
||||
numeric_level: Final[str] = getattr(logging, log_level.upper())
|
||||
numeric_level: Final[int] = resolve_log_level(log_level)
|
||||
handler: Final = LevelRoutingStreamHandler()
|
||||
handler.setLevel(numeric_level)
|
||||
handler.addFilter(_secret_filter)
|
||||
|
|
|
|||
|
|
@ -1203,6 +1203,24 @@ def _without_provider_stated_cost(usage: Usage | None) -> Usage | None:
|
|||
return usage.model_copy(update=MappingProxyType({"cost": None}))
|
||||
|
||||
|
||||
def _split_responses_ws_logging_object_by_service_tier(
|
||||
completion_response: LiteLLMRealtimeStreamLoggingObject,
|
||||
) -> tuple[LiteLLMRealtimeStreamLoggingObject, ...] | None:
|
||||
partition: Final = ResponsesWebSocketTokenUsageProcessor.partition_results_by_service_tier(
|
||||
cast(Sequence[Mapping[str, object]], completion_response.results)
|
||||
)
|
||||
if len(partition) <= 1:
|
||||
return None
|
||||
return tuple(
|
||||
LiteLLMRealtimeStreamLoggingObject(
|
||||
results=cast(OpenAIRealtimeStreamList, list(group)),
|
||||
usage=ResponsesWebSocketTokenUsageProcessor.collect_and_combine_usage_from_responses_ws_results(group),
|
||||
service_tier=tier,
|
||||
)
|
||||
for tier, group in partition.items()
|
||||
)
|
||||
|
||||
|
||||
def completion_cost(
|
||||
completion_response: object | None = None,
|
||||
model: str | None = None,
|
||||
|
|
@ -1266,6 +1284,41 @@ def completion_cost(
|
|||
try:
|
||||
call_type = _infer_call_type(call_type, completion_response) or "completion"
|
||||
|
||||
if call_type == CallTypes.aresponses_websocket.value and isinstance(
|
||||
completion_response, LiteLLMRealtimeStreamLoggingObject
|
||||
):
|
||||
ws_tier_parts: Final = _split_responses_ws_logging_object_by_service_tier(completion_response)
|
||||
if ws_tier_parts is not None:
|
||||
return sum(
|
||||
completion_cost(
|
||||
completion_response=part,
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
messages=messages,
|
||||
completion=completion,
|
||||
total_time=total_time,
|
||||
call_type=call_type,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
region_name=region_name,
|
||||
size=size,
|
||||
quality=quality,
|
||||
n=n,
|
||||
custom_cost_per_token=custom_cost_per_token,
|
||||
custom_cost_per_second=custom_cost_per_second,
|
||||
optional_params=optional_params,
|
||||
custom_pricing=custom_pricing,
|
||||
base_model=base_model,
|
||||
standard_built_in_tools_params=standard_built_in_tools_params,
|
||||
litellm_model_name=litellm_model_name,
|
||||
router_model_id=router_model_id,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
service_tier=service_tier,
|
||||
data_residency=data_residency,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
for part in ws_tier_parts
|
||||
)
|
||||
|
||||
if (
|
||||
(call_type == "aimage_generation" or call_type == "image_generation")
|
||||
and model is not None
|
||||
|
|
@ -1466,12 +1519,15 @@ def completion_cost(
|
|||
duration_seconds = usage_obj.get("duration_seconds", None)
|
||||
_vr = usage_obj.get("video_resolution", None)
|
||||
provider_reported_cost = usage_obj.get("provider_reported_cost_usd", None)
|
||||
_vc = usage_obj.get("video_count", None)
|
||||
else:
|
||||
duration_seconds = getattr(usage_obj, "duration_seconds", None)
|
||||
_vr = getattr(usage_obj, "video_resolution", None)
|
||||
provider_reported_cost = getattr(usage_obj, "provider_reported_cost_usd", None)
|
||||
_vc = getattr(usage_obj, "video_count", None)
|
||||
if _vr is not None:
|
||||
video_resolution = str(_vr).strip().lower()
|
||||
video_count = _vc if isinstance(_vc, int) and not isinstance(_vc, bool) and _vc > 1 else 1
|
||||
|
||||
if _video_model_info is None and provider_reported_cost is not None:
|
||||
return float(provider_reported_cost)
|
||||
|
|
@ -1482,12 +1538,15 @@ def completion_cost(
|
|||
video_generation_cost,
|
||||
)
|
||||
|
||||
return video_generation_cost(
|
||||
model=model,
|
||||
duration_seconds=duration_seconds,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_info=_video_model_info,
|
||||
video_resolution=video_resolution,
|
||||
return (
|
||||
video_generation_cost(
|
||||
model=model,
|
||||
duration_seconds=duration_seconds,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_info=_video_model_info,
|
||||
video_resolution=video_resolution,
|
||||
)
|
||||
* video_count
|
||||
)
|
||||
# Fallback to default video cost calculation if no duration available
|
||||
return default_video_cost_calculator(
|
||||
|
|
@ -2558,6 +2617,7 @@ _RESPONSES_WS_BILLABLE_EVENT_TYPES: Final = frozenset({"response.completed", "re
|
|||
|
||||
class _ResponsesWsEventResponse(BaseModel):
|
||||
usage: Mapping[str, object] | None = None
|
||||
service_tier: str | None = None
|
||||
|
||||
|
||||
class _ResponsesWsEvent(BaseModel):
|
||||
|
|
@ -2565,20 +2625,39 @@ class _ResponsesWsEvent(BaseModel):
|
|||
response: _ResponsesWsEventResponse | None = None
|
||||
|
||||
|
||||
def _billable_responses_ws_events(
|
||||
results: Sequence[Mapping[str, object]],
|
||||
) -> tuple[tuple[Mapping[str, object], _ResponsesWsEventResponse], ...]:
|
||||
return tuple(
|
||||
(result, event.response)
|
||||
for result in results
|
||||
if (event := _ResponsesWsEvent.model_validate(result)).type in _RESPONSES_WS_BILLABLE_EVENT_TYPES
|
||||
and event.response is not None
|
||||
and event.response.usage is not None
|
||||
)
|
||||
|
||||
|
||||
class ResponsesWebSocketTokenUsageProcessor(BaseTokenUsageProcessor):
|
||||
@staticmethod
|
||||
def collect_usage_from_responses_ws_results(
|
||||
results: Sequence[Mapping[str, object]],
|
||||
) -> tuple[Usage, ...]:
|
||||
events: Final = tuple(_ResponsesWsEvent.model_validate(result) for result in results)
|
||||
return tuple(
|
||||
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( # pyright: ignore[reportPrivateUsage] # same shared transform the realtime processor uses
|
||||
event.response.usage
|
||||
response.usage
|
||||
)
|
||||
for event in events
|
||||
if event.type in _RESPONSES_WS_BILLABLE_EVENT_TYPES
|
||||
and event.response is not None
|
||||
and event.response.usage is not None
|
||||
for _, response in _billable_responses_ws_events(results)
|
||||
if response.usage is not None
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def partition_results_by_service_tier(
|
||||
results: Sequence[Mapping[str, object]],
|
||||
) -> Mapping[str | None, tuple[Mapping[str, object], ...]]:
|
||||
billable: Final = _billable_responses_ws_events(results)
|
||||
tiers: Final = dict.fromkeys(response.service_tier for _, response in billable)
|
||||
return MappingProxyType(
|
||||
{tier: tuple(result for result, response in billable if response.service_tier == tier) for tier in tiers}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -2101,9 +2101,14 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
results=result # pyright: ignore[reportUnknownArgumentType] # raw event dicts from the WS stream
|
||||
)
|
||||
)
|
||||
ws_tier_partition: Final = ResponsesWebSocketTokenUsageProcessor.partition_results_by_service_tier(
|
||||
results=result # pyright: ignore[reportUnknownArgumentType] # raw event dicts from the WS stream
|
||||
)
|
||||
ws_service_tier: Final = next(iter(ws_tier_partition)) if len(ws_tier_partition) == 1 else None
|
||||
logging_result = LiteLLMRealtimeStreamLoggingObject(
|
||||
usage=combined_ws_usage,
|
||||
results=result, # pyright: ignore[reportUnknownArgumentType] # raw event dicts from the WS stream
|
||||
service_tier=ws_service_tier,
|
||||
)
|
||||
|
||||
elif (
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from types import MappingProxyType
|
|||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.azure_ai.common_utils import (
|
||||
|
|
@ -18,6 +18,8 @@ from litellm.llms.base_llm.passthrough.transformation import (
|
|||
BasePassthroughConfig,
|
||||
RelayShape,
|
||||
logged_relay_shape,
|
||||
model_group_from,
|
||||
relayed_body,
|
||||
strip_leading_model_segment,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
|
@ -35,19 +37,6 @@ if TYPE_CHECKING:
|
|||
EMPTY_QUERY: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
class PassthroughMetadata(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
model_group: str = ""
|
||||
|
||||
|
||||
def model_group_from(litellm_params: Mapping[str, object]) -> str:
|
||||
try:
|
||||
return PassthroughMetadata.model_validate(litellm_params.get("litellm_metadata")).model_group
|
||||
except ValidationError:
|
||||
return ""
|
||||
|
||||
|
||||
def api_version_from(litellm_params: Mapping[str, object]) -> str | None:
|
||||
try:
|
||||
return TypeAdapter(str | None).validate_python(litellm_params.get("api_version"))
|
||||
|
|
@ -96,14 +85,6 @@ def relay_query_params(
|
|||
return MappingProxyType({**(request_query_params or EMPTY_QUERY), "api-version": api_version})
|
||||
|
||||
|
||||
def relayed_body(httpx_response: Response) -> str | dict:
|
||||
try:
|
||||
body: Final[object] = httpx_response.json()
|
||||
except ValueError:
|
||||
return httpx_response.text
|
||||
return body if isinstance(body, dict) else httpx_response.text
|
||||
|
||||
|
||||
FOUNDRY_RELAY_SHAPES: Final = (
|
||||
RelayShape("/rerank", CallTypes.arerank, RerankResponse.model_validate),
|
||||
RelayShape("/providers/blackforestlabs/v1/flux-2-pro", CallTypes.aimage_generation, ImageResponse.model_validate),
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from collections.abc import Callable, Mapping, Sequence
|
|||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final, Protocol, TypeAlias
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
|
@ -29,6 +29,19 @@ if TYPE_CHECKING:
|
|||
RELAYED_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
class PassthroughMetadata(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
model_group: str = ""
|
||||
|
||||
|
||||
def model_group_from(litellm_params: Mapping[str, object]) -> str:
|
||||
try:
|
||||
return PassthroughMetadata.model_validate(litellm_params.get("litellm_metadata")).model_group
|
||||
except ValidationError:
|
||||
return ""
|
||||
|
||||
|
||||
def strip_leading_model_segment(endpoint: str, model_names: tuple[str, ...]) -> str:
|
||||
path: Final = endpoint.lstrip("/")
|
||||
for model_name in model_names:
|
||||
|
|
@ -55,6 +68,14 @@ def relayed_json_object(httpx_response: Response) -> Mapping[str, object] | None
|
|||
return None
|
||||
|
||||
|
||||
def relayed_body(httpx_response: Response) -> str | dict:
|
||||
try:
|
||||
body: Final[object] = httpx_response.json()
|
||||
except ValueError:
|
||||
return httpx_response.text
|
||||
return body if isinstance(body, dict) else httpx_response.text
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RelayShape:
|
||||
path_suffix: str
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import litellm
|
|||
from litellm.constants import DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS
|
||||
from litellm.images.utils import ImageEditRequestUtils
|
||||
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
|
||||
from litellm.llms.vertex_ai.videos.transformation import veo_video_count_from_parameters
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.gemini import (
|
||||
GeminiLongRunningOperationResponse,
|
||||
|
|
@ -354,6 +355,9 @@ class GeminiVideoConfig(BaseVideoConfig):
|
|||
video_resolution: Final = _usage_video_resolution_from_parameters(parameters)
|
||||
if video_resolution is not None:
|
||||
usage_data["video_resolution"] = video_resolution
|
||||
video_count: Final = veo_video_count_from_parameters(parameters)
|
||||
if video_count is not None:
|
||||
usage_data["video_count"] = video_count
|
||||
|
||||
video_obj.usage = usage_data
|
||||
return video_obj
|
||||
|
|
|
|||
0
litellm/llms/nvidia_nim/passthrough/__init__.py
Normal file
0
litellm/llms/nvidia_nim/passthrough/__init__.py
Normal file
139
litellm/llms/nvidia_nim/passthrough/transformation.py
Normal file
139
litellm/llms/nvidia_nim/passthrough/transformation.py
Normal file
|
|
@ -0,0 +1,139 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from collections.abc import Collection, Iterable, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.passthrough.transformation import (
|
||||
BasePassthroughConfig,
|
||||
model_group_from,
|
||||
relayed_body,
|
||||
strip_leading_model_segment,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.router import DeploymentTypedDict
|
||||
from litellm.types.utils import LlmProviders, StandardPassThroughResponseObject
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from httpx import URL, Response
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.llms.base_llm.passthrough.transformation import LoggedRelayResponse
|
||||
|
||||
|
||||
API_VERSION_SEGMENT: Final = re.compile(r"^v\d+$")
|
||||
NVIDIA_NIM_MODEL_PREFIX: Final = f"{LlmProviders.NVIDIA_NIM.value}/"
|
||||
NVIDIA_NIM_ROUTE_PREFIX: Final = re.compile(rf"^/{LlmProviders.NVIDIA_NIM.value}/", re.IGNORECASE)
|
||||
|
||||
|
||||
def is_nvidia_nim_deployment(deployment: DeploymentTypedDict) -> bool:
|
||||
litellm_params: Final = deployment["litellm_params"]
|
||||
return litellm_params.get("custom_llm_provider") == LlmProviders.NVIDIA_NIM.value or litellm_params.get(
|
||||
"model", ""
|
||||
).startswith(NVIDIA_NIM_MODEL_PREFIX)
|
||||
|
||||
|
||||
def nvidia_nim_model_groups(deployments: Iterable[DeploymentTypedDict] | None) -> frozenset[str]:
|
||||
listed: Final = tuple(deployments or ())
|
||||
nim_groups: Final = frozenset(d["model_name"] for d in listed if is_nvidia_nim_deployment(d))
|
||||
other_groups: Final = frozenset(d["model_name"] for d in listed if not is_nvidia_nim_deployment(d))
|
||||
return nim_groups - other_groups
|
||||
|
||||
|
||||
def nvidia_nim_model_group_in_path(path: str, deployments: Iterable[DeploymentTypedDict] | None) -> str | None:
|
||||
return nvidia_nim_router_model_in_endpoint(
|
||||
NVIDIA_NIM_ROUTE_PREFIX.sub("", path), nvidia_nim_model_groups(deployments)
|
||||
)
|
||||
|
||||
|
||||
def nvidia_nim_router_model_in_endpoint(endpoint: str, router_models: Collection[str]) -> str | None:
|
||||
segments: Final = tuple(segment for segment in endpoint.split("/") if segment)
|
||||
return next(
|
||||
(
|
||||
"/".join(segments[:length])
|
||||
for length in range(len(segments), 0, -1)
|
||||
if "/".join(segments[:length]) in router_models
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def without_repeated_version_prefix(api_base: str, native_endpoint: str) -> str:
|
||||
url: Final = httpx.URL(api_base)
|
||||
base_segments: Final = tuple(segment for segment in url.path.split("/") if segment)
|
||||
first_native_segment: Final = native_endpoint.lstrip("/").split("/", 1)[0]
|
||||
repeated: Final = (
|
||||
bool(base_segments)
|
||||
and API_VERSION_SEGMENT.match(first_native_segment) is not None
|
||||
and base_segments[-1] == first_native_segment
|
||||
)
|
||||
kept_segments: Final = base_segments[:-1] if repeated else base_segments
|
||||
return str(url.copy_with(path="/" + "/".join(kept_segments), query=None)).rstrip("/")
|
||||
|
||||
|
||||
class NvidiaNimPassthroughConfig(BasePassthroughConfig):
|
||||
def is_streaming_request(self, endpoint: str, request_data: dict) -> bool:
|
||||
return bool(request_data.get("stream", False))
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
endpoint: str,
|
||||
request_query_params: dict | None,
|
||||
litellm_params: dict,
|
||||
) -> tuple[URL, str]:
|
||||
base_target_url: Final = self.get_api_base(api_base)
|
||||
if base_target_url is None:
|
||||
raise ValueError("NVIDIA NIM api base not found: set `api_base` on the deployment or NVIDIA_NIM_API_BASE")
|
||||
native_endpoint: Final = strip_leading_model_segment(endpoint, (model, model_group_from(litellm_params)))
|
||||
root: Final = without_repeated_version_prefix(base_target_url, native_endpoint)
|
||||
return (self.format_url(native_endpoint, root, request_query_params), root)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
model: str,
|
||||
messages: Sequence[AllMessageValues],
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict[str, str]: # mutable-ok: base class contract returns dict for httpx
|
||||
if api_key is None:
|
||||
return dict(headers) # mutable-ok: base class contract returns dict for httpx
|
||||
return {
|
||||
**headers,
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
} # mutable-ok: base class contract returns dict for httpx
|
||||
|
||||
@staticmethod
|
||||
def get_api_base(api_base: str | None = None) -> str | None:
|
||||
return api_base or get_secret_str("NVIDIA_NIM_API_BASE")
|
||||
|
||||
@staticmethod
|
||||
def get_api_key(api_key: str | None = None) -> str | None:
|
||||
return api_key or get_secret_str("NVIDIA_NIM_API_KEY")
|
||||
|
||||
@staticmethod
|
||||
def get_base_model(model: str) -> str | None:
|
||||
return model
|
||||
|
||||
def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]:
|
||||
return []
|
||||
|
||||
def logging_non_streaming_response(
|
||||
self,
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
httpx_response: Response,
|
||||
request_data: Mapping[str, object],
|
||||
logging_obj: Logging,
|
||||
endpoint: str,
|
||||
) -> LoggedRelayResponse | OCRResponse | StandardPassThroughResponseObject | None:
|
||||
return StandardPassThroughResponseObject(response=relayed_body(httpx_response))
|
||||
|
|
@ -70,10 +70,17 @@ def _parse_veo_operation(raw_response: httpx.Response) -> _VeoOperation:
|
|||
return operation
|
||||
|
||||
|
||||
def veo_video_count_from_parameters(parameters: Mapping[str, object]) -> int | None:
|
||||
sample_count: Final = parameters.get("sampleCount")
|
||||
if isinstance(sample_count, bool) or not isinstance(sample_count, int) or sample_count < 1:
|
||||
return None
|
||||
return sample_count
|
||||
|
||||
|
||||
def _build_vertex_video_usage_from_request_data(
|
||||
request_data: dict[str, Any] | None,
|
||||
) -> dict[str, float | str]:
|
||||
"""Build usage metadata (duration, resolution) for video cost calculation."""
|
||||
"""Build usage metadata (duration, resolution, video count) for video cost calculation."""
|
||||
usage_data: Final[dict[str, float | str]] = {}
|
||||
if not request_data:
|
||||
return usage_data
|
||||
|
|
@ -88,6 +95,9 @@ def _build_vertex_video_usage_from_request_data(
|
|||
res: Final = parameters.get("resolution")
|
||||
if res is not None and str(res).strip() != "":
|
||||
usage_data["video_resolution"] = str(res).strip().lower()
|
||||
video_count: Final = veo_video_count_from_parameters(parameters)
|
||||
if video_count is not None:
|
||||
usage_data["video_count"] = video_count
|
||||
return usage_data
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -58151,6 +58151,23 @@
|
|||
"model_info": {
|
||||
"supports_reasoning": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "gemini-chat-baseline",
|
||||
"pattern": "gemini-(?!.*(?:-tts|-image|-live|-audio|-embedding|-computer-use|-robotics|-transcribe|-translate))(?:2[.-][5-9]|[3-9](?:[.-]\\d{1,2})?)-(?:pro|flash)(?:-lite)?(?![a-z])",
|
||||
"description": "Any Gemini text-chat id at 2.5 or higher under any namespace, including bare ids, gemini/, vertex_ai/, openrouter/google/, deepinfra/google/, vercel_ai_gateway/google/, oci/google., and databricks-gemini-<major>-<minor>: gemini-<major>[.minor]-(pro|flash)[-lite] with any trailing preview, date or variant tag. The capability flags were verified against each of those providers' own catalogs and docs. The lookahead excludes the tts, image, live, audio, embedding, computer-use, robotics, transcribe and translate lines, which are different modes with different capabilities. Provider-specific deviations, such as Perplexity's Agent API serving these as mode responses, are carried by their exact map entries, which always win over this rule. Carries no token limits or pricing, so those stay on the standard unmapped behavior rather than a guessed number. Source check 2026-09-15: all 45 first-party 2.5+ text-chat entries in this map carry every field below, and the OpenRouter (openrouter.ai/api/v1/models), Vercel AI Gateway (ai-gateway.vercel.sh/v1/models), DeepInfra (api.deepinfra.com/models/list), OCI and Databricks model docs list reasoning, tools and image input for the same models.",
|
||||
"model_info": {
|
||||
"mode": "chat",
|
||||
"supports_reasoning": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_vision": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_web_search": true
|
||||
}
|
||||
}
|
||||
]
|
||||
},
|
||||
|
|
|
|||
|
|
@ -205,6 +205,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
|
|||
"/gigachat/",
|
||||
"/milvus/",
|
||||
"/mistral/",
|
||||
"/nvidia_nim/",
|
||||
"/openai/",
|
||||
"/openai_passthrough/",
|
||||
"/vertex-ai/",
|
||||
|
|
|
|||
|
|
@ -9986,7 +9986,7 @@
|
|||
},
|
||||
"unreachable_fallback": {
|
||||
"default": "fail_closed",
|
||||
"description": "Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.",
|
||||
"description": "Behavior when a guardrail endpoint is unreachable due to network errors. Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. 'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed.",
|
||||
"enum": [
|
||||
"fail_closed",
|
||||
"fail_open"
|
||||
|
|
@ -10948,6 +10948,18 @@
|
|||
"description": "Custom advisory message template used when on_flagged='inject_system_message'. Must contain a {reason} placeholder. Defaults to a generic advisory message if unset.",
|
||||
"title": "Advisory System Message"
|
||||
},
|
||||
"agent_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Agent identity reported to Agent 365 with every tool evaluation. When unset, the caller's key alias is used.",
|
||||
"title": "Agent Id"
|
||||
},
|
||||
"akto_account_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
@ -11450,6 +11462,30 @@
|
|||
"title": "Chunk Budget Chars",
|
||||
"type": "integer"
|
||||
},
|
||||
"client_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Client id of the gateway's Entra app registration (a confidential client). Falls back to the AGENT365_CLIENT_ID environment variable.",
|
||||
"title": "Client Id"
|
||||
},
|
||||
"client_secret": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Client secret of the gateway's Entra app registration, used to perform the On-Behalf-Of exchange. Falls back to the AGENT365_CLIENT_SECRET environment variable.",
|
||||
"title": "Client Secret"
|
||||
},
|
||||
"confidence_threshold": {
|
||||
"default": 0.5,
|
||||
"default_value": 0.5,
|
||||
|
|
@ -12496,6 +12532,18 @@
|
|||
"description": "The message the bot speaks aloud when a /v1/realtime guardrail fires. Falls back to violation_message_template if not set.",
|
||||
"title": "Realtime Violation Message"
|
||||
},
|
||||
"resource_app_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Application id of the Agent 365 resource the OBO token is minted for. Defaults to the production resource ea9ffc3e-8a23-4a7d-836d-234d7c7565c1; the Test and PreProd environments use a different id. Falls back to the AGENT365_RESOURCE_APP_ID environment variable.",
|
||||
"title": "Resource App Id"
|
||||
},
|
||||
"rules": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
@ -12733,6 +12781,18 @@
|
|||
"description": "The ID of your Model Armor template",
|
||||
"title": "Template Id"
|
||||
},
|
||||
"tenant_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"description": "Entra tenant id used for the On-Behalf-Of token exchange. Falls back to the AGENT365_TENANT_ID environment variable.",
|
||||
"title": "Tenant Id"
|
||||
},
|
||||
"timeout": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
@ -18945,6 +19005,228 @@
|
|||
]
|
||||
}
|
||||
},
|
||||
"/nvidia_nim/{endpoint}": {
|
||||
"delete": {
|
||||
"description": "Relay a native NVIDIA NIM request through a LiteLLM model group.\n\n`{PROXY_BASE_URL}/nvidia_nim/{model_group}/v1/infer` forwards the body unchanged to the deployment's\n`api_base`, so object detection and OCR NIMs whose payload carries no `model` field still go through\nvirtual key auth, model access checks, and spend logging.",
|
||||
"operationId": "nvidia_nim_proxy_route_nvidia_nim__endpoint__delete",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
"name": "endpoint",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"title": "Endpoint",
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Nvidia Nim Proxy Route",
|
||||
"tags": [
|
||||
"llm_passthrough"
|
||||
]
|
||||
},
|
||||
"get": {
|
||||
"description": "Relay a native NVIDIA NIM request through a LiteLLM model group.\n\n`{PROXY_BASE_URL}/nvidia_nim/{model_group}/v1/infer` forwards the body unchanged to the deployment's\n`api_base`, so object detection and OCR NIMs whose payload carries no `model` field still go through\nvirtual key auth, model access checks, and spend logging.",
|
||||
"operationId": "nvidia_nim_proxy_route_nvidia_nim__endpoint__get",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
"name": "endpoint",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"title": "Endpoint",
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Nvidia Nim Proxy Route",
|
||||
"tags": [
|
||||
"llm_passthrough"
|
||||
]
|
||||
},
|
||||
"patch": {
|
||||
"description": "Relay a native NVIDIA NIM request through a LiteLLM model group.\n\n`{PROXY_BASE_URL}/nvidia_nim/{model_group}/v1/infer` forwards the body unchanged to the deployment's\n`api_base`, so object detection and OCR NIMs whose payload carries no `model` field still go through\nvirtual key auth, model access checks, and spend logging.",
|
||||
"operationId": "nvidia_nim_proxy_route_nvidia_nim__endpoint__patch",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
"name": "endpoint",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"title": "Endpoint",
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Nvidia Nim Proxy Route",
|
||||
"tags": [
|
||||
"llm_passthrough"
|
||||
]
|
||||
},
|
||||
"post": {
|
||||
"description": "Relay a native NVIDIA NIM request through a LiteLLM model group.\n\n`{PROXY_BASE_URL}/nvidia_nim/{model_group}/v1/infer` forwards the body unchanged to the deployment's\n`api_base`, so object detection and OCR NIMs whose payload carries no `model` field still go through\nvirtual key auth, model access checks, and spend logging.",
|
||||
"operationId": "nvidia_nim_proxy_route_nvidia_nim__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
"name": "endpoint",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"title": "Endpoint",
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Nvidia Nim Proxy Route",
|
||||
"tags": [
|
||||
"llm_passthrough"
|
||||
]
|
||||
},
|
||||
"put": {
|
||||
"description": "Relay a native NVIDIA NIM request through a LiteLLM model group.\n\n`{PROXY_BASE_URL}/nvidia_nim/{model_group}/v1/infer` forwards the body unchanged to the deployment's\n`api_base`, so object detection and OCR NIMs whose payload carries no `model` field still go through\nvirtual key auth, model access checks, and spend logging.",
|
||||
"operationId": "nvidia_nim_proxy_route_nvidia_nim__endpoint__put",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
"name": "endpoint",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"title": "Endpoint",
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Nvidia Nim Proxy Route",
|
||||
"tags": [
|
||||
"llm_passthrough"
|
||||
]
|
||||
}
|
||||
},
|
||||
"/openai/deployments/{model}/chat/completions": {
|
||||
"post": {
|
||||
"description": "Follows the exact same API spec as `OpenAI's Chat API https://platform.openai.com/docs/api-reference/chat`\n\n```bash\ncurl -X POST http://localhost:4000/v1/chat/completions \n-H \"Content-Type: application/json\" \n-H \"Authorization: Bearer sk-1234\" \n-d '{\n \"model\": \"gpt-4o\",\n \"messages\": [\n {\n \"role\": \"user\",\n \"content\": \"Hello!\"\n }\n ]\n}'\n```",
|
||||
|
|
|
|||
|
|
@ -246,6 +246,7 @@ class Litellm_EntityType(enum.Enum):
|
|||
TEAM = "team"
|
||||
TEAM_MEMBER = "team_member"
|
||||
ORGANIZATION = "organization"
|
||||
ORGANIZATION_MEMBER = "organization_member"
|
||||
PROJECT = "project"
|
||||
TAG = "tag"
|
||||
AGENT = "agent"
|
||||
|
|
@ -485,6 +486,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/milvus",
|
||||
"/gigachat",
|
||||
"/watsonx",
|
||||
"/nvidia_nim",
|
||||
]
|
||||
|
||||
#########################################################
|
||||
|
|
@ -5256,6 +5258,7 @@ class DBSpendUpdateTransactions(TypedDict):
|
|||
team_list_transactions: dict[str, float] | None
|
||||
team_member_list_transactions: dict[str, float] | None
|
||||
org_list_transactions: dict[str, float] | None
|
||||
org_member_list_transactions: ReadOnly[dict[str, float] | None]
|
||||
tag_list_transactions: dict[str, float] | None
|
||||
agent_list_transactions: dict[str, float] | None
|
||||
model_access_group_list_transactions: ReadOnly[dict[str, float] | None]
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from litellm.litellm_core_utils.url_utils import (
|
|||
validate_url,
|
||||
)
|
||||
from litellm.llms.azure.passthrough.transformation import azure_router_model_in_endpoint
|
||||
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.common_utils.http_parsing_utils import extract_nested_form_metadata
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
|
|
@ -976,6 +977,26 @@ def _get_deployment_default_tpm_limit(model_name: str) -> int | None:
|
|||
return _get_deployment_default_limit(model_name, "default_api_key_tpm_limit")
|
||||
|
||||
|
||||
def get_key_own_model_rate_limit(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"],
|
||||
) -> dict[str, int] | None:
|
||||
if user_api_key_dict.metadata:
|
||||
result: Final = user_api_key_dict.metadata.get(rate_limit_key)
|
||||
if result:
|
||||
return result
|
||||
|
||||
if not user_api_key_dict.model_max_budget:
|
||||
return None
|
||||
budget_key: Final = "rpm_limit" if rate_limit_key == "model_rpm_limit" else "tpm_limit"
|
||||
model_limit: Final = {
|
||||
model: budget[budget_key]
|
||||
for model, budget in user_api_key_dict.model_max_budget.items()
|
||||
if isinstance(budget, dict) and budget.get(budget_key) is not None
|
||||
}
|
||||
return model_limit or None
|
||||
|
||||
|
||||
def get_key_model_rpm_limit(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
model_name: str | None = None,
|
||||
|
|
@ -989,20 +1010,9 @@ def get_key_model_rpm_limit(
|
|||
3. Team metadata (model_rpm_limit)
|
||||
4. Deployment default_api_key_rpm_limit (when model_name is provided)
|
||||
"""
|
||||
# 1. Check key metadata first (takes priority)
|
||||
if user_api_key_dict.metadata:
|
||||
result: Final = user_api_key_dict.metadata.get("model_rpm_limit")
|
||||
if result:
|
||||
return result
|
||||
|
||||
# 2. Check model_max_budget
|
||||
if user_api_key_dict.model_max_budget:
|
||||
model_rpm_limit: Final[dict[str, int]] = {}
|
||||
for model, budget in user_api_key_dict.model_max_budget.items():
|
||||
if isinstance(budget, dict) and budget.get("rpm_limit") is not None:
|
||||
model_rpm_limit[model] = budget["rpm_limit"]
|
||||
if model_rpm_limit:
|
||||
return model_rpm_limit
|
||||
key_own_limit: Final = get_key_own_model_rate_limit(user_api_key_dict, "model_rpm_limit")
|
||||
if key_own_limit is not None:
|
||||
return key_own_limit
|
||||
|
||||
# 3. Fallback to team metadata
|
||||
if user_api_key_dict.team_metadata:
|
||||
|
|
@ -1032,20 +1042,9 @@ def get_key_model_tpm_limit(
|
|||
3. Team metadata (model_tpm_limit)
|
||||
4. Deployment default_api_key_tpm_limit (when model_name is provided)
|
||||
"""
|
||||
# 1. Check key metadata first (takes priority)
|
||||
if user_api_key_dict.metadata:
|
||||
result: Final = user_api_key_dict.metadata.get("model_tpm_limit")
|
||||
if result:
|
||||
return result
|
||||
|
||||
# 2. Check model_max_budget (iterate per-model like RPM does)
|
||||
if user_api_key_dict.model_max_budget:
|
||||
model_tpm_limit: Final[dict[str, int]] = {}
|
||||
for model, budget in user_api_key_dict.model_max_budget.items():
|
||||
if isinstance(budget, dict) and budget.get("tpm_limit") is not None:
|
||||
model_tpm_limit[model] = budget["tpm_limit"]
|
||||
if model_tpm_limit:
|
||||
return model_tpm_limit
|
||||
key_own_limit: Final = get_key_own_model_rate_limit(user_api_key_dict, "model_tpm_limit")
|
||||
if key_own_limit is not None:
|
||||
return key_own_limit
|
||||
|
||||
# 3. Fallback to team metadata
|
||||
if user_api_key_dict.team_metadata:
|
||||
|
|
@ -2045,6 +2044,12 @@ def get_model_from_request(
|
|||
azure_model: Final = _router_model_from_azure_route(route, llm_router)
|
||||
return model if azure_model is None else azure_model
|
||||
|
||||
if route.lower().startswith("/nvidia_nim/"):
|
||||
nvidia_nim_model: Final = (
|
||||
nvidia_nim_model_group_in_path(route, llm_router.get_model_list()) if llm_router else None
|
||||
)
|
||||
return model if nvidia_nim_model is None else nvidia_nim_model
|
||||
|
||||
return model
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -155,8 +155,8 @@ class LicenseCheck:
|
|||
|
||||
def auto_router_capability_limit(self) -> int | None:
|
||||
"""
|
||||
How many auto-routers may claim each licensed capability (heuristic_v2, operator-defined
|
||||
tier_definitions): unlimited (None) only when the signed license lists the auto_router
|
||||
How many auto-routers may claim each gated classifier or customization capability:
|
||||
unlimited (None) only when the signed license lists the auto_router
|
||||
feature, otherwise one per capability. A license verified through the API carries no
|
||||
feature list, so it does not lift the limit either.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from collections.abc import Mapping, Sequence
|
|||
from datetime import datetime, timedelta, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload
|
||||
from urllib.parse import quote, unquote
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -85,6 +86,10 @@ else:
|
|||
RESPONSES_SESSION_CALL_TYPES: Final = frozenset({CallTypes.responses.value, CallTypes.aresponses.value})
|
||||
|
||||
|
||||
def _org_member_transaction_key(org_id: str, user_id: str) -> str:
|
||||
return f"organization_id::{quote(org_id, safe='')}::user_id::{quote(user_id, safe='')}"
|
||||
|
||||
|
||||
def _is_batch_cost_row(payload: SpendLogsPayload) -> bool:
|
||||
return payload.get("call_type") == CallTypes.aretrieve_batch.value and payload.get("status") == "success"
|
||||
|
||||
|
|
@ -110,6 +115,7 @@ class _SpendBatch(Protocol):
|
|||
litellm_teamtable: BatchTable
|
||||
litellm_teammembership: BatchTable
|
||||
litellm_organizationtable: BatchTable
|
||||
litellm_organizationmembership: BatchTable
|
||||
litellm_tagtable: BatchTable
|
||||
litellm_agentstable: BatchTable
|
||||
litellm_modelaccessgroupbudgettable: BatchTable
|
||||
|
|
@ -666,6 +672,7 @@ class DBSpendUpdateWriter:
|
|||
await self._update_org_db(
|
||||
response_cost=response_cost,
|
||||
org_id=org_id,
|
||||
user_id=user_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
except Exception:
|
||||
|
|
@ -900,6 +907,7 @@ class DBSpendUpdateWriter:
|
|||
self,
|
||||
response_cost: float | None,
|
||||
org_id: str | None,
|
||||
user_id: str | None,
|
||||
prisma_client: PrismaClient | None,
|
||||
):
|
||||
try:
|
||||
|
|
@ -916,6 +924,15 @@ class DBSpendUpdateWriter:
|
|||
response_cost=response_cost,
|
||||
)
|
||||
)
|
||||
|
||||
if user_id is not None:
|
||||
await self.spend_update_queue.add_update(
|
||||
update=SpendUpdateQueueItem(
|
||||
entity_type=Litellm_EntityType.ORGANIZATION_MEMBER,
|
||||
entity_id=_org_member_transaction_key(org_id, user_id),
|
||||
response_cost=response_cost,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
spend_log_error(
|
||||
"Spend tracking - failed to enqueue org spend update. org_id=%s, response_cost=%s - %s",
|
||||
|
|
@ -1163,14 +1180,15 @@ class DBSpendUpdateWriter:
|
|||
if db_spend_update_transactions is not None:
|
||||
verbose_proxy_logger.info(
|
||||
"Spend tracking - committing spend updates from Redis to DB: "
|
||||
"keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, tags=%d, agents=%d, "
|
||||
"model_access_groups=%d",
|
||||
"keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, org_members=%d, tags=%d, "
|
||||
"agents=%d, model_access_groups=%d",
|
||||
len(db_spend_update_transactions.get("key_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("user_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("team_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("org_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("end_user_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("team_member_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("org_member_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("tag_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("agent_list_transactions") or {}),
|
||||
len(db_spend_update_transactions.get("model_access_group_list_transactions") or {}),
|
||||
|
|
@ -1708,6 +1726,29 @@ class DBSpendUpdateWriter:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
org_member_list_transactions: Final = db_spend_update_transactions.get("org_member_list_transactions")
|
||||
verbose_proxy_logger.debug("Org Membership Spend transactions: %s", org_member_list_transactions)
|
||||
if org_member_list_transactions is not None and len(org_member_list_transactions.keys()) > 0:
|
||||
for i in range(n_retry_times + 1):
|
||||
start_time = time.time()
|
||||
try:
|
||||
async with _spend_update_tx(prisma_client) as transaction, transaction.batch_() as batcher:
|
||||
for key, response_cost in sorted(org_member_list_transactions.items()):
|
||||
_, quoted_org_id, _, quoted_user_id = key.split("::")
|
||||
batcher.litellm_organizationmembership.update_many(
|
||||
where={"organization_id": unquote(quoted_org_id), "user_id": unquote(quoted_user_id)},
|
||||
data={"spend": {"increment": response_cost}},
|
||||
)
|
||||
break
|
||||
except Exception as e:
|
||||
await self._handle_spend_update_failure(
|
||||
e=e,
|
||||
attempt=i,
|
||||
n_retry_times=n_retry_times,
|
||||
start_time=start_time,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
### UPDATE TAG TABLE ###
|
||||
tag_list_transactions: Final = db_spend_update_transactions["tag_list_transactions"]
|
||||
await DBSpendUpdateWriter._update_entity_spend_in_db(
|
||||
|
|
|
|||
|
|
@ -69,6 +69,7 @@ _SpendTransactionField: TypeAlias = Literal[
|
|||
"team_list_transactions",
|
||||
"team_member_list_transactions",
|
||||
"org_list_transactions",
|
||||
"org_member_list_transactions",
|
||||
"tag_list_transactions",
|
||||
"agent_list_transactions",
|
||||
"model_access_group_list_transactions",
|
||||
|
|
@ -81,6 +82,7 @@ _SPEND_TRANSACTION_FIELDS: Final[tuple[_SpendTransactionField, ...]] = (
|
|||
"team_list_transactions",
|
||||
"team_member_list_transactions",
|
||||
"org_list_transactions",
|
||||
"org_member_list_transactions",
|
||||
"tag_list_transactions",
|
||||
"agent_list_transactions",
|
||||
"model_access_group_list_transactions",
|
||||
|
|
@ -412,6 +414,10 @@ class RedisUpdateBuffer:
|
|||
Litellm_EntityType.ORGANIZATION,
|
||||
db_spend_update_transactions.get("org_list_transactions"),
|
||||
),
|
||||
(
|
||||
Litellm_EntityType.ORGANIZATION_MEMBER,
|
||||
db_spend_update_transactions.get("org_member_list_transactions"),
|
||||
),
|
||||
(
|
||||
Litellm_EntityType.TAG,
|
||||
db_spend_update_transactions.get("tag_list_transactions"),
|
||||
|
|
@ -876,6 +882,9 @@ class RedisUpdateBuffer:
|
|||
list_of_transactions, "team_member_list_transactions"
|
||||
),
|
||||
org_list_transactions=_merged_entity_transactions(list_of_transactions, "org_list_transactions"),
|
||||
org_member_list_transactions=_merged_entity_transactions(
|
||||
list_of_transactions, "org_member_list_transactions"
|
||||
),
|
||||
tag_list_transactions=_merged_entity_transactions(list_of_transactions, "tag_list_transactions"),
|
||||
agent_list_transactions=_merged_entity_transactions(list_of_transactions, "agent_list_transactions"),
|
||||
model_access_group_list_transactions=_merged_entity_transactions(
|
||||
|
|
|
|||
|
|
@ -137,6 +137,7 @@ class SpendUpdateQueue(BaseUpdateQueue):
|
|||
team_list_transactions={},
|
||||
team_member_list_transactions={},
|
||||
org_list_transactions={},
|
||||
org_member_list_transactions={},
|
||||
tag_list_transactions={},
|
||||
agent_list_transactions={},
|
||||
model_access_group_list_transactions={},
|
||||
|
|
@ -150,6 +151,7 @@ class SpendUpdateQueue(BaseUpdateQueue):
|
|||
Litellm_EntityType.TEAM: "team_list_transactions",
|
||||
Litellm_EntityType.TEAM_MEMBER: "team_member_list_transactions",
|
||||
Litellm_EntityType.ORGANIZATION: "org_list_transactions",
|
||||
Litellm_EntityType.ORGANIZATION_MEMBER: "org_member_list_transactions",
|
||||
Litellm_EntityType.TAG: "tag_list_transactions",
|
||||
Litellm_EntityType.AGENT: "agent_list_transactions",
|
||||
Litellm_EntityType.MODEL_ACCESS_GROUP: "model_access_group_list_transactions",
|
||||
|
|
@ -188,6 +190,8 @@ class SpendUpdateQueue(BaseUpdateQueue):
|
|||
transactions_dict = db_spend_update_transactions["team_member_list_transactions"]
|
||||
elif dict_key == "org_list_transactions":
|
||||
transactions_dict = db_spend_update_transactions["org_list_transactions"]
|
||||
elif dict_key == "org_member_list_transactions":
|
||||
transactions_dict = db_spend_update_transactions["org_member_list_transactions"]
|
||||
elif dict_key == "tag_list_transactions":
|
||||
transactions_dict = db_spend_update_transactions["tag_list_transactions"]
|
||||
elif dict_key == "agent_list_transactions":
|
||||
|
|
|
|||
|
|
@ -0,0 +1,63 @@
|
|||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import (
|
||||
AGENT_365_PROD_API_BASE,
|
||||
AGENT_365_PROD_RESOURCE_APP_ID,
|
||||
)
|
||||
|
||||
from .agent_365 import Agent365Guardrail
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail") -> Agent365Guardrail:
|
||||
import litellm
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
tenant_id: Final = litellm_params.tenant_id or get_secret_str("AGENT365_TENANT_ID")
|
||||
client_id: Final = litellm_params.client_id or get_secret_str("AGENT365_CLIENT_ID")
|
||||
client_secret: Final = (
|
||||
litellm_params.client_secret or litellm_params.api_key or get_secret_str("AGENT365_CLIENT_SECRET")
|
||||
)
|
||||
api_base: Final = litellm_params.api_base or get_secret_str("AGENT365_API_BASE")
|
||||
resource_app_id: Final = litellm_params.resource_app_id or get_secret_str("AGENT365_RESOURCE_APP_ID")
|
||||
|
||||
if not tenant_id:
|
||||
raise ValueError("Microsoft Agent 365: tenant_id is required")
|
||||
if not client_id:
|
||||
raise ValueError("Microsoft Agent 365: client_id is required")
|
||||
if not client_secret:
|
||||
raise ValueError(
|
||||
"Microsoft Agent 365: client secret is required. Set client_secret, api_key, or AGENT365_CLIENT_SECRET"
|
||||
)
|
||||
|
||||
guardrail_name: Final = guardrail.get("guardrail_name")
|
||||
if not guardrail_name:
|
||||
raise ValueError("Microsoft Agent 365: guardrail_name is required")
|
||||
|
||||
agent_365_guardrail: Final = Agent365Guardrail(
|
||||
guardrail_name=guardrail_name,
|
||||
tenant_id=tenant_id,
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
api_base=api_base or AGENT_365_PROD_API_BASE,
|
||||
resource_app_id=resource_app_id or AGENT_365_PROD_RESOURCE_APP_ID,
|
||||
agent_id=litellm_params.agent_id,
|
||||
request_timeout=litellm_params.timeout if litellm_params.timeout is not None else 10.0,
|
||||
unreachable_fallback=litellm_params.unreachable_fallback,
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(agent_365_guardrail)
|
||||
return agent_365_guardrail
|
||||
|
||||
|
||||
guardrail_initializer_registry: Final = { # mutable-ok: registry auto-discovery requires a dict instance
|
||||
SupportedGuardrailIntegrations.AGENT_365.value: initialize_guardrail,
|
||||
}
|
||||
|
||||
guardrail_class_registry: Final = { # mutable-ok: registry auto-discovery requires a dict instance
|
||||
SupportedGuardrailIntegrations.AGENT_365.value: Agent365Guardrail,
|
||||
}
|
||||
637
litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py
Normal file
637
litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py
Normal file
|
|
@ -0,0 +1,637 @@
|
|||
"""Microsoft Agent 365 governance guardrail for MCP tool calls.
|
||||
|
||||
Before the gateway executes an MCP tool, the pending call is sent to the
|
||||
Agent 365 tool-evaluation endpoint, where Microsoft Defender scores it and
|
||||
Agent 365 records it for observability. The returned allow/block verdict is
|
||||
enforced here. Authentication is the Entra On-Behalf-Of flow: the caller's
|
||||
incoming bearer token (audienced to this gateway's app registration) is
|
||||
exchanged for a delegated Agent 365 token, so Defender evaluates and audits
|
||||
as the signed-in user.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import Timeout as LitellmTimeout
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import (
|
||||
AGENT_365_PROD_API_BASE,
|
||||
AGENT_365_PROD_RESOURCE_APP_ID,
|
||||
AGENT_365_SCOPE_NAME,
|
||||
Agent365GuardrailConfigModel,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
from litellm.types.utils import GuardrailStatus
|
||||
|
||||
TOKEN_ENDPOINT_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token"
|
||||
EVALUATE_PATH: Final = "/agents/tool-evaluation/evaluate"
|
||||
MCP_SESSION_ID_HEADER: Final = "mcp-session-id"
|
||||
DEFENDER_STATUS_EVALUATED: Final = "Evaluated"
|
||||
_GATEWAY_OWNED_TOKEN_ERRORS: Final = frozenset(
|
||||
{"invalid_client", "unauthorized_client", "invalid_scope", "invalid_resource"}
|
||||
)
|
||||
# Entra reports a malformed or unverifiable assertion as ``invalid_client`` too; only its AADSTS50027xx
|
||||
# (InvalidJwtToken) sub-codes tell that apart from a bad gateway secret.
|
||||
_INVALID_ASSERTION_AADSTS_PREFIX: Final = "50027"
|
||||
_AADSTS_CODES_ADAPTER: Final = TypeAdapter(tuple[int, ...])
|
||||
_MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool")
|
||||
_OBO_CACHE_MAX_ENTRIES: Final = 1000
|
||||
_DEFAULT_TOKEN_TTL_SECONDS: Final = 3599.0
|
||||
_TOKEN_EXPIRY_SLACK_SECONDS: Final = 60.0
|
||||
|
||||
|
||||
def _parse_expires_in(raw: object) -> float:
|
||||
if not isinstance(raw, (int, float, str)):
|
||||
return _DEFAULT_TOKEN_TTL_SECONDS
|
||||
try:
|
||||
return float(raw)
|
||||
except ValueError:
|
||||
return _DEFAULT_TOKEN_TTL_SECONDS
|
||||
|
||||
|
||||
def _parse_aadsts_codes(raw: object) -> tuple[int, ...]:
|
||||
try:
|
||||
return _AADSTS_CODES_ADAPTER.validate_python(raw)
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def entra_assertion(value: object) -> str | None:
|
||||
"""``value`` when it is a compact JWS, the only bearer shape the OBO exchange accepts as its assertion.
|
||||
A LiteLLM virtual key, session bearer, or opaque upstream token in ``Authorization`` yields ``None``."""
|
||||
return value if isinstance(value, str) and value.count(".") == 2 else None
|
||||
|
||||
|
||||
class _DefenderResult(TypedDict, total=False):
|
||||
status: ReadOnly[str]
|
||||
verdict: ReadOnly[str | None]
|
||||
message: ReadOnly[str | None]
|
||||
|
||||
|
||||
class _EvaluateResponse(TypedDict, total=False):
|
||||
allowed: ReadOnly[bool]
|
||||
defender: ReadOnly[_DefenderResult]
|
||||
correlationId: ReadOnly[str]
|
||||
|
||||
|
||||
class _UnavailableDetail(TypedDict):
|
||||
error: ReadOnly[str]
|
||||
message: ReadOnly[str]
|
||||
tool: ReadOnly[str]
|
||||
|
||||
|
||||
class _BlockedDetail(TypedDict):
|
||||
error: ReadOnly[str]
|
||||
message: ReadOnly[str]
|
||||
tool: ReadOnly[str]
|
||||
correlation_id: ReadOnly[str | None]
|
||||
|
||||
|
||||
class Agent365TokenExchangeError(Exception):
|
||||
def __init__(self, status_code: int, error_code: str, description: str, aadsts_codes: tuple[int, ...] = ()) -> None:
|
||||
super().__init__(f"{error_code}: {description}")
|
||||
self.status_code = status_code
|
||||
self.error_code = error_code
|
||||
self.description = description
|
||||
self.aadsts_codes = aadsts_codes
|
||||
|
||||
@property
|
||||
def gateway_owned(self) -> bool:
|
||||
"""Whether the gateway's own client credentials, scope or resource were refused, as opposed to the
|
||||
caller's assertion. The caller cannot fix a gateway-owned rejection by signing in again."""
|
||||
if self.error_code not in _GATEWAY_OWNED_TOKEN_ERRORS:
|
||||
return False
|
||||
return not any(str(code).startswith(_INVALID_ASSERTION_AADSTS_PREFIX) for code in self.aadsts_codes)
|
||||
|
||||
|
||||
class Agent365MalformedResponseError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class Agent365ThrottledError(Exception):
|
||||
def __init__(self, status_code: int) -> None:
|
||||
super().__init__(f"HTTP {status_code}")
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
class Agent365Guardrail(CustomGuardrail):
|
||||
"""Pre-MCP-call guardrail enforcing Microsoft Agent 365 tool-evaluation verdicts.
|
||||
|
||||
Block-only: it never rewrites the call, so it runs in the post-sequential phase and judges the
|
||||
arguments the sequential guardrails hand upstream, whatever order the guardrails list uses."""
|
||||
|
||||
records_own_guardrail_information: ClassVar[bool] = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
guardrail_name: str,
|
||||
tenant_id: str,
|
||||
client_id: str,
|
||||
client_secret: str,
|
||||
api_base: str = AGENT_365_PROD_API_BASE,
|
||||
resource_app_id: str = AGENT_365_PROD_RESOURCE_APP_ID,
|
||||
agent_id: str | None = None,
|
||||
request_timeout: float = 10.0,
|
||||
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
|
||||
async_handler: AsyncHTTPHandler | None = None,
|
||||
**kwargs, # noqa: ANN003 # kwargs-ok: forwarded verbatim to CustomGuardrail (event_hook, default_on)
|
||||
) -> None:
|
||||
super().__init__(
|
||||
guardrail_name=guardrail_name,
|
||||
supported_event_hooks=self.get_supported_event_hooks(),
|
||||
run_in_parallel=True,
|
||||
**kwargs,
|
||||
)
|
||||
self.guardrail_provider = "agent_365"
|
||||
self.tenant_id = tenant_id
|
||||
self.client_id = client_id
|
||||
self.client_secret = client_secret
|
||||
self.api_base = api_base.rstrip("/")
|
||||
self.resource_app_id = resource_app_id
|
||||
self.agent_id = agent_id
|
||||
self.request_timeout = request_timeout
|
||||
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
|
||||
"fail_open" if unreachable_fallback == "fail_open" else "fail_closed"
|
||||
)
|
||||
self.async_handler = async_handler or get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.GuardrailCallback
|
||||
)
|
||||
self._obo_token_cache: OrderedDict[str, tuple[str, float]] = OrderedDict() # mutable-ok: lock-guarded LRU
|
||||
self._obo_cache_lock = threading.Lock()
|
||||
verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name)
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> "type[GuardrailConfigModel] | None":
|
||||
return Agent365GuardrailConfigModel
|
||||
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: CustomGuardrail contract
|
||||
return [GuardrailEventHooks.pre_mcp_call] # mutable-ok: CustomGuardrail contract expects a list
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
cache: "DualCache",
|
||||
data: dict, # mutable-ok: hook contract; guardrail logging appends into the request metadata in place
|
||||
call_type: str,
|
||||
) -> Exception | str | dict | None: # mutable-ok: CustomGuardrail.async_pre_call_hook contract
|
||||
if call_type not in _MCP_CALL_TYPES:
|
||||
return data
|
||||
if "mcp_tool_name" not in data:
|
||||
return data
|
||||
if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_mcp_call) is not True:
|
||||
return data
|
||||
|
||||
tool_name: Final = str(data.get("mcp_tool_name") or "")
|
||||
assertion: Final = entra_assertion(data.get("incoming_bearer_token"))
|
||||
if assertion is None:
|
||||
self._handle_caller_fault(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
status_code=401,
|
||||
reason=(
|
||||
"the caller did not present an Entra bearer token; the Agent 365 guardrail "
|
||||
"authorizes tool calls On-Behalf-Of the signed-in user"
|
||||
),
|
||||
)
|
||||
|
||||
try:
|
||||
obo_token: Final = await self._get_obo_token(assertion)
|
||||
except Agent365TokenExchangeError as exc:
|
||||
if exc.gateway_owned:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=(
|
||||
f"Entra rejected the gateway's own Agent 365 credentials ({exc.error_code}); "
|
||||
"check the guardrail's client_id, client_secret and resource_app_id"
|
||||
),
|
||||
)
|
||||
self._handle_caller_fault(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
status_code=401,
|
||||
reason=f"the Entra On-Behalf-Of token exchange was rejected ({exc.error_code})",
|
||||
)
|
||||
except Agent365ThrottledError as exc:
|
||||
self._handle_throttled(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=f"the Entra token endpoint returned HTTP {exc.status_code}",
|
||||
latency_ms=None,
|
||||
)
|
||||
except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=f"the Entra token endpoint could not be reached ({type(exc).__name__})",
|
||||
)
|
||||
except Agent365MalformedResponseError as exc:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=str(exc),
|
||||
)
|
||||
|
||||
start: Final = time.perf_counter()
|
||||
try:
|
||||
response: Final = await self._post_allowing_error_status(
|
||||
url=f"{self.api_base}{EVALUATE_PATH}",
|
||||
json=self._build_evaluate_payload(data=data, user_api_key_dict=user_api_key_dict),
|
||||
headers={"Authorization": f"Bearer {obo_token}"}, # mutable-ok: httpx header dict
|
||||
)
|
||||
except (httpx.HTTPError, LitellmTimeout, TimeoutError) as exc:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=f"the Agent 365 endpoint could not be reached ({type(exc).__name__})",
|
||||
)
|
||||
latency_ms: Final = (time.perf_counter() - start) * 1000.0
|
||||
fallback: Final = self._handle_evaluate_error(
|
||||
data=data, tool_name=tool_name, assertion=assertion, response=response, latency_ms=latency_ms
|
||||
)
|
||||
if fallback is not None:
|
||||
return fallback
|
||||
return self._enforce_verdict(data=data, tool_name=tool_name, response=response, latency_ms=latency_ms)
|
||||
|
||||
def _handle_evaluate_error(
|
||||
self,
|
||||
data: dict, # mutable-ok: guardrail logging appends into the request metadata in place
|
||||
tool_name: str,
|
||||
assertion: str,
|
||||
response: httpx.Response,
|
||||
latency_ms: float,
|
||||
) -> dict | None: # mutable-ok: returns the request data dict per hook contract on fail_open
|
||||
if response.status_code in (408, 429):
|
||||
self._handle_throttled(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=f"the Agent 365 endpoint returned HTTP {response.status_code}",
|
||||
latency_ms=latency_ms,
|
||||
)
|
||||
if 400 <= response.status_code < 500:
|
||||
if response.status_code == 401:
|
||||
self._evict_obo_token(assertion)
|
||||
self._record_verdict(
|
||||
data=data,
|
||||
verdict="Rejected",
|
||||
guardrail_status="guardrail_intervened",
|
||||
defender_status=None,
|
||||
correlation_id=None,
|
||||
latency_ms=latency_ms,
|
||||
reason=f"HTTP {response.status_code}: {response.text[:512]}",
|
||||
)
|
||||
rejected_detail: Final[_UnavailableDetail] = {
|
||||
"error": "Agent 365 rejected the tool evaluation request",
|
||||
"message": response.text[:512]
|
||||
if response.status_code == 400
|
||||
else f"the Agent 365 evaluation request failed with HTTP {response.status_code}",
|
||||
"tool": tool_name,
|
||||
}
|
||||
raise HTTPException(status_code=400, detail=rejected_detail)
|
||||
if response.status_code != 200:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=f"the Agent 365 endpoint returned HTTP {response.status_code}",
|
||||
)
|
||||
return None
|
||||
|
||||
def _enforce_verdict(
|
||||
self,
|
||||
data: dict, # mutable-ok: guardrail logging appends into the request metadata in place
|
||||
tool_name: str,
|
||||
response: httpx.Response,
|
||||
latency_ms: float,
|
||||
) -> dict: # mutable-ok: returns the request data dict per hook contract
|
||||
try:
|
||||
parsed_verdict: Final = response.json()
|
||||
except ValueError:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason="the Agent 365 endpoint returned a non-JSON body",
|
||||
)
|
||||
if not isinstance(parsed_verdict, dict):
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason="the Agent 365 endpoint returned a non-object JSON body",
|
||||
)
|
||||
verdict: Final[_EvaluateResponse] = parsed_verdict
|
||||
allowed: Final = verdict.get("allowed")
|
||||
if not isinstance(allowed, bool):
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason="the Agent 365 endpoint returned a verdict without a boolean 'allowed' field",
|
||||
)
|
||||
raw_defender: Final = verdict.get("defender")
|
||||
defender: Final = raw_defender if isinstance(raw_defender, dict) else _DefenderResult()
|
||||
raw_correlation_id: Final = verdict.get("correlationId")
|
||||
correlation_id: Final = raw_correlation_id if isinstance(raw_correlation_id, str) else None
|
||||
defender_status: Final = defender.get("status")
|
||||
if allowed and defender_status != DEFENDER_STATUS_EVALUATED:
|
||||
return self._handle_unavailable(
|
||||
data=data,
|
||||
tool_name=tool_name,
|
||||
reason=f"Microsoft Defender did not evaluate the call (defender.status={defender_status or 'missing'})",
|
||||
defender_status=defender_status,
|
||||
correlation_id=correlation_id,
|
||||
latency_ms=latency_ms,
|
||||
)
|
||||
self._record_verdict(
|
||||
data=data,
|
||||
verdict="Allow" if allowed else "Block",
|
||||
guardrail_status="success" if allowed else "guardrail_intervened",
|
||||
defender_status=defender_status,
|
||||
correlation_id=correlation_id,
|
||||
latency_ms=latency_ms,
|
||||
)
|
||||
if not allowed:
|
||||
blocked_detail: Final[_BlockedDetail] = {
|
||||
"error": "Blocked by Microsoft Defender",
|
||||
"message": (
|
||||
defender.get("message")
|
||||
or f"Invocation of '{tool_name}' is blocked by Microsoft Threat Detection policies "
|
||||
"configured by your administrator."
|
||||
),
|
||||
"tool": tool_name,
|
||||
"correlation_id": correlation_id,
|
||||
}
|
||||
raise HTTPException(status_code=400, detail=blocked_detail)
|
||||
return data
|
||||
|
||||
def _build_evaluate_payload(
|
||||
self,
|
||||
data: Mapping[str, object],
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
) -> dict[str, object]: # mutable-ok: JSON body for AsyncHTTPHandler.post, which requires dict
|
||||
tool_name: Final = str(data.get("mcp_tool_name") or "")
|
||||
arguments: Final = data.get("mcp_arguments")
|
||||
server_name: Final = str(data.get("mcp_server_name") or "litellm")
|
||||
agent_id: Final = self.agent_id or user_api_key_dict.key_alias
|
||||
payload: Final[dict[str, object]] = { # mutable-ok: JSON body with optional fields added below
|
||||
"tool": {"name": tool_name},
|
||||
"serverName": server_name,
|
||||
"conversationId": self._resolve_conversation_id(data),
|
||||
}
|
||||
if isinstance(arguments, dict):
|
||||
payload["arguments"] = arguments
|
||||
if agent_id:
|
||||
payload["agentId"] = str(agent_id)
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _resolve_conversation_id(data: Mapping[str, object]) -> str:
|
||||
"""The MCP session groups every tool call of one client conversation, so it is the conversation id
|
||||
when the transport carries one; stateless calls fall back to the per-call id."""
|
||||
raw_logging_obj: Final = data.get("litellm_logging_obj")
|
||||
logging_obj: Final = raw_logging_obj if isinstance(raw_logging_obj, LiteLLMLoggingObj) else None
|
||||
if logging_obj is not None:
|
||||
tool_call_metadata: Final = logging_obj.model_call_details.get("mcp_tool_call_metadata")
|
||||
session_from_logging: Final = (
|
||||
tool_call_metadata.get("mcp_session_id") if isinstance(tool_call_metadata, Mapping) else None
|
||||
)
|
||||
if isinstance(session_from_logging, str) and session_from_logging:
|
||||
return session_from_logging
|
||||
metadata: Final = next(
|
||||
(m for m in (data.get("metadata"), data.get("litellm_metadata")) if isinstance(m, Mapping)),
|
||||
None,
|
||||
)
|
||||
headers: Final = metadata.get("headers") if isinstance(metadata, Mapping) else None
|
||||
if isinstance(headers, Mapping):
|
||||
session_id: Final = next(
|
||||
(value for name, value in headers.items() if str(name).lower() == MCP_SESSION_ID_HEADER),
|
||||
None,
|
||||
)
|
||||
if isinstance(session_id, str) and session_id:
|
||||
return session_id
|
||||
call_id: Final = data.get("litellm_call_id") or (logging_obj.litellm_call_id if logging_obj else None)
|
||||
if isinstance(call_id, str) and call_id:
|
||||
return call_id
|
||||
return str(uuid.uuid4())
|
||||
|
||||
async def _get_obo_token(self, assertion: str) -> str:
|
||||
cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest()
|
||||
now: Final = time.time()
|
||||
with self._obo_cache_lock:
|
||||
cached: Final = self._obo_token_cache.get(cache_key)
|
||||
if cached and cached[1] > now + _TOKEN_EXPIRY_SLACK_SECONDS:
|
||||
self._obo_token_cache.move_to_end(cache_key)
|
||||
return cached[0]
|
||||
|
||||
response: Final = await self._post_allowing_error_status(
|
||||
url=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id),
|
||||
data={ # mutable-ok: OAuth form body; AsyncHTTPHandler.post requires dict
|
||||
"grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer",
|
||||
"client_id": self.client_id,
|
||||
"client_secret": self.client_secret,
|
||||
"assertion": assertion,
|
||||
"scope": f"{self.resource_app_id}/{AGENT_365_SCOPE_NAME}",
|
||||
"requested_token_use": "on_behalf_of",
|
||||
},
|
||||
headers={"Content-Type": "application/x-www-form-urlencoded"}, # mutable-ok: httpx header dict
|
||||
)
|
||||
if response.status_code in (408, 429):
|
||||
raise Agent365ThrottledError(status_code=response.status_code)
|
||||
if response.status_code >= 500:
|
||||
raise httpx.HTTPStatusError(
|
||||
f"Entra token endpoint returned {response.status_code}",
|
||||
request=response.request,
|
||||
response=response,
|
||||
)
|
||||
try:
|
||||
parsed_body: Final = response.json()
|
||||
except ValueError as exc:
|
||||
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-JSON body") from exc
|
||||
if not isinstance(parsed_body, dict):
|
||||
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-object JSON body")
|
||||
body: Final = parsed_body
|
||||
if response.status_code >= 400:
|
||||
raise Agent365TokenExchangeError(
|
||||
status_code=response.status_code,
|
||||
error_code=str(body.get("error", "invalid_grant")),
|
||||
description=str(body.get("error_description", ""))[:512],
|
||||
aadsts_codes=_parse_aadsts_codes(body.get("error_codes")),
|
||||
)
|
||||
if "access_token" not in body:
|
||||
raise Agent365MalformedResponseError("the Entra token endpoint returned no access_token")
|
||||
raw_access_token: Final = body.get("access_token")
|
||||
if not isinstance(raw_access_token, str) or not raw_access_token:
|
||||
raise Agent365MalformedResponseError("the Entra token endpoint returned a non-string access_token")
|
||||
access_token: Final = raw_access_token
|
||||
expires_at: Final = time.time() + _parse_expires_in(body.get("expires_in", 3599))
|
||||
with self._obo_cache_lock:
|
||||
self._obo_token_cache[cache_key] = (access_token, expires_at)
|
||||
self._obo_token_cache.move_to_end(cache_key)
|
||||
while len(self._obo_token_cache) > _OBO_CACHE_MAX_ENTRIES:
|
||||
self._obo_token_cache.popitem(last=False)
|
||||
return access_token
|
||||
|
||||
async def _post_allowing_error_status(
|
||||
self,
|
||||
url: str,
|
||||
headers: dict[str, str], # mutable-ok: AsyncHTTPHandler.post requires dict
|
||||
data: dict[str, str] | None = None, # mutable-ok: AsyncHTTPHandler.post requires dict
|
||||
json: dict[str, object] | None = None, # mutable-ok: AsyncHTTPHandler.post requires dict
|
||||
) -> httpx.Response:
|
||||
try:
|
||||
return await self.async_handler.post(
|
||||
url=url,
|
||||
data=data,
|
||||
json=json,
|
||||
headers=headers,
|
||||
timeout=self.request_timeout,
|
||||
)
|
||||
except httpx.HTTPStatusError as exc:
|
||||
return exc.response
|
||||
|
||||
def _handle_caller_fault(
|
||||
self,
|
||||
data: dict, # mutable-ok: guardrail logging appends into the request metadata in place
|
||||
tool_name: str,
|
||||
status_code: int,
|
||||
reason: str,
|
||||
) -> NoReturn:
|
||||
self._record_verdict(
|
||||
data=data,
|
||||
verdict="Rejected",
|
||||
guardrail_status="guardrail_intervened",
|
||||
defender_status=None,
|
||||
correlation_id=None,
|
||||
latency_ms=None,
|
||||
reason=reason,
|
||||
)
|
||||
caller_fault_detail: Final[_UnavailableDetail] = {
|
||||
"error": "Agent 365 guardrail rejected the tool call",
|
||||
"message": f"Tool call '{tool_name}' was blocked because {reason}.",
|
||||
"tool": tool_name,
|
||||
}
|
||||
raise HTTPException(status_code=status_code, detail=caller_fault_detail)
|
||||
|
||||
def _handle_throttled(
|
||||
self,
|
||||
data: dict, # mutable-ok: guardrail logging appends into the request metadata in place
|
||||
tool_name: str,
|
||||
reason: str,
|
||||
latency_ms: float | None,
|
||||
) -> NoReturn:
|
||||
self._record_verdict(
|
||||
data=data,
|
||||
verdict="Throttled",
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
defender_status=None,
|
||||
correlation_id=None,
|
||||
latency_ms=latency_ms,
|
||||
reason=reason,
|
||||
)
|
||||
throttled_detail: Final[_UnavailableDetail] = {
|
||||
"error": "Agent 365 guardrail could not authorize the tool call",
|
||||
"message": f"Tool call '{tool_name}' was blocked because {reason}; "
|
||||
"throttled evaluations block regardless of unreachable_fallback.",
|
||||
"tool": tool_name,
|
||||
}
|
||||
raise HTTPException(status_code=503, detail=throttled_detail)
|
||||
|
||||
def _evict_obo_token(self, assertion: str) -> None:
|
||||
cache_key: Final = hashlib.sha256(assertion.encode("utf-8")).hexdigest()
|
||||
with self._obo_cache_lock:
|
||||
self._obo_token_cache.pop(cache_key, None)
|
||||
|
||||
def _handle_unavailable(
|
||||
self,
|
||||
data: dict, # mutable-ok: guardrail logging appends into the request metadata in place
|
||||
tool_name: str,
|
||||
reason: str,
|
||||
defender_status: str | None = None,
|
||||
correlation_id: str | None = None,
|
||||
latency_ms: float | None = None,
|
||||
) -> dict: # mutable-ok: returns the request data dict per hook contract
|
||||
if self.unreachable_fallback == "fail_open":
|
||||
verbose_proxy_logger.warning(
|
||||
"Agent 365 guardrail (%s): %s; unreachable_fallback='fail_open', allowing tool call '%s' unscanned",
|
||||
self.guardrail_name,
|
||||
reason,
|
||||
tool_name,
|
||||
)
|
||||
self._record_verdict(
|
||||
data=data,
|
||||
verdict="Unscanned",
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
defender_status=defender_status,
|
||||
correlation_id=correlation_id,
|
||||
latency_ms=latency_ms,
|
||||
reason=reason,
|
||||
)
|
||||
return data
|
||||
self._record_verdict(
|
||||
data=data,
|
||||
verdict="Unavailable",
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
defender_status=defender_status,
|
||||
correlation_id=correlation_id,
|
||||
latency_ms=latency_ms,
|
||||
reason=reason,
|
||||
)
|
||||
unavailable_detail: Final[_UnavailableDetail] = {
|
||||
"error": "Agent 365 guardrail could not authorize the tool call",
|
||||
"message": f"Tool call '{tool_name}' was blocked because {reason} and unreachable_fallback is "
|
||||
"'fail_closed'.",
|
||||
"tool": tool_name,
|
||||
}
|
||||
raise HTTPException(status_code=503, detail=unavailable_detail)
|
||||
|
||||
def _record_verdict(
|
||||
self,
|
||||
data: dict[str, object], # mutable-ok: standard guardrail logging appends into the request metadata in place
|
||||
verdict: str,
|
||||
guardrail_status: "GuardrailStatus",
|
||||
defender_status: str | None,
|
||||
correlation_id: str | None,
|
||||
latency_ms: float | None,
|
||||
reason: str | None = None,
|
||||
) -> None:
|
||||
payload: Final[dict[str, object]] = {"verdict": verdict} # mutable-ok: optional fields added below
|
||||
if defender_status:
|
||||
payload["defender_status"] = defender_status
|
||||
if correlation_id:
|
||||
payload["correlation_id"] = correlation_id
|
||||
if latency_ms is not None:
|
||||
payload["latency_ms"] = round(latency_ms, 1)
|
||||
if reason:
|
||||
payload["reason"] = reason
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response=payload,
|
||||
request_data=data,
|
||||
guardrail_status=guardrail_status,
|
||||
duration=(latency_ms / 1000.0) if latency_ms is not None else None,
|
||||
guardrail_provider=self.guardrail_provider,
|
||||
event_type=GuardrailEventHooks.pre_mcp_call,
|
||||
)
|
||||
|
|
@ -58,6 +58,11 @@ if TYPE_CHECKING:
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
def _metadata_bucket(request_data: Mapping[str, object], key: str) -> Mapping[str, object]:
|
||||
bucket: Final = request_data.get(key)
|
||||
return bucket if isinstance(bucket, Mapping) else {}
|
||||
|
||||
|
||||
class CustomCodeGuardrailError(Exception):
|
||||
"""Raised when custom code guardrail execution fails."""
|
||||
|
||||
|
|
@ -280,12 +285,16 @@ class CustomCodeGuardrail(CustomGuardrail):
|
|||
Returns:
|
||||
Safe subset of request data
|
||||
"""
|
||||
metadata: Final = {
|
||||
**_metadata_bucket(request_data, "metadata"),
|
||||
**_metadata_bucket(request_data, "litellm_metadata"),
|
||||
}
|
||||
return {
|
||||
"model": request_data.get("model"),
|
||||
"user_id": request_data.get("user_api_key_user_id"),
|
||||
"team_id": request_data.get("user_api_key_team_id"),
|
||||
"end_user_id": request_data.get("user_api_key_end_user_id"),
|
||||
"metadata": request_data.get("metadata", {}),
|
||||
"user_id": metadata.get("user_api_key_user_id"),
|
||||
"team_id": metadata.get("user_api_key_team_id"),
|
||||
"end_user_id": metadata.get("user_api_key_end_user_id"),
|
||||
"metadata": metadata,
|
||||
}
|
||||
|
||||
def _process_result(
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.proxy.auth.auth_utils import (
|
||||
ESTIMATED_OUTPUT_TOKENS_FIELD,
|
||||
get_estimated_output_tokens,
|
||||
get_key_own_model_rate_limit,
|
||||
get_key_tag_rpm_limit,
|
||||
get_model_rate_limit_from_metadata,
|
||||
)
|
||||
|
|
@ -2892,41 +2893,67 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
return batch_limiter
|
||||
return None
|
||||
|
||||
def _key_owns_model_limit(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
requested_model: str,
|
||||
rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"],
|
||||
) -> bool:
|
||||
key_own_limits: Final = get_key_own_model_rate_limit(user_api_key_dict, rate_limit_key)
|
||||
return key_own_limits is not None and key_own_limits.get(requested_model) is not None
|
||||
|
||||
def _inherited_team_model_limit(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
requested_model: str,
|
||||
rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"],
|
||||
) -> int | None:
|
||||
team_limits: Final = get_model_rate_limit_from_metadata(user_api_key_dict, "team_metadata", rate_limit_key)
|
||||
team_limit: Final = team_limits.get(requested_model) if team_limits else None
|
||||
if team_limit is None:
|
||||
return None
|
||||
if self._key_owns_model_limit(user_api_key_dict, requested_model, rate_limit_key):
|
||||
return None
|
||||
return team_limit
|
||||
|
||||
def _key_owns_model_tpm_limit_from_request_metadata(
|
||||
self,
|
||||
request_metadata: Mapping[str, object],
|
||||
model_group: str | None,
|
||||
) -> bool:
|
||||
if model_group is None:
|
||||
return False
|
||||
key_view: Final = UserAPIKeyAuth.model_validate(
|
||||
{
|
||||
"metadata": request_metadata.get("user_api_key_metadata") or {},
|
||||
"model_max_budget": request_metadata.get("user_api_key_model_max_budget") or {},
|
||||
}
|
||||
)
|
||||
return self._key_owns_model_limit(key_view, model_group, "model_tpm_limit")
|
||||
|
||||
def _add_team_model_rate_limit_descriptor_from_metadata(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
requested_model: str | None,
|
||||
descriptors: list[RateLimitDescriptor],
|
||||
) -> None:
|
||||
"""Add team model rate limit descriptor from team_metadata if applicable."""
|
||||
if (
|
||||
get_model_rate_limit_from_metadata(user_api_key_dict, "team_metadata", "model_rpm_limit") is not None
|
||||
or get_model_rate_limit_from_metadata(user_api_key_dict, "team_metadata", "model_tpm_limit") is not None
|
||||
):
|
||||
_tpm_limit_for_team_model: Final = (
|
||||
get_model_rate_limit_from_metadata(user_api_key_dict, "team_metadata", "model_tpm_limit") or {}
|
||||
if requested_model is None:
|
||||
return
|
||||
team_rpm_limit: Final = self._inherited_team_model_limit(user_api_key_dict, requested_model, "model_rpm_limit")
|
||||
team_tpm_limit: Final = self._inherited_team_model_limit(user_api_key_dict, requested_model, "model_tpm_limit")
|
||||
if team_rpm_limit is None and team_tpm_limit is None:
|
||||
return
|
||||
descriptors.append(
|
||||
RateLimitDescriptor(
|
||||
key="model_per_team",
|
||||
value=f"{user_api_key_dict.team_id}:{requested_model}",
|
||||
rate_limit={
|
||||
"requests_per_unit": team_rpm_limit,
|
||||
"tokens_per_unit": team_tpm_limit,
|
||||
"window_size": self.window_size,
|
||||
},
|
||||
)
|
||||
_rpm_limit_for_team_model: Final = (
|
||||
get_model_rate_limit_from_metadata(user_api_key_dict, "team_metadata", "model_rpm_limit") or {}
|
||||
)
|
||||
should_check_rate_limit: Final = (
|
||||
requested_model in _tpm_limit_for_team_model or requested_model in _rpm_limit_for_team_model
|
||||
)
|
||||
|
||||
if should_check_rate_limit and requested_model is not None:
|
||||
model_specific_tpm_limit: Final = _tpm_limit_for_team_model.get(requested_model)
|
||||
model_specific_rpm_limit: Final = _rpm_limit_for_team_model.get(requested_model)
|
||||
descriptors.append(
|
||||
RateLimitDescriptor(
|
||||
key="model_per_team",
|
||||
value=f"{user_api_key_dict.team_id}:{requested_model}",
|
||||
rate_limit={
|
||||
"requests_per_unit": model_specific_rpm_limit,
|
||||
"tokens_per_unit": model_specific_tpm_limit,
|
||||
"window_size": self.window_size,
|
||||
},
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
def _add_project_model_rate_limit_descriptor_from_metadata(
|
||||
self,
|
||||
|
|
@ -4459,6 +4486,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
kwargs=kwargs,
|
||||
model_group=reconcile_model,
|
||||
)
|
||||
charged_targets: Final = (
|
||||
[target for target in targets if target[0] != "model_per_team"]
|
||||
if self._key_owns_model_tpm_limit_from_request_metadata(request_metadata, reconcile_model)
|
||||
else targets
|
||||
)
|
||||
if reserved_tokens > 0 and total_tokens < reserved_tokens:
|
||||
verbose_proxy_logger.debug(
|
||||
"Releasing unused TPM budget on success: reserved=%s, actual=%s, release=%s",
|
||||
|
|
@ -4468,7 +4500,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
)
|
||||
pipeline_operations.extend(
|
||||
self._build_reservation_aware_tpm_ops(
|
||||
targets=targets,
|
||||
targets=charged_targets,
|
||||
reserved_scopes=reserved_scopes,
|
||||
actual_tokens=total_tokens,
|
||||
reserved_tokens=reserved_tokens,
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
|
|||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
from litellm.llms.azure.passthrough.transformation import foreign_azure_deployment
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
|
||||
from litellm.proxy._types import *
|
||||
|
|
@ -1550,6 +1551,26 @@ async def _relay_azure_router_model(
|
|||
"put the model group name in the deployments segment"
|
||||
}
|
||||
raise HTTPException(status_code=400, detail=rejection)
|
||||
return await _relay_router_model(
|
||||
llm_router=llm_router,
|
||||
model=model,
|
||||
endpoint=endpoint,
|
||||
request=request,
|
||||
request_body=request_body,
|
||||
is_streaming_request=is_streaming_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
|
||||
async def _relay_router_model(
|
||||
llm_router: litellm.Router,
|
||||
model: str,
|
||||
endpoint: str,
|
||||
request: Request,
|
||||
request_body: Mapping[str, object],
|
||||
is_streaming_request: bool,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Response:
|
||||
try:
|
||||
result: Final = await llm_router.allm_passthrough_route(
|
||||
model=model,
|
||||
|
|
@ -1599,6 +1620,65 @@ async def _relay_azure_router_model(
|
|||
)
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/nvidia_nim/{endpoint:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
|
||||
tags=["NVIDIA NIM Pass-through", "pass-through"],
|
||||
)
|
||||
async def nvidia_nim_proxy_route(
|
||||
endpoint: str,
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
):
|
||||
"""
|
||||
Relay a native NVIDIA NIM request through a LiteLLM model group.
|
||||
|
||||
`{PROXY_BASE_URL}/nvidia_nim/{model_group}/v1/infer` forwards the body unchanged to the deployment's
|
||||
`api_base`, so object detection and OCR NIMs whose payload carries no `model` field still go through
|
||||
virtual key auth, model access checks, and spend logging.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
return await relay_nvidia_nim_request(
|
||||
llm_router=llm_router,
|
||||
endpoint=endpoint,
|
||||
request=request,
|
||||
request_body=await get_request_body(request),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
|
||||
async def relay_nvidia_nim_request(
|
||||
llm_router: litellm.Router | None,
|
||||
endpoint: str,
|
||||
request: Request,
|
||||
request_body: Mapping[str, object],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Response:
|
||||
model_group: Final = nvidia_nim_model_group_in_path(endpoint, llm_router.get_model_list()) if llm_router else None
|
||||
if llm_router is None or model_group is None:
|
||||
rejection: Final[RelayRejection] = {
|
||||
"error": "no NVIDIA NIM model group in the path; call /nvidia_nim/{model_group}/v1/infer with a model "
|
||||
"group from your `model_list` whose deployments all use `nvidia_nim/` models"
|
||||
}
|
||||
raise HTTPException(status_code=400, detail=rejection)
|
||||
|
||||
is_streaming_request: Final = is_passthrough_request_streaming(request_body)
|
||||
return await open_sse_before_first_byte(
|
||||
_relay_router_model(
|
||||
llm_router=llm_router,
|
||||
model=model_group,
|
||||
endpoint=endpoint,
|
||||
request=request,
|
||||
request_body=request_body,
|
||||
is_streaming_request=is_streaming_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
),
|
||||
ping_interval_seconds=(litellm.sse_keepalive_ping_interval_seconds if is_streaming_request else None),
|
||||
)
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/azure_ai/{endpoint:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
|
||||
|
|
|
|||
|
|
@ -12,6 +12,9 @@ from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
|||
ModelResponseIterator as GeminiModelResponseIterator,
|
||||
)
|
||||
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import (
|
||||
VertexPassthroughLoggingHandler,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
ModelResponse,
|
||||
TextCompletionResponse,
|
||||
|
|
@ -40,6 +43,17 @@ class GeminiPassthroughLoggingHandler:
|
|||
request_body: dict,
|
||||
**kwargs,
|
||||
) -> PassThroughEndpointLoggingTypedDict:
|
||||
if VertexPassthroughLoggingHandler.is_interactions_route(url_route):
|
||||
return VertexPassthroughLoggingHandler.interactions_passthrough_handler(
|
||||
httpx_response=httpx_response,
|
||||
request_body=request_body,
|
||||
logging_obj=logging_obj,
|
||||
kwargs=kwargs,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
custom_llm_provider="gemini",
|
||||
vertex_location=None,
|
||||
)
|
||||
if "predictLongRunning" in url_route:
|
||||
model = GeminiPassthroughLoggingHandler.extract_model_from_url(url_route)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,15 +1,20 @@
|
|||
import asyncio
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import VERTEX_BATCH_PREDICTION_JOBS_ROUTE
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import (
|
||||
InteractionsUsageObjectTransformation,
|
||||
)
|
||||
from litellm.llms.vertex_ai.common_utils import (
|
||||
get_vertex_ai_lyria_generation_cost,
|
||||
get_vertex_location_from_url,
|
||||
|
|
@ -49,8 +54,73 @@ else:
|
|||
|
||||
EndpointType = Any
|
||||
|
||||
_VERTEX_INTERACTIONS_PATH: Final = re.compile(r"/projects/[^/]+/locations/[^/]+/interactions/?$")
|
||||
_INTERACTIONS_RESPONSE_BODY: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def _interactions_model(
|
||||
response_body: Mapping[str, object],
|
||||
request_body: Mapping[str, object] | None,
|
||||
) -> str | None:
|
||||
response_model: Final = response_body.get("model")
|
||||
if isinstance(response_model, str) and response_model:
|
||||
return response_model
|
||||
request_model: Final = (request_body or {}).get("model")
|
||||
if isinstance(request_model, str) and request_model:
|
||||
return request_model
|
||||
return None
|
||||
|
||||
|
||||
class VertexPassthroughLoggingHandler:
|
||||
@staticmethod
|
||||
def is_interactions_route(url_route: str) -> bool:
|
||||
return urlparse(url_route).path.rstrip("/").endswith("/interactions")
|
||||
|
||||
@staticmethod
|
||||
def is_vertex_interactions_route(url_route: str) -> bool:
|
||||
return _VERTEX_INTERACTIONS_PATH.search(urlparse(url_route).path) is not None
|
||||
|
||||
@staticmethod
|
||||
def interactions_passthrough_handler(
|
||||
httpx_response: httpx.Response,
|
||||
request_body: Mapping[str, object] | None,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
kwargs: dict[str, object],
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
custom_llm_provider: Literal["vertex_ai", "gemini"],
|
||||
vertex_location: str | None,
|
||||
) -> PassThroughEndpointLoggingTypedDict:
|
||||
response_body: Final = _INTERACTIONS_RESPONSE_BODY.validate_python(httpx_response.json())
|
||||
usage_object: Final = response_body.get("usage")
|
||||
model: Final = _interactions_model(response_body, request_body)
|
||||
if model is None or not InteractionsUsageObjectTransformation.is_interactions_usage_object(usage_object):
|
||||
return {"result": None, "kwargs": kwargs}
|
||||
|
||||
litellm_model_response: Final = ModelResponse(
|
||||
model=model,
|
||||
usage=InteractionsUsageObjectTransformation.transform_interactions_usage_object(
|
||||
cast(Mapping[str, Any], usage_object)
|
||||
),
|
||||
)
|
||||
logging_obj.custom_llm_provider = custom_llm_provider
|
||||
logging_kwargs: Final = (
|
||||
VertexPassthroughLoggingHandler._create_vertex_response_logging_payload_for_generate_content(
|
||||
litellm_model_response=litellm_model_response,
|
||||
model=model,
|
||||
kwargs=kwargs,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
)
|
||||
return {
|
||||
"result": litellm_model_response,
|
||||
"kwargs": {**logging_kwargs, "custom_llm_provider": custom_llm_provider},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def vertex_passthrough_handler(
|
||||
httpx_response: httpx.Response,
|
||||
|
|
@ -66,6 +136,17 @@ class VertexPassthroughLoggingHandler:
|
|||
vertex_location: Final = get_vertex_location_from_url(url_route)
|
||||
if vertex_location is not None:
|
||||
logging_obj.optional_params["vertex_location"] = vertex_location
|
||||
if VertexPassthroughLoggingHandler.is_interactions_route(url_route):
|
||||
return VertexPassthroughLoggingHandler.interactions_passthrough_handler(
|
||||
httpx_response=httpx_response,
|
||||
request_body=request_body,
|
||||
logging_obj=logging_obj,
|
||||
kwargs=kwargs,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
custom_llm_provider="vertex_ai",
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
if "predictLongRunning" in url_route:
|
||||
model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
|
||||
|
||||
|
|
|
|||
|
|
@ -361,7 +361,9 @@ class PassThroughEndpointLogging:
|
|||
def is_vertex_route(self, url_route: str) -> bool:
|
||||
if any(f":{method}" in url_route for method in self.TRACKED_VERTEX_METHOD_ROUTES):
|
||||
return True
|
||||
return any(resource in url_route for resource in self.TRACKED_VERTEX_RESOURCE_ROUTES)
|
||||
if any(resource in url_route for resource in self.TRACKED_VERTEX_RESOURCE_ROUTES):
|
||||
return True
|
||||
return VertexPassthroughLoggingHandler.is_vertex_interactions_route(url_route)
|
||||
|
||||
def is_anthropic_route(self, url_route: str):
|
||||
for route in self.TRACKED_ANTHROPIC_ROUTES:
|
||||
|
|
@ -434,8 +436,12 @@ class PassThroughEndpointLogging:
|
|||
|
||||
def is_gemini_route(self, url_route: str, custom_llm_provider: str | None = None):
|
||||
"""Check if the URL route is a Gemini API route."""
|
||||
if custom_llm_provider != "gemini":
|
||||
return False
|
||||
if VertexPassthroughLoggingHandler.is_interactions_route(url_route):
|
||||
return True
|
||||
for route in self.TRACKED_GEMINI_ROUTES:
|
||||
if route in url_route and custom_llm_provider == "gemini":
|
||||
if route in url_route:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
|
|
|||
|
|
@ -261,7 +261,7 @@ class ProxyInitializationHelpers:
|
|||
import uvicorn
|
||||
|
||||
import litellm
|
||||
from litellm._logging import _get_uvicorn_json_log_config
|
||||
from litellm._logging import _get_uvicorn_json_log_config, resolve_log_level
|
||||
|
||||
uvicorn_args: Final = {
|
||||
"app": "litellm.proxy.proxy_server:app",
|
||||
|
|
@ -275,6 +275,8 @@ class ProxyInitializationHelpers:
|
|||
elif litellm.json_logs:
|
||||
# Use JSON log config for uvicorn to ensure all logs (including exceptions) are JSON
|
||||
uvicorn_args["log_config"] = _get_uvicorn_json_log_config()
|
||||
elif litellm_log := os.environ.get("LITELLM_LOG"):
|
||||
uvicorn_args["log_level"] = resolve_log_level(litellm_log)
|
||||
if keepalive_timeout is not None:
|
||||
uvicorn_args["timeout_keep_alive"] = keepalive_timeout
|
||||
if timeout_worker_healthcheck is not None:
|
||||
|
|
|
|||
|
|
@ -176,6 +176,7 @@ class _SessionSpendRow(TypedDict):
|
|||
api_key: ReadOnly[str]
|
||||
session_total_count: ReadOnly[int]
|
||||
session_total_spend: float
|
||||
session_total_duration_ms: ReadOnly[int]
|
||||
mcp_tool_call_count: int
|
||||
mcp_tool_call_spend: float
|
||||
session_cache_hit_count: ReadOnly[int]
|
||||
|
|
@ -194,6 +195,7 @@ _SESSION_MODEL_NAME_MAX_LEN: Final = 256
|
|||
class _SessionSpendStats(NamedTuple):
|
||||
session_total_count: int
|
||||
session_total_spend: float
|
||||
session_total_duration_ms: int
|
||||
mcp_tool_call_count: int
|
||||
mcp_tool_call_spend: float
|
||||
session_cache_hit_count: int
|
||||
|
|
@ -4543,6 +4545,12 @@ async def _build_ui_spend_logs_response(
|
|||
SELECT session_id, api_key,
|
||||
COUNT(*)::int AS session_total_count,
|
||||
COALESCE(SUM(spend), 0)::double precision AS session_total_spend,
|
||||
COALESCE(SUM(
|
||||
COALESCE(
|
||||
request_duration_ms,
|
||||
(EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER
|
||||
)
|
||||
), 0)::bigint AS session_total_duration_ms,
|
||||
COUNT(*) FILTER (
|
||||
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
|
||||
)::int AS mcp_tool_call_count,
|
||||
|
|
@ -4584,6 +4592,7 @@ async def _build_ui_spend_logs_response(
|
|||
(row["session_id"], row["api_key"]): _SessionSpendStats(
|
||||
session_total_count=int(row.get("session_total_count") or 0),
|
||||
session_total_spend=float(row.get("session_total_spend") or 0.0),
|
||||
session_total_duration_ms=int(row.get("session_total_duration_ms") or 0),
|
||||
mcp_tool_call_count=int(row.get("mcp_tool_call_count") or 0),
|
||||
mcp_tool_call_spend=float(row.get("mcp_tool_call_spend") or 0.0),
|
||||
session_cache_hit_count=int(row.get("session_cache_hit_count") or 0),
|
||||
|
|
@ -4615,6 +4624,7 @@ async def _build_ui_spend_logs_response(
|
|||
row_dict["session_total_count"] = session_stats.session_total_count if session_stats else 1
|
||||
if session_stats:
|
||||
row_dict["session_total_spend"] = session_stats.session_total_spend
|
||||
row_dict["session_total_duration_ms"] = session_stats.session_total_duration_ms
|
||||
if session_stats.mcp_tool_call_count:
|
||||
row_dict["mcp_tool_call_count"] = session_stats.mcp_tool_call_count
|
||||
row_dict["mcp_tool_call_spend"] = session_stats.mcp_tool_call_spend
|
||||
|
|
|
|||
|
|
@ -1269,7 +1269,12 @@ class ProxyLogging:
|
|||
# (e.g. MCPJWTSigner) to independently verify the caller's identity
|
||||
# before re-signing an outbound token (FR-5 verify+re-sign).
|
||||
"incoming_bearer_token": kwargs.get("incoming_bearer_token"),
|
||||
"metadata": {"headers": kwargs.get("headers") or {}},
|
||||
"metadata": {
|
||||
"headers": kwargs.get("headers") or {},
|
||||
"user_api_key_user_id": kwargs.get("user_api_key_user_id"),
|
||||
"user_api_key_team_id": kwargs.get("user_api_key_team_id"),
|
||||
"user_api_key_end_user_id": kwargs.get("user_api_key_end_user_id"),
|
||||
},
|
||||
}
|
||||
user_api_key_auth: Final = kwargs.get("user_api_key_auth")
|
||||
if isinstance(user_api_key_auth, UserAPIKeyAuth):
|
||||
|
|
|
|||
|
|
@ -218,9 +218,8 @@ class GatedAutoRouterCapability:
|
|||
stored ``litellm_params`` (``{config}`` is the caller's expression for the normalized
|
||||
``complexity_router_config`` jsonb, substituted as many times as the predicate needs); they live
|
||||
on one record so they cannot drift apart. ``subject`` and ``remedy`` build the shared refusal
|
||||
message. A validated config claims at most one capability, and the validator is what makes that
|
||||
true: tier_definitions rejects every heuristic classifier_type, and it also rejects the
|
||||
classifier system_prompt, which in turn only applies to the classifier types heuristic_v2 is not.
|
||||
message. A validated config claims at most one capability: gated classifier types cannot be
|
||||
combined with operator-defined tiers or classifier prompts.
|
||||
"""
|
||||
|
||||
key: str
|
||||
|
|
@ -238,6 +237,22 @@ HEURISTIC_V2_CAPABILITY: Final = GatedAutoRouterCapability(
|
|||
sql_config_predicate="{config} ->> 'classifier_type' = 'heuristic_v2'",
|
||||
)
|
||||
|
||||
CAPABILITY_CLASSIFIER_CAPABILITY: Final = GatedAutoRouterCapability(
|
||||
key="capability",
|
||||
subject="with classifier_type 'capability' (Capability)",
|
||||
remedy="Use a different classifier or remove an existing Capability router.",
|
||||
uses=lambda config: _mapping(config).get("classifier_type") == "capability",
|
||||
sql_config_predicate="{config} ->> 'classifier_type' = 'capability'",
|
||||
)
|
||||
|
||||
LLM_V2_CAPABILITY: Final = GatedAutoRouterCapability(
|
||||
key="llm_v2",
|
||||
subject="with classifier_type 'llm_v2' (Fuse v2)",
|
||||
remedy="Use a different classifier or remove an existing Fuse v2 router.",
|
||||
uses=lambda config: _mapping(config).get("classifier_type") == "llm_v2",
|
||||
sql_config_predicate="{config} ->> 'classifier_type' = 'llm_v2'",
|
||||
)
|
||||
|
||||
_OPERATOR_PROMPT_FIELDS_SQL: Final = " OR ".join(
|
||||
f"{{config}} ->> '{field}' IS NOT NULL" for field in OPERATOR_CLASSIFIER_PROMPT_FIELDS
|
||||
)
|
||||
|
|
@ -258,7 +273,12 @@ CUSTOMIZATION_CAPABILITY: Final = GatedAutoRouterCapability(
|
|||
),
|
||||
)
|
||||
|
||||
GATED_AUTO_ROUTER_CAPABILITIES: Final = (HEURISTIC_V2_CAPABILITY, CUSTOMIZATION_CAPABILITY)
|
||||
GATED_AUTO_ROUTER_CAPABILITIES: Final = (
|
||||
HEURISTIC_V2_CAPABILITY,
|
||||
CAPABILITY_CLASSIFIER_CAPABILITY,
|
||||
LLM_V2_CAPABILITY,
|
||||
CUSTOMIZATION_CAPABILITY,
|
||||
)
|
||||
|
||||
|
||||
def claimed_capability(complexity_router_config: object) -> GatedAutoRouterCapability | None:
|
||||
|
|
|
|||
|
|
@ -8,6 +8,9 @@ from pydantic import BaseModel, ConfigDict, Field, field_validator, model_valida
|
|||
from typing_extensions import Required, TypedDict
|
||||
|
||||
from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import (
|
||||
Agent365GuardrailConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.akto import (
|
||||
AktoConfigModel,
|
||||
)
|
||||
|
|
@ -137,6 +140,7 @@ class SupportedGuardrailIntegrations(Enum):
|
|||
COMPRESR = "compresr"
|
||||
STRAIKER = "straiker"
|
||||
ALICE = "alice"
|
||||
AGENT_365 = "agent_365"
|
||||
CONDUCT = "conduct"
|
||||
|
||||
|
||||
|
|
@ -1045,7 +1049,7 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up
|
|||
default="fail_closed",
|
||||
description=(
|
||||
"Behavior when a guardrail endpoint is unreachable due to network errors. "
|
||||
"Implemented by guardrail='generic_guardrail_api', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. "
|
||||
"Implemented by guardrail='generic_guardrail_api', 'agent_365', 'akto', 'vigil_guard', 'repelloai', 'headroom', and 'compresr'. "
|
||||
"'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed."
|
||||
),
|
||||
)
|
||||
|
|
@ -1183,6 +1187,7 @@ class LitellmParams( # pyright: ignore[reportIncompatibleVariableOverride] # o
|
|||
QostodianNexusConfigModel,
|
||||
VigilGuardGuardrailConfigModel,
|
||||
SingulrGuardrailConfigModel,
|
||||
Agent365GuardrailConfigModel,
|
||||
):
|
||||
guardrail: str = Field(description="The type of guardrail integration to use")
|
||||
mode: str | list[str] | Mode = Field(
|
||||
|
|
|
|||
66
litellm/types/proxy/guardrails/guardrail_hooks/agent_365.py
Normal file
66
litellm/types/proxy/guardrails/guardrail_hooks/agent_365.py
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
from typing import Final
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from .base import GuardrailConfigModel
|
||||
|
||||
AGENT_365_PROD_API_BASE: Final = "https://agent365.svc.cloud.microsoft"
|
||||
AGENT_365_PROD_RESOURCE_APP_ID: Final = "ea9ffc3e-8a23-4a7d-836d-234d7c7565c1"
|
||||
AGENT_365_SCOPE_NAME: Final = "ThreatProtection.Evaluate.All"
|
||||
|
||||
|
||||
class Agent365GuardrailConfigModel(GuardrailConfigModel):
|
||||
tenant_id: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Entra tenant id used for the On-Behalf-Of token exchange. "
|
||||
"Falls back to the AGENT365_TENANT_ID environment variable."
|
||||
),
|
||||
)
|
||||
|
||||
client_id: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Client id of the gateway's Entra app registration (a confidential client). "
|
||||
"Falls back to the AGENT365_CLIENT_ID environment variable."
|
||||
),
|
||||
)
|
||||
|
||||
client_secret: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Client secret of the gateway's Entra app registration, used to perform the "
|
||||
"On-Behalf-Of exchange. Falls back to the AGENT365_CLIENT_SECRET environment variable."
|
||||
),
|
||||
)
|
||||
|
||||
api_base: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Base URL of the Microsoft Agent 365 tool-evaluation endpoint. "
|
||||
f"Defaults to the production endpoint {AGENT_365_PROD_API_BASE}. "
|
||||
"Falls back to the AGENT365_API_BASE environment variable."
|
||||
),
|
||||
)
|
||||
|
||||
resource_app_id: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Application id of the Agent 365 resource the OBO token is minted for. "
|
||||
f"Defaults to the production resource {AGENT_365_PROD_RESOURCE_APP_ID}; "
|
||||
"the Test and PreProd environments use a different id. "
|
||||
"Falls back to the AGENT365_RESOURCE_APP_ID environment variable."
|
||||
),
|
||||
)
|
||||
|
||||
agent_id: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Agent identity reported to Agent 365 with every tool evaluation. "
|
||||
"When unset, the caller's key alias is used."
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "Microsoft Agent 365"
|
||||
|
|
@ -4275,6 +4275,7 @@ class LiteLLMRealtimeStreamLoggingObject(LiteLLMPydanticObjectBase):
|
|||
# rate_limits.updated), blocks the event loop, and discards the session usage.
|
||||
results: SkipValidation[OpenAIRealtimeStreamList]
|
||||
usage: Usage
|
||||
service_tier: str | None = None
|
||||
_hidden_params: dict = {}
|
||||
|
||||
@field_serializer("results")
|
||||
|
|
|
|||
|
|
@ -8998,6 +8998,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return WatsonxPassthroughConfig()
|
||||
elif LlmProviders.NVIDIA_NIM == provider:
|
||||
from litellm.llms.nvidia_nim.passthrough.transformation import (
|
||||
NvidiaNimPassthroughConfig,
|
||||
)
|
||||
|
||||
return NvidiaNimPassthroughConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -58151,6 +58151,23 @@
|
|||
"model_info": {
|
||||
"supports_reasoning": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "gemini-chat-baseline",
|
||||
"pattern": "gemini-(?!.*(?:-tts|-image|-live|-audio|-embedding|-computer-use|-robotics|-transcribe|-translate))(?:2[.-][5-9]|[3-9](?:[.-]\\d{1,2})?)-(?:pro|flash)(?:-lite)?(?![a-z])",
|
||||
"description": "Any Gemini text-chat id at 2.5 or higher under any namespace, including bare ids, gemini/, vertex_ai/, openrouter/google/, deepinfra/google/, vercel_ai_gateway/google/, oci/google., and databricks-gemini-<major>-<minor>: gemini-<major>[.minor]-(pro|flash)[-lite] with any trailing preview, date or variant tag. The capability flags were verified against each of those providers' own catalogs and docs. The lookahead excludes the tts, image, live, audio, embedding, computer-use, robotics, transcribe and translate lines, which are different modes with different capabilities. Provider-specific deviations, such as Perplexity's Agent API serving these as mode responses, are carried by their exact map entries, which always win over this rule. Carries no token limits or pricing, so those stay on the standard unmapped behavior rather than a guessed number. Source check 2026-09-15: all 45 first-party 2.5+ text-chat entries in this map carry every field below, and the OpenRouter (openrouter.ai/api/v1/models), Vercel AI Gateway (ai-gateway.vercel.sh/v1/models), DeepInfra (api.deepinfra.com/models/list), OCI and Databricks model docs list reasoning, tools and image input for the same models.",
|
||||
"model_info": {
|
||||
"mode": "chat",
|
||||
"supports_reasoning": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_vision": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_web_search": true
|
||||
}
|
||||
}
|
||||
]
|
||||
},
|
||||
|
|
|
|||
|
|
@ -67,6 +67,7 @@ IGNORE_FUNCTIONS = [
|
|||
"_redact_agent_params_tree", # max depth set (default 10), same shape as _redact_sensitive_litellm_params.
|
||||
"_restore_redacted_nested_value", # max depth set (default 10), mirrors _redact_agent_params_tree on the write side.
|
||||
"_unqualified", # bounded by the qualifier depth of a static TypedDict annotation (Annotated, Required/NotRequired, ReadOnly around one type, no cycles possible).
|
||||
"completion_cost", # max depth 1: recursion only fires for mixed-tier Responses WS logging objects, and each split part carries a single service_tier so _split_responses_ws_logging_object_by_service_tier returns None.
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@
|
|||
-- Idempotent: deletes all e2e-* rows then re-inserts deterministic data.
|
||||
|
||||
-- 1. Clean up in dependency order
|
||||
DELETE FROM "LiteLLM_InvitationLink"
|
||||
WHERE "user_id" LIKE 'e2e-%' OR "created_by" LIKE 'e2e-%' OR "updated_by" LIKE 'e2e-%';
|
||||
DELETE FROM "LiteLLM_TeamMembership" WHERE "user_id" LIKE 'e2e-%';
|
||||
DELETE FROM "LiteLLM_VerificationToken" WHERE token LIKE 'e2e-%';
|
||||
DELETE FROM "LiteLLM_TeamTable" WHERE "team_id" LIKE 'e2e-%';
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import { chromium, expect, request } from "@playwright/test";
|
||||
import { users, Role, STORAGE_PATHS } from "./fixtures/users";
|
||||
import { ARTIFACT_DIR, UI_BASE_URL } from "./constants";
|
||||
import { expectUnrestrictedDashboard, setInvitedUserPassword } from "./helpers/userOnboarding";
|
||||
import * as fs from "fs";
|
||||
import * as path from "path";
|
||||
|
||||
|
|
@ -30,32 +31,37 @@ async function globalSetup() {
|
|||
throw new Error(`Enabling enable_projects_ui failed (${settingsRes.status()}): ${await settingsRes.text()}`);
|
||||
}
|
||||
|
||||
for (const { email, password, seedApiRole } of Object.values(users)) {
|
||||
if (!seedApiRole) {
|
||||
continue;
|
||||
}
|
||||
const createRes = await api.post(`${UI_BASE_URL}${rootPath}/user/new`, {
|
||||
headers: { Authorization: `Bearer ${masterKey}` },
|
||||
data: { user_email: email, user_role: seedApiRole, auto_create_key: false },
|
||||
});
|
||||
if (!createRes.ok() && createRes.status() !== 409) {
|
||||
throw new Error(`Seeding user ${email} failed (${createRes.status()}): ${await createRes.text()}`);
|
||||
}
|
||||
const passwordRes = await api.post(`${UI_BASE_URL}${rootPath}/user/update`, {
|
||||
headers: { Authorization: `Bearer ${masterKey}` },
|
||||
data: { user_email: email, password },
|
||||
});
|
||||
if (!passwordRes.ok()) {
|
||||
throw new Error(`Setting password for ${email} failed (${passwordRes.status()}): ${await passwordRes.text()}`);
|
||||
}
|
||||
}
|
||||
await api.dispose();
|
||||
|
||||
for (const role of Object.values(Role)) {
|
||||
const { email, password } = users[role];
|
||||
const roles = [Role.ProxyAdmin, ...Object.values(Role).filter((role) => role !== Role.ProxyAdmin)];
|
||||
for (const role of roles) {
|
||||
const { email, password, seedApiRole } = users[role];
|
||||
const storagePath = STORAGE_PATHS[role];
|
||||
const page = await browser.newPage();
|
||||
try {
|
||||
if (seedApiRole) {
|
||||
const createRes = await api.post(`${UI_BASE_URL}${rootPath}/user/new`, {
|
||||
headers: { Authorization: `Bearer ${masterKey}` },
|
||||
data: { user_email: email, user_role: seedApiRole, auto_create_key: false },
|
||||
});
|
||||
if (!createRes.ok() && createRes.status() !== 409) {
|
||||
throw new Error(`Seeding user ${email} failed (${createRes.status()}): ${await createRes.text()}`);
|
||||
}
|
||||
const userId = createRes.ok()
|
||||
? (await createRes.json()).user_id
|
||||
: await (async () => {
|
||||
const existing = await api.get(`${UI_BASE_URL}${rootPath}/user/list`, {
|
||||
headers: { Authorization: `Bearer ${masterKey}` },
|
||||
params: { user_email: email },
|
||||
});
|
||||
expect(existing.ok(), `Find seeded user ${email}: HTTP ${existing.status()}`).toBe(true);
|
||||
const matches = (await existing.json()).users.filter(
|
||||
(user: { user_email: string }) => user.user_email === email,
|
||||
);
|
||||
expect(matches, `Exactly one seeded user for ${email}`).toHaveLength(1);
|
||||
return matches[0].user_id;
|
||||
})();
|
||||
expect(typeof userId, `User ID for ${email}`).toBe("string");
|
||||
await setInvitedUserPassword(api, userId, password);
|
||||
}
|
||||
await page.goto(`${UI_BASE_URL}${rootPath}/ui/login`);
|
||||
await page.getByPlaceholder("Enter your username").fill(email);
|
||||
await page.getByPlaceholder("Enter your password").fill(password);
|
||||
|
|
@ -63,7 +69,7 @@ async function globalSetup() {
|
|||
await page.waitForURL((url) => url.pathname.startsWith(`${rootPath}/ui`) && !url.pathname.includes("/login"), {
|
||||
timeout: 30_000,
|
||||
});
|
||||
await expect(page.locator("a", { hasText: "Virtual Keys" })).toBeVisible({ timeout: 30_000 });
|
||||
await expectUnrestrictedDashboard(page);
|
||||
// Dismiss feedback popup if present
|
||||
const dismiss = page.getByText("Don't ask me again");
|
||||
if (await dismiss.isVisible({ timeout: 1_500 }).catch(() => false)) {
|
||||
|
|
@ -100,6 +106,7 @@ async function globalSetup() {
|
|||
}
|
||||
}
|
||||
|
||||
await api.dispose();
|
||||
await browser.close();
|
||||
}
|
||||
|
||||
|
|
|
|||
59
tests/e2e/ui/helpers/userOnboarding.ts
Normal file
59
tests/e2e/ui/helpers/userOnboarding.ts
Normal file
|
|
@ -0,0 +1,59 @@
|
|||
import { expect, type APIRequestContext, type Page } from "@playwright/test";
|
||||
import { UI_BASE_URL } from "../constants";
|
||||
import { masterKey, rootPath } from "./traffic";
|
||||
|
||||
const endpoint = (route: string): string => `${UI_BASE_URL}${rootPath()}${route}`;
|
||||
|
||||
export async function setInvitedUserPassword(
|
||||
request: APIRequestContext,
|
||||
userId: string,
|
||||
password: string,
|
||||
): Promise<void> {
|
||||
const invitation = await request.post(endpoint("/invitation/new"), {
|
||||
headers: { Authorization: `Bearer ${masterKey()}` },
|
||||
data: { user_id: userId },
|
||||
});
|
||||
expect(invitation.ok(), `Create invitation for ${userId}: HTTP ${invitation.status()}`).toBe(true);
|
||||
const { id } = await invitation.json();
|
||||
expect(typeof id, "invitation ID").toBe("string");
|
||||
|
||||
const onboarding = await request.get(endpoint("/onboarding/get_token"), {
|
||||
params: { invite_link: id },
|
||||
});
|
||||
expect(onboarding.ok(), `Get onboarding session for ${userId}: HTTP ${onboarding.status()}`).toBe(true);
|
||||
const { token } = await onboarding.json();
|
||||
const payload = JSON.parse(Buffer.from(token.split(".")[1], "base64url").toString("utf-8"));
|
||||
expect(typeof payload.key, "onboarding credential").toBe("string");
|
||||
const claimed = await request.post(endpoint("/onboarding/claim_token"), {
|
||||
headers: { Authorization: `Bearer ${payload.key}` },
|
||||
data: { invitation_link: id, user_id: userId, password },
|
||||
});
|
||||
expect(claimed.ok(), `Claim invitation for ${userId}: HTTP ${claimed.status()}`).toBe(true);
|
||||
}
|
||||
|
||||
export async function readDashboardSession(page: Page): Promise<{
|
||||
key: string;
|
||||
user_id: string;
|
||||
password_reset_required?: boolean;
|
||||
}> {
|
||||
await expect.poll(async () => (await page.context().cookies()).some((cookie) => cookie.name === "token")).toBe(true);
|
||||
const cookie = (await page.context().cookies()).find((candidate) => candidate.name === "token")!;
|
||||
return JSON.parse(Buffer.from(cookie.value.split(".")[1], "base64url").toString("utf-8"));
|
||||
}
|
||||
|
||||
export async function expectUnrestrictedDashboard(page: Page): Promise<void> {
|
||||
const virtualKeys = page.getByRole("complementary").getByRole("link", { name: "Virtual Keys", exact: true });
|
||||
await expect(virtualKeys).toBeVisible({ timeout: 30_000 });
|
||||
const session = await readDashboardSession(page);
|
||||
expect(session.password_reset_required === true, "login must not require a password reset").toBe(false);
|
||||
await virtualKeys.click();
|
||||
await expect(page.getByRole("main").getByRole("heading", { name: "Virtual Keys", exact: true })).toBeVisible({
|
||||
timeout: 30_000,
|
||||
});
|
||||
const info = await page.request.get(endpoint("/user/info"), {
|
||||
headers: { Authorization: `Bearer ${session.key}` },
|
||||
params: { user_id: session.user_id },
|
||||
});
|
||||
expect(info.ok(), `Read own user with dashboard session: HTTP ${info.status()}`).toBe(true);
|
||||
expect((await info.json()).user_id).toBe(session.user_id);
|
||||
}
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
import { expectUnrestrictedDashboard, setInvitedUserPassword } from "../../helpers/userOnboarding";
|
||||
import { test, expect, type APIRequestContext } from "@playwright/test";
|
||||
import { Page } from "../../fixtures/pages";
|
||||
import {
|
||||
|
|
@ -76,10 +77,7 @@ test.describe("Internal User - own team key model scope", () => {
|
|||
user_role: "internal_user",
|
||||
auto_create_key: false,
|
||||
});
|
||||
await postAsMaster(request, "/user/update", {
|
||||
user_id: userId,
|
||||
password: MEMBER_PASSWORD,
|
||||
});
|
||||
await setInvitedUserPassword(request, userId, MEMBER_PASSWORD);
|
||||
await postAsMaster(request, "/team/member_add", {
|
||||
team_id: teamId,
|
||||
member: { role: "user", user_id: userId },
|
||||
|
|
@ -99,10 +97,7 @@ test.describe("Internal User - own team key model scope", () => {
|
|||
.getByPlaceholder("Enter your password")
|
||||
.fill(MEMBER_PASSWORD);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
await expect(
|
||||
page.locator("a", { hasText: "Virtual Keys" }),
|
||||
`${email} never reached the dashboard`,
|
||||
).toBeVisible({ timeout: 30_000 });
|
||||
await expectUnrestrictedDashboard(page);
|
||||
await dismissFeedbackPopup(page);
|
||||
|
||||
await navigateToPage(page, Page.ApiKeys);
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import { expectUnrestrictedDashboard, setInvitedUserPassword } from "../../helpers/userOnboarding";
|
||||
import { test, expect } from "@playwright/test";
|
||||
import { ADMIN_STORAGE_PATH } from "../../constants";
|
||||
import { Page } from "../../fixtures/pages";
|
||||
|
|
@ -46,19 +47,13 @@ test.describe("Second proxy admin", () => {
|
|||
|
||||
const userId = await inviteAdminUser();
|
||||
try {
|
||||
const passwordRes = await request.post("/user/update", {
|
||||
headers: auth,
|
||||
data: { user_email: email, password },
|
||||
});
|
||||
expect(passwordRes.ok(), `setting password failed (${passwordRes.status()}): ${await passwordRes.text()}`).toBe(
|
||||
true,
|
||||
);
|
||||
await setInvitedUserPassword(request, userId, password);
|
||||
|
||||
await page.goto("/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill(email);
|
||||
await page.getByPlaceholder("Enter your password").fill(password);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
await expect(page.locator("a", { hasText: "Virtual Keys" })).toBeVisible({ timeout: 30_000 });
|
||||
await expectUnrestrictedDashboard(page);
|
||||
await dismissFeedbackPopup(page);
|
||||
|
||||
await navigateToPage(page, Page.ApiKeys);
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import { expectUnrestrictedDashboard, setInvitedUserPassword } from "../../helpers/userOnboarding";
|
||||
import { test, expect, type Browser, type BrowserContext, type Page as PlaywrightPage } from "@playwright/test";
|
||||
import { Page } from "../../fixtures/pages";
|
||||
import { navigateToPage, dismissFeedbackPopup, clickTeamId } from "../../helpers/navigation";
|
||||
|
|
@ -24,7 +25,7 @@ async function signIn(browser: Browser, email: string): Promise<BrowserContext>
|
|||
await page.getByPlaceholder("Enter your username").fill(email);
|
||||
await page.getByPlaceholder("Enter your password").fill(PASSWORD);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
await expect(page.locator("a", { hasText: "Virtual Keys" })).toBeVisible({ timeout: 30_000 });
|
||||
await expectUnrestrictedDashboard(page);
|
||||
await dismissFeedbackPopup(page);
|
||||
return context;
|
||||
}
|
||||
|
|
@ -49,11 +50,7 @@ test.describe("Team Admin - Member permissions", () => {
|
|||
data: { user_id: userId, user_email: email, user_role: "internal_user", auto_create_key: false },
|
||||
});
|
||||
expect(created.ok(), `POST /user/new for ${userId} (${created.status()}): ${await created.text()}`).toBe(true);
|
||||
const password = await request.post("/user/update", {
|
||||
headers: auth(),
|
||||
data: { user_id: userId, password: PASSWORD },
|
||||
});
|
||||
expect(password.ok(), `POST /user/update for ${userId} (${password.status()})`).toBe(true);
|
||||
await setInvitedUserPassword(request, userId, PASSWORD);
|
||||
};
|
||||
|
||||
let teamId = "";
|
||||
|
|
|
|||
|
|
@ -0,0 +1,37 @@
|
|||
import importlib
|
||||
import logging
|
||||
from collections.abc import Iterator
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm_proxy_extras._logging as extras_logging
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fresh_extras_logger() -> Iterator[logging.Logger]:
|
||||
logger = logging.getLogger("litellm_proxy_extras")
|
||||
saved_handlers = logger.handlers[:]
|
||||
saved_level = logger.level
|
||||
logger.handlers[:] = []
|
||||
try:
|
||||
yield logger
|
||||
finally:
|
||||
logger.handlers[:] = saved_handlers
|
||||
logger.setLevel(saved_level)
|
||||
|
||||
|
||||
def test_litellm_log_error_silences_extras_info_lines(monkeypatch, fresh_extras_logger):
|
||||
monkeypatch.setenv("LITELLM_LOG", "ERROR")
|
||||
reloaded = importlib.reload(extras_logging).logger
|
||||
assert reloaded is fresh_extras_logger
|
||||
assert reloaded.isEnabledFor(logging.INFO) is False
|
||||
assert reloaded.isEnabledFor(logging.ERROR) is True
|
||||
|
||||
|
||||
@pytest.mark.parametrize("litellm_log", [None, "info", "DEBUG"])
|
||||
def test_unset_or_verbose_litellm_log_keeps_extras_info_lines(monkeypatch, fresh_extras_logger, litellm_log):
|
||||
if litellm_log is None:
|
||||
monkeypatch.delenv("LITELLM_LOG", raising=False)
|
||||
else:
|
||||
monkeypatch.setenv("LITELLM_LOG", litellm_log)
|
||||
assert importlib.reload(extras_logging).logger.isEnabledFor(logging.INFO) is True
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -135,9 +135,7 @@ def test_fill_missing_requires_per_rule_opt_in(restore_generalizations):
|
|||
"supports_vision": True,
|
||||
}
|
||||
|
||||
restore_generalizations(
|
||||
[{"name": "base", "pattern": r"^acme-", "model_info": {"supports_reasoning": True}}]
|
||||
)
|
||||
restore_generalizations([{"name": "base", "pattern": r"^acme-", "model_info": {"supports_reasoning": True}}])
|
||||
assert match_fill_missing_generalizations("acme-1", "openai") is None
|
||||
|
||||
restore_generalizations(
|
||||
|
|
@ -451,6 +449,94 @@ def shipped_cost_map(monkeypatch):
|
|||
set_fallback_generalizations(previous_rules)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,provider",
|
||||
[
|
||||
("gemini-4-pro", "gemini"),
|
||||
("gemini/gemini-4-pro", None),
|
||||
("gemini-3.9-flash-lite-preview-09-2026", "vertex_ai"),
|
||||
("vertex_ai/gemini-4-pro", None),
|
||||
("gemini-4-pro-preview-customtools", "gemini"),
|
||||
("google/gemini-4-pro", "openrouter"),
|
||||
("google/gemini-4-pro", "deepinfra"),
|
||||
("google/gemini-4-pro", "vercel_ai_gateway"),
|
||||
("google.gemini-4-pro", "oci"),
|
||||
("databricks-gemini-4-1-pro", "databricks"),
|
||||
],
|
||||
)
|
||||
def test_shipped_gemini_chat_baseline_resolves_unmapped_ids(shipped_cost_map, model, provider):
|
||||
assert model not in litellm.model_cost
|
||||
if provider == "gemini":
|
||||
assert f"gemini/{model}" not in litellm.model_cost
|
||||
elif provider in {"openrouter", "deepinfra", "vercel_ai_gateway", "oci", "databricks"}:
|
||||
assert f"{provider}/{model}" not in litellm.model_cost
|
||||
|
||||
info = litellm.get_model_info(model, custom_llm_provider=provider)
|
||||
assert info["litellm_provider"] == (provider or model.split("/")[0])
|
||||
assert info["mode"] == "chat"
|
||||
assert not info.get("max_input_tokens")
|
||||
assert info["supports_reasoning"] is True
|
||||
assert info["supports_function_calling"] is True
|
||||
assert info["supports_tool_choice"] is True
|
||||
assert info["supports_system_messages"] is True
|
||||
assert info["supports_vision"] is True
|
||||
assert info["supports_response_schema"] is True
|
||||
assert info["supports_pdf_input"] is True
|
||||
assert info["supports_prompt_caching"] is True
|
||||
assert info["supports_web_search"] is True
|
||||
assert not info.get("input_cost_per_token")
|
||||
assert not info.get("output_cost_per_token")
|
||||
|
||||
|
||||
def test_shipped_gemini_chat_baseline_loses_to_perplexity_exact_entries(shipped_cost_map):
|
||||
info = litellm.get_model_info("google/gemini-2.5-pro", custom_llm_provider="perplexity")
|
||||
entry = litellm.model_cost["perplexity/google/gemini-2.5-pro"]
|
||||
assert info["mode"] == "responses"
|
||||
assert entry["supports_reasoning"] is False
|
||||
|
||||
|
||||
def test_shipped_gemini_chat_baseline_skips_non_chat_and_pre_2_5_ids(shipped_cost_map):
|
||||
for model in (
|
||||
"gemini/gemini-4-flash-image",
|
||||
"gemini/gemini-3.9-flash-preview-tts",
|
||||
"gemini/gemini-4-flash-live-preview",
|
||||
"gemini/gemini-4-flash-native-audio",
|
||||
"gemini/gemini-embedding-4",
|
||||
"gemini/gemini-2.5-computer-use-preview-12-2026",
|
||||
"gemini/gemini-2.0-flash-new",
|
||||
"gemini/gemini-1.5-pro-new",
|
||||
"gemini/gemini-4-flashy",
|
||||
"gemini/gemini-4-flash-transcribe",
|
||||
"gemini/gemini-4-flash-live-translate-preview",
|
||||
"databricks-gemini-3-1-flash-image",
|
||||
"openrouter/google/gemini-2.0-flash-001",
|
||||
):
|
||||
assert match_capability_generalizations(model) is None, model
|
||||
|
||||
|
||||
def test_shipped_gemini_chat_baseline_keeps_reasoning_effort_on_unmapped_model(shipped_cost_map):
|
||||
assert litellm.supports_reasoning(model="gemini-4-pro", custom_llm_provider="gemini") is True
|
||||
|
||||
optional_params = litellm.utils.get_optional_params(
|
||||
model="gemini-4-pro",
|
||||
custom_llm_provider="gemini",
|
||||
reasoning_effort="medium",
|
||||
drop_params=False,
|
||||
)
|
||||
assert isinstance(optional_params, dict)
|
||||
assert optional_params["thinkingConfig"]["thinkingBudget"] > 0
|
||||
assert optional_params["thinkingConfig"]["includeThoughts"] is True
|
||||
|
||||
|
||||
def test_shipped_gemini_chat_baseline_loses_to_exact_entries(shipped_cost_map):
|
||||
model = "gemini-2.5-flash-lite"
|
||||
info = litellm.get_model_info(model, custom_llm_provider="gemini")
|
||||
entry = litellm.model_cost["gemini/gemini-2.5-flash-lite"]
|
||||
assert info["max_tokens"] == entry["max_tokens"]
|
||||
assert info["input_cost_per_token"] == entry["input_cost_per_token"]
|
||||
assert entry["input_cost_per_token"] > 0
|
||||
|
||||
|
||||
def test_shipped_bare_claude_id_routes_to_anthropic(shipped_cost_map):
|
||||
_, provider, _, _ = litellm.get_llm_provider(model="claude-haiku-4-6")
|
||||
assert provider == "anthropic"
|
||||
|
|
|
|||
|
|
@ -1,7 +1,5 @@
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
from litellm.litellm_core_utils.get_supported_openai_params import (
|
||||
get_supported_openai_params,
|
||||
)
|
||||
|
|
@ -33,9 +31,7 @@ def test_base_model_label_alone_lacks_bedrock_tools():
|
|||
"""The label by itself does not advertise tools; this is what made the union
|
||||
necessary. Guards against the discrepancy disappearing (and the regression test
|
||||
above silently passing for the wrong reason)."""
|
||||
params = get_supported_openai_params(
|
||||
model=BEDROCK_LABEL, custom_llm_provider="bedrock"
|
||||
)
|
||||
params = get_supported_openai_params(model=BEDROCK_LABEL, custom_llm_provider="bedrock")
|
||||
|
||||
assert params is not None
|
||||
assert "tools" not in params
|
||||
|
|
@ -46,14 +42,8 @@ def test_base_model_is_additive_not_replacement():
|
|||
|
||||
Bedrock: real id supports ``tools`` but not the label's reasoning hint; the union
|
||||
must contain the real model's ``tools`` regardless of the label being a subset."""
|
||||
real_only = set(
|
||||
get_supported_openai_params(
|
||||
model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock"
|
||||
)
|
||||
)
|
||||
label_only = set(
|
||||
get_supported_openai_params(model=BEDROCK_LABEL, custom_llm_provider="bedrock")
|
||||
)
|
||||
real_only = set(get_supported_openai_params(model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock"))
|
||||
label_only = set(get_supported_openai_params(model=BEDROCK_LABEL, custom_llm_provider="bedrock"))
|
||||
combined = set(
|
||||
get_supported_openai_params(
|
||||
model=BEDROCK_REAL_MODEL,
|
||||
|
|
@ -70,19 +60,15 @@ def test_base_model_is_additive_not_replacement():
|
|||
def test_base_model_adds_capabilities_the_real_model_lacks():
|
||||
"""Regression for #27717 (the behavior the union must preserve).
|
||||
|
||||
``gemini-3.1-pro`` isn't in the cost map so it advertises no reasoning support,
|
||||
``gemini-exp-9999`` isn't in the cost map so it advertises no reasoning support,
|
||||
but the registered ``gemini-3.1-pro-preview`` base_model does. The hint must add
|
||||
``reasoning_effort``/``thinking`` without the call erroring."""
|
||||
real_only = set(
|
||||
get_supported_openai_params(
|
||||
model="gemini-3.1-pro", custom_llm_provider="gemini"
|
||||
)
|
||||
)
|
||||
real_only = set(get_supported_openai_params(model="gemini-exp-9999", custom_llm_provider="gemini"))
|
||||
assert "reasoning_effort" not in real_only
|
||||
|
||||
combined = set(
|
||||
get_supported_openai_params(
|
||||
model="gemini-3.1-pro",
|
||||
model="gemini-exp-9999",
|
||||
custom_llm_provider="gemini",
|
||||
base_model="gemini-3.1-pro-preview",
|
||||
)
|
||||
|
|
@ -93,21 +79,15 @@ def test_base_model_adds_capabilities_the_real_model_lacks():
|
|||
|
||||
def test_no_base_model_is_unchanged():
|
||||
"""Omitting ``base_model`` must resolve purely from ``model``."""
|
||||
with_none = get_supported_openai_params(
|
||||
model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock", base_model=None
|
||||
)
|
||||
plain = get_supported_openai_params(
|
||||
model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock"
|
||||
)
|
||||
with_none = get_supported_openai_params(model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock", base_model=None)
|
||||
plain = get_supported_openai_params(model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock")
|
||||
|
||||
assert with_none == plain
|
||||
|
||||
|
||||
def test_base_model_equal_to_model_is_unchanged():
|
||||
"""A ``base_model`` identical to ``model`` must not double-resolve or reorder."""
|
||||
plain = get_supported_openai_params(
|
||||
model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock"
|
||||
)
|
||||
plain = get_supported_openai_params(model=BEDROCK_REAL_MODEL, custom_llm_provider="bedrock")
|
||||
same = get_supported_openai_params(
|
||||
model=BEDROCK_REAL_MODEL,
|
||||
custom_llm_provider="bedrock",
|
||||
|
|
@ -152,14 +132,10 @@ def test_bedrock_converse_alias_resolves_like_bedrock():
|
|||
params saw no Bedrock capabilities for a Converse model invoked via the alias."""
|
||||
anthropic_model = "bedrock/converse/us.anthropic.claude-sonnet-4-6"
|
||||
|
||||
via_alias = get_supported_openai_params(
|
||||
model=anthropic_model, custom_llm_provider="bedrock_converse"
|
||||
)
|
||||
via_alias = get_supported_openai_params(model=anthropic_model, custom_llm_provider="bedrock_converse")
|
||||
|
||||
assert via_alias is not None
|
||||
assert via_alias == get_supported_openai_params(
|
||||
model=anthropic_model, custom_llm_provider="bedrock"
|
||||
)
|
||||
assert via_alias == get_supported_openai_params(model=anthropic_model, custom_llm_provider="bedrock")
|
||||
assert "web_search_options" not in via_alias
|
||||
assert "tools" in via_alias
|
||||
|
||||
|
|
@ -167,9 +143,7 @@ def test_bedrock_converse_alias_resolves_like_bedrock():
|
|||
def test_bedrock_converse_alias_keeps_nova_web_search_options():
|
||||
"""Nova on the ``bedrock_converse`` alias still advertises web_search_options, proving the
|
||||
alias routes through the model-aware config rather than a blanket Bedrock default."""
|
||||
nova_params = get_supported_openai_params(
|
||||
model="amazon.nova-pro-v1:0", custom_llm_provider="bedrock_converse"
|
||||
)
|
||||
nova_params = get_supported_openai_params(model="amazon.nova-pro-v1:0", custom_llm_provider="bedrock_converse")
|
||||
|
||||
assert nova_params is not None
|
||||
assert "web_search_options" in nova_params
|
||||
|
|
|
|||
|
|
@ -6554,9 +6554,9 @@ async def test_prompt_hook_injection_marker_recorded_for_every_surface(logging_o
|
|||
assert pre_choice["metadata"]["litellm_gateway_injected_cache"] == ""
|
||||
|
||||
|
||||
def _responses_ws_logging_obj() -> LitellmLogging:
|
||||
def _responses_ws_logging_obj(model: str = "gpt-4o") -> LitellmLogging:
|
||||
return LitellmLogging(
|
||||
model="gpt-4o",
|
||||
model=model,
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type=CallTypes.aresponses_websocket.value,
|
||||
|
|
@ -6638,6 +6638,62 @@ def test_normalize_logging_result_bills_incomplete_responses_websocket_turns():
|
|||
assert normalized.usage.total_tokens == 75
|
||||
|
||||
|
||||
def test_normalize_logging_result_prices_responses_websocket_at_returned_service_tier():
|
||||
"""Issue #41299: a WebSocket turn billed at priority tier reported it on
|
||||
response.completed.response.service_tier, but the logging object dropped it and the
|
||||
session was priced at the default tier."""
|
||||
events = [
|
||||
{"type": "response.created", "response": {}},
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"service_tier": "priority",
|
||||
"usage": {"input_tokens": 100, "output_tokens": 40, "total_tokens": 140},
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
normalized = _responses_ws_logging_obj(model="gpt-5.4").normalize_logging_result(result=events)
|
||||
|
||||
assert isinstance(normalized, LiteLLMRealtimeStreamLoggingObject)
|
||||
assert normalized.service_tier == "priority"
|
||||
|
||||
usage = ResponseAPIUsage(input_tokens=100, output_tokens=40, total_tokens=140)
|
||||
ws_cost = litellm.completion_cost(
|
||||
completion_response=normalized,
|
||||
model="gpt-5.4",
|
||||
call_type=CallTypes.aresponses_websocket.value,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
priority_http_cost = litellm.completion_cost(
|
||||
completion_response=ResponsesAPIResponse(
|
||||
id="resp-priority",
|
||||
created_at=1700000000,
|
||||
output=[],
|
||||
service_tier="priority",
|
||||
usage=usage,
|
||||
),
|
||||
model="gpt-5.4",
|
||||
call_type=CallTypes.aresponses.value,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
default_http_cost = litellm.completion_cost(
|
||||
completion_response=ResponsesAPIResponse(
|
||||
id="resp-default",
|
||||
created_at=1700000000,
|
||||
output=[],
|
||||
service_tier="default",
|
||||
usage=usage,
|
||||
),
|
||||
model="gpt-5.4",
|
||||
call_type=CallTypes.aresponses.value,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
assert ws_cost == priority_http_cost
|
||||
assert priority_http_cost > default_http_cost
|
||||
|
||||
|
||||
def test_get_standard_logging_object_payload_reads_overhead_from_logging_obj_for_dict_results(logging_obj):
|
||||
"""LIT-5466: /v1/messages returns a plain dict with no _hidden_params, so the overhead
|
||||
recorded on the logging object must reach hidden_params.litellm_overhead_time_ms (SpendLogs)."""
|
||||
|
|
|
|||
|
|
@ -4,18 +4,13 @@ from typing import NamedTuple
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
from litellm.llms.bedrock.common_utils import BedrockModelInfo
|
||||
from litellm.utils import _get_model_info_helper
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
Message,
|
||||
ModelResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -31,8 +26,7 @@ def local_model_cost_map(monkeypatch):
|
|||
litellm.bedrock_converse_models.update(
|
||||
key
|
||||
for key, value in litellm.model_cost.items()
|
||||
if isinstance(value, dict)
|
||||
and value.get("litellm_provider") == "bedrock_converse"
|
||||
if isinstance(value, dict) and value.get("litellm_provider") == "bedrock_converse"
|
||||
)
|
||||
yield
|
||||
finally:
|
||||
|
|
@ -56,45 +50,69 @@ class GptProfile(NamedTuple):
|
|||
GPT_5_6_PROFILES = [
|
||||
GptProfile(
|
||||
model_id="us.openai.gpt-5.6-sol",
|
||||
input_cost=4.4e-06, input_cost_above_272k=8.8e-06,
|
||||
cache_write=5.5e-06, cache_write_above_272k=1.1e-05,
|
||||
cache_read=4.4e-07, cache_read_above_272k=8.8e-07,
|
||||
output_cost=2.2e-05, output_cost_above_272k=3.3e-05,
|
||||
input_cost=4.4e-06,
|
||||
input_cost_above_272k=8.8e-06,
|
||||
cache_write=5.5e-06,
|
||||
cache_write_above_272k=1.1e-05,
|
||||
cache_read=4.4e-07,
|
||||
cache_read_above_272k=8.8e-07,
|
||||
output_cost=2.2e-05,
|
||||
output_cost_above_272k=3.3e-05,
|
||||
),
|
||||
GptProfile(
|
||||
model_id="global.openai.gpt-5.6-sol",
|
||||
input_cost=4e-06, input_cost_above_272k=8e-06,
|
||||
cache_write=5e-06, cache_write_above_272k=1e-05,
|
||||
cache_read=4e-07, cache_read_above_272k=8e-07,
|
||||
output_cost=2e-05, output_cost_above_272k=3e-05,
|
||||
input_cost=4e-06,
|
||||
input_cost_above_272k=8e-06,
|
||||
cache_write=5e-06,
|
||||
cache_write_above_272k=1e-05,
|
||||
cache_read=4e-07,
|
||||
cache_read_above_272k=8e-07,
|
||||
output_cost=2e-05,
|
||||
output_cost_above_272k=3e-05,
|
||||
),
|
||||
GptProfile(
|
||||
model_id="us.openai.gpt-5.6-terra",
|
||||
input_cost=2.2e-06, input_cost_above_272k=4.4e-06,
|
||||
cache_write=2.75e-06, cache_write_above_272k=5.5e-06,
|
||||
cache_read=2.2e-07, cache_read_above_272k=4.4e-07,
|
||||
output_cost=1.32e-05, output_cost_above_272k=1.98e-05,
|
||||
input_cost=2.2e-06,
|
||||
input_cost_above_272k=4.4e-06,
|
||||
cache_write=2.75e-06,
|
||||
cache_write_above_272k=5.5e-06,
|
||||
cache_read=2.2e-07,
|
||||
cache_read_above_272k=4.4e-07,
|
||||
output_cost=1.32e-05,
|
||||
output_cost_above_272k=1.98e-05,
|
||||
),
|
||||
GptProfile(
|
||||
model_id="global.openai.gpt-5.6-terra",
|
||||
input_cost=2e-06, input_cost_above_272k=4e-06,
|
||||
cache_write=2.5e-06, cache_write_above_272k=5e-06,
|
||||
cache_read=2e-07, cache_read_above_272k=4e-07,
|
||||
output_cost=1.2e-05, output_cost_above_272k=1.8e-05,
|
||||
input_cost=2e-06,
|
||||
input_cost_above_272k=4e-06,
|
||||
cache_write=2.5e-06,
|
||||
cache_write_above_272k=5e-06,
|
||||
cache_read=2e-07,
|
||||
cache_read_above_272k=4e-07,
|
||||
output_cost=1.2e-05,
|
||||
output_cost_above_272k=1.8e-05,
|
||||
),
|
||||
GptProfile(
|
||||
model_id="us.openai.gpt-5.6-luna",
|
||||
input_cost=2.2e-07, input_cost_above_272k=4.4e-07,
|
||||
cache_write=2.75e-07, cache_write_above_272k=5.5e-07,
|
||||
cache_read=2.2e-08, cache_read_above_272k=4.4e-08,
|
||||
output_cost=1.32e-06, output_cost_above_272k=1.98e-06,
|
||||
input_cost=2.2e-07,
|
||||
input_cost_above_272k=4.4e-07,
|
||||
cache_write=2.75e-07,
|
||||
cache_write_above_272k=5.5e-07,
|
||||
cache_read=2.2e-08,
|
||||
cache_read_above_272k=4.4e-08,
|
||||
output_cost=1.32e-06,
|
||||
output_cost_above_272k=1.98e-06,
|
||||
),
|
||||
GptProfile(
|
||||
model_id="global.openai.gpt-5.6-luna",
|
||||
input_cost=2e-07, input_cost_above_272k=4e-07,
|
||||
cache_write=2.5e-07, cache_write_above_272k=5e-07,
|
||||
cache_read=2e-08, cache_read_above_272k=4e-08,
|
||||
output_cost=1.2e-06, output_cost_above_272k=1.8e-06,
|
||||
input_cost=2e-07,
|
||||
input_cost_above_272k=4e-07,
|
||||
cache_write=2.5e-07,
|
||||
cache_write_above_272k=5e-07,
|
||||
cache_read=2e-08,
|
||||
cache_read_above_272k=4e-08,
|
||||
output_cost=1.2e-06,
|
||||
output_cost_above_272k=1.8e-06,
|
||||
),
|
||||
]
|
||||
|
||||
|
|
@ -116,112 +134,18 @@ def _bedrock_response(model, usage):
|
|||
)
|
||||
|
||||
|
||||
def test_proxy_cost_calculation_scenario():
|
||||
"""Test exact GitHub issue scenario: proxy cost calculation"""
|
||||
model = "litellm_proxy/bedrock/us.anthropic.claude-3-5-haiku-20241022-v1:0"
|
||||
|
||||
# Test model info lookup works
|
||||
model_info = _get_model_info_helper(
|
||||
model=model, custom_llm_provider="litellm_proxy"
|
||||
)
|
||||
assert model_info is not None
|
||||
|
||||
# Test cost calculation works
|
||||
response = ModelResponse(
|
||||
id="test",
|
||||
created=1234567890,
|
||||
model=model,
|
||||
object="chat.completion",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(content="Test", role="assistant"),
|
||||
)
|
||||
],
|
||||
usage=Usage(total_tokens=150, prompt_tokens=100, completion_tokens=50),
|
||||
)
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=response, model=model, custom_llm_provider="litellm_proxy"
|
||||
)
|
||||
expected_cost = (100 * 8e-07) + (50 * 4e-06)
|
||||
assert cost == expected_cost
|
||||
|
||||
|
||||
@pytest.mark.parametrize("profile", GPT_5_6_PROFILES, ids=lambda p: p.model_id)
|
||||
def test_bedrock_gpt_5_6_profiles_route_to_converse(profile, local_model_cost_map):
|
||||
"""GPT-5.6 is served by Converse on bedrock-runtime, never by Invoke."""
|
||||
assert BedrockModelInfo.get_bedrock_route(f"bedrock/{profile.model_id}") == "converse"
|
||||
|
||||
|
||||
def test_bedrock_gpt_5_6_above_272k_tier_applies_to_cost(local_model_cost_map):
|
||||
"""A prompt over 272K tokens is billed at the long-context rate, not the base rate."""
|
||||
response = _bedrock_response(
|
||||
"bedrock/us.openai.gpt-5.6-sol",
|
||||
Usage(prompt_tokens=300000, completion_tokens=1000, total_tokens=301000),
|
||||
)
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=response,
|
||||
model="bedrock/us.openai.gpt-5.6-sol",
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx((300000 * 8.8e-06) + (1000 * 3.3e-05), rel=1e-9)
|
||||
|
||||
|
||||
def test_bedrock_gpt_5_6_bills_cache_read_tokens(local_model_cost_map):
|
||||
"""Bedrock caches long prefixes implicitly and reports them, so a cache-read turn
|
||||
must be billed at the cache rate rather than dropped to zero."""
|
||||
usage = Usage(
|
||||
prompt_tokens=15611,
|
||||
completion_tokens=5,
|
||||
total_tokens=15616,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=15609),
|
||||
)
|
||||
response = _bedrock_response("bedrock/us.openai.gpt-5.6-sol", usage)
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=response,
|
||||
model="bedrock/us.openai.gpt-5.6-sol",
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
|
||||
expected = (2 * 4.4e-06) + (15609 * 4.4e-07) + (5 * 2.2e-05)
|
||||
assert cost == pytest.approx(expected, rel=1e-9)
|
||||
# Without cache_read_input_token_cost the cached prefix bills at zero.
|
||||
assert cost > (15611 * 4.4e-06) * 0.1
|
||||
|
||||
|
||||
def test_bedrock_gpt_5_6_bills_cache_write_tokens(local_model_cost_map):
|
||||
"""The write side of the same cache cycle is billed at the 30m cache-write rate."""
|
||||
usage = Usage(
|
||||
prompt_tokens=15611,
|
||||
completion_tokens=5,
|
||||
total_tokens=15616,
|
||||
cache_creation_input_tokens=15609,
|
||||
)
|
||||
response = _bedrock_response("bedrock/us.openai.gpt-5.6-sol", usage)
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=response,
|
||||
model="bedrock/us.openai.gpt-5.6-sol",
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
|
||||
expected = (2 * 4.4e-06) + (15609 * 5.5e-06) + (5 * 2.2e-05)
|
||||
assert cost == pytest.approx(expected, rel=1e-9)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("profile", GPT_5_6_PROFILES, ids=lambda p: p.model_id)
|
||||
def test_bedrock_gpt_5_6_offers_tools_and_reasoning_effort_but_not_thinking(profile, local_model_cost_map):
|
||||
"""GPT-5.x on Converse maps reasoning_effort to reasoning.effort, so reasoning_effort
|
||||
is offered while the Anthropic-only thinking/output_config are not, alongside the tool
|
||||
params these models accept."""
|
||||
supported = AmazonConverseConfig().get_supported_openai_params(
|
||||
model=f"bedrock/{profile.model_id}"
|
||||
)
|
||||
supported = AmazonConverseConfig().get_supported_openai_params(model=f"bedrock/{profile.model_id}")
|
||||
|
||||
assert "tools" in supported
|
||||
assert "tool_choice" in supported
|
||||
|
|
|
|||
|
|
@ -3,10 +3,8 @@ from pathlib import Path
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse, OCRUsageInfo
|
||||
|
||||
COST_PER_PAGE = 0.0015
|
||||
REPO_ROOT = Path(__file__).parents[5]
|
||||
COST_MAPS = [
|
||||
REPO_ROOT / "model_prices_and_context_window.json",
|
||||
|
|
@ -28,17 +26,3 @@ def test_model_info_resolves_ocr_mode_and_price(local_model_cost_map, model: str
|
|||
info = litellm.get_model_info(model=model, custom_llm_provider=provider)
|
||||
|
||||
assert info["mode"] == "ocr"
|
||||
assert info["ocr_cost_per_page"] == COST_PER_PAGE
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model, provider", MODELS)
|
||||
@pytest.mark.parametrize("pages_processed", [1, 3])
|
||||
def test_cost_scales_with_billed_pages(local_model_cost_map, model: str, provider: str, pages_processed: int) -> None:
|
||||
cost = completion_cost(
|
||||
completion_response=_ocr_response(model.split("/", 1)[1], pages_processed),
|
||||
model=model,
|
||||
custom_llm_provider=provider,
|
||||
call_type="ocr",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(COST_PER_PAGE * pages_processed)
|
||||
|
|
|
|||
|
|
@ -215,7 +215,6 @@ def test_every_model_without_published_cache_dbu_bills_cache_at_its_own_input_ra
|
|||
and model not in PUBLISHED_DBU_PER_MILLION
|
||||
]
|
||||
|
||||
assert len(without_published_rates) == 14
|
||||
for model in without_published_rates:
|
||||
info = _model_info(model)
|
||||
for field in CACHE_FIELDS:
|
||||
|
|
|
|||
|
|
@ -1,51 +0,0 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
|
||||
|
||||
def test_databricks_pricing_integrity():
|
||||
"""
|
||||
Verifies that for all Databricks models in model_prices_and_context_window.json:
|
||||
USD Price == DBU Price * 0.07
|
||||
"""
|
||||
json_path = os.path.join(
|
||||
os.path.dirname(__file__), "../../../../model_prices_and_context_window.json"
|
||||
)
|
||||
|
||||
# Verify file exists
|
||||
assert os.path.exists(
|
||||
json_path
|
||||
), f"Could not find model_prices_and_context_window.json at {json_path}"
|
||||
|
||||
with open(json_path, "r") as f:
|
||||
data = json.load(f)
|
||||
|
||||
conversion_rate = 0.07 # 1 DBU = 0.07 USD
|
||||
errors = []
|
||||
|
||||
for model, info in data.items():
|
||||
if info.get("litellm_provider") == "databricks":
|
||||
# Check Input Cost
|
||||
input_usd = info.get("input_cost_per_token")
|
||||
input_dbu = info.get("input_dbu_cost_per_token")
|
||||
|
||||
if input_usd is not None and input_dbu is not None:
|
||||
expected = input_dbu * conversion_rate
|
||||
# Allow small floating point difference
|
||||
if abs(input_usd - expected) > 1e-9:
|
||||
errors.append(
|
||||
f"{model} input mismatch: USD={input_usd}, DBU={input_dbu}, Expected={expected}"
|
||||
)
|
||||
|
||||
# Check Output Cost
|
||||
output_usd = info.get("output_cost_per_token")
|
||||
output_dbu = info.get("output_dbu_cost_per_token")
|
||||
|
||||
if output_usd is not None and output_dbu is not None:
|
||||
expected = output_dbu * conversion_rate
|
||||
if abs(output_usd - expected) > 1e-9:
|
||||
errors.append(
|
||||
f"{model} output mismatch: USD={output_usd}, DBU={output_dbu}, Expected={expected}"
|
||||
)
|
||||
|
||||
assert not errors, "\n" + "\n".join(errors)
|
||||
|
|
@ -1,10 +1,6 @@
|
|||
|
||||
import math
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.llms.fireworks_ai.cost_calculator import cost_per_token
|
||||
from litellm.types.utils import OffPeakPricing, PromptTokensDetailsWrapper, Usage
|
||||
|
|
@ -26,49 +22,16 @@ def _usage(prompt_tokens: int, cached_tokens: int, completion_tokens: int) -> Us
|
|||
)
|
||||
|
||||
|
||||
def test_cached_prompt_tokens_billed_at_cache_read_rate():
|
||||
prompt_tokens = 7036
|
||||
cached_tokens = 7020
|
||||
completion_tokens = 8
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model=MODEL, usage=_usage(prompt_tokens, cached_tokens, completion_tokens)
|
||||
)
|
||||
|
||||
expected_prompt_cost = (prompt_tokens - cached_tokens) * INPUT_COST + cached_tokens * CACHE_READ_COST
|
||||
assert prompt_cost == pytest.approx(expected_prompt_cost)
|
||||
assert completion_cost == pytest.approx(completion_tokens * OUTPUT_COST)
|
||||
|
||||
full_rate_cost = prompt_tokens * INPUT_COST
|
||||
assert prompt_cost < full_rate_cost
|
||||
|
||||
|
||||
def test_warm_call_cheaper_than_cold_call():
|
||||
prompt_tokens = 7036
|
||||
completion_tokens = 8
|
||||
|
||||
cold_prompt_cost, _ = cost_per_token(
|
||||
model=MODEL, usage=_usage(prompt_tokens, 16, completion_tokens)
|
||||
)
|
||||
warm_prompt_cost, _ = cost_per_token(
|
||||
model=MODEL, usage=_usage(prompt_tokens, 7020, completion_tokens)
|
||||
)
|
||||
cold_prompt_cost, _ = cost_per_token(model=MODEL, usage=_usage(prompt_tokens, 16, completion_tokens))
|
||||
warm_prompt_cost, _ = cost_per_token(model=MODEL, usage=_usage(prompt_tokens, 7020, completion_tokens))
|
||||
|
||||
assert warm_prompt_cost < cold_prompt_cost
|
||||
|
||||
|
||||
def test_no_cached_tokens_matches_full_input_rate():
|
||||
prompt_tokens = 100
|
||||
completion_tokens = 10
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model=MODEL, usage=_usage(prompt_tokens, 0, completion_tokens)
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(prompt_tokens * INPUT_COST)
|
||||
assert completion_cost == pytest.approx(completion_tokens * OUTPUT_COST)
|
||||
|
||||
|
||||
OFF_PEAK_MODEL = "accounts/fireworks/models/off-peak-test"
|
||||
OFF_PEAK_WINDOW = "14:00-00:00"
|
||||
INSIDE_WINDOW = datetime(2026, 9, 3, 17, 25, tzinfo=timezone.utc)
|
||||
|
|
@ -78,7 +41,9 @@ STANDARD_OUTPUT_COST = 6e-07
|
|||
STANDARD_CACHE_READ_COST = 1.5e-08
|
||||
|
||||
|
||||
def _register_off_peak_model(off_peak_pricing: OffPeakPricing, cache_read_cost: float | None = STANDARD_CACHE_READ_COST) -> None:
|
||||
def _register_off_peak_model(
|
||||
off_peak_pricing: OffPeakPricing, cache_read_cost: float | None = STANDARD_CACHE_READ_COST
|
||||
) -> None:
|
||||
litellm.model_cost[f"fireworks_ai/{OFF_PEAK_MODEL}"] = {
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"mode": "chat",
|
||||
|
|
@ -151,7 +116,9 @@ def test_off_peak_window_bills_cached_tokens_at_the_off_peak_input_rate_without_
|
|||
def test_off_peak_defaults_to_the_current_time():
|
||||
"""The proxy's cost dispatch passes no clock, so an all-day window has to apply on the
|
||||
default current time."""
|
||||
_register_off_peak_model({"hours_utc": "00:00-00:00", "input_cost_per_token": 1e-08, "output_cost_per_token": 2e-08})
|
||||
_register_off_peak_model(
|
||||
{"hours_utc": "00:00-00:00", "input_cost_per_token": 1e-08, "output_cost_per_token": 2e-08}
|
||||
)
|
||||
usage = _usage(prompt_tokens=1000, cached_tokens=0, completion_tokens=200)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model=OFF_PEAK_MODEL, usage=usage)
|
||||
|
|
|
|||
|
|
@ -1,65 +0,0 @@
|
|||
"""
|
||||
Regression test for Fireworks Kimi K2.5 / K2.6 / K2.7 context and output limits.
|
||||
|
||||
Fireworks publishes a 262144-token context window for every Kimi K2.5, K2.6 and
|
||||
K2.7 model, but caps generation well below that. A previous bulk edit had flattened
|
||||
max_output_tokens/max_tokens to 262144 (equal to the context window), which let the
|
||||
pre-call context-window check admit requests asking for a full 262144-token
|
||||
completion that Fireworks then rejects. These assertions pin the corrected per-alias
|
||||
limits so a future bulk edit can't silently flatten them again.
|
||||
"""
|
||||
|
||||
import json
|
||||
from importlib.resources import files
|
||||
|
||||
import pytest
|
||||
|
||||
CONTEXT_WINDOW = 262144
|
||||
OUTPUT_LIMIT = 32768
|
||||
|
||||
KIMI_ALIASES = (
|
||||
"fireworks_ai/kimi-k2p5",
|
||||
"fireworks_ai/kimi-k2p6",
|
||||
"fireworks_ai/kimi-k2p6-fast",
|
||||
"fireworks_ai/kimi-k2p7-code",
|
||||
"fireworks_ai/kimi-k2p7-code-fast",
|
||||
"fireworks_ai/accounts/fireworks/models/kimi-k2p5",
|
||||
"fireworks_ai/accounts/fireworks/models/kimi-k2p6",
|
||||
"fireworks_ai/accounts/fireworks/models/kimi-k2p7-code",
|
||||
"fireworks_ai/accounts/fireworks/routers/kimi-k2p6-fast",
|
||||
"fireworks_ai/accounts/fireworks/routers/kimi-k2p7-code-fast",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def use_local_model_cost_map():
|
||||
monkeypatch = pytest.MonkeyPatch()
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
|
||||
import litellm
|
||||
from litellm.utils import _invalidate_model_cost_lowercase_map
|
||||
|
||||
original_model_cost = litellm.model_cost
|
||||
litellm.model_cost = json.loads(
|
||||
files("litellm")
|
||||
.joinpath("model_prices_and_context_window_backup.json")
|
||||
.read_text(encoding="utf-8")
|
||||
)
|
||||
litellm.get_model_info.cache_clear()
|
||||
_invalidate_model_cost_lowercase_map()
|
||||
try:
|
||||
yield litellm
|
||||
finally:
|
||||
litellm.model_cost = original_model_cost
|
||||
litellm.get_model_info.cache_clear()
|
||||
_invalidate_model_cost_lowercase_map()
|
||||
monkeypatch.undo()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("alias", KIMI_ALIASES)
|
||||
def test_fireworks_kimi_get_model_info_limits(use_local_model_cost_map, alias):
|
||||
model_info = use_local_model_cost_map.get_model_info(model=alias)
|
||||
|
||||
assert model_info["max_input_tokens"] == CONTEXT_WINDOW
|
||||
assert model_info["max_output_tokens"] == OUTPUT_LIMIT
|
||||
assert model_info["max_tokens"] == OUTPUT_LIMIT
|
||||
|
|
@ -4,7 +4,6 @@ import json
|
|||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.llms.gemini.audio_transcription.transformation import (
|
||||
GeminiAudioTranscriptionConfig,
|
||||
|
|
@ -318,15 +317,3 @@ class TestCostRegression:
|
|||
assert live_entry["input_cost_per_token"] == 3.5e-06
|
||||
assert live_entry["output_cost_per_token"] == 2.1e-05
|
||||
assert live_entry["supported_endpoints"] == ["/v1/realtime"]
|
||||
|
||||
def test_completion_cost_bills_provider_reported_tokens(self, config, local_cost_map):
|
||||
payload = json.loads(json.dumps(COMPLETED_RESPONSE))
|
||||
payload["usage"]["total_output_tokens"] = 10
|
||||
payload["usage"]["total_tokens"] = 210
|
||||
response = config.transform_audio_transcription_response(make_response(payload))
|
||||
cost = litellm.completion_cost(
|
||||
completion_response=response,
|
||||
model="gemini/gemini-3.5-transcribe",
|
||||
call_type="transcription",
|
||||
)
|
||||
assert cost == pytest.approx(199 * 2e-06 + 1 * 2e-06 + 10 * 1.2e-05)
|
||||
|
|
|
|||
|
|
@ -430,6 +430,25 @@ class TestGeminiVideoConfig:
|
|||
assert result.usage["video_resolution"] == "1080p"
|
||||
assert result.usage["duration_seconds"] == 8.0
|
||||
|
||||
def test_transform_video_create_response_usage_includes_video_count(self):
|
||||
"""Regression for LIT-6896: sampleCount (number of generated videos) is copied into usage for billing."""
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {"name": "operations/generate_1234567890"}
|
||||
request_data = {
|
||||
"instances": [{"prompt": "Test"}],
|
||||
"parameters": {"durationSeconds": 8, "sampleCount": 3},
|
||||
}
|
||||
result = self.config.transform_video_create_response(
|
||||
model="gemini/veo-3.1-fast-generate-preview",
|
||||
raw_response=mock_response,
|
||||
logging_obj=self.mock_logging_obj,
|
||||
custom_llm_provider="gemini",
|
||||
request_data=request_data,
|
||||
)
|
||||
assert result.usage is not None
|
||||
assert result.usage["video_count"] == 3
|
||||
assert result.usage["duration_seconds"] == 8.0
|
||||
|
||||
def test_transform_video_create_response_cost_tracking_with_different_durations(
|
||||
self,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -1,128 +0,0 @@
|
|||
"""
|
||||
Cost tests for Mistral OCR models against the real litellm cost map
|
||||
(no monkeypatching of get_model_info). These regress the pricing entries
|
||||
for mistral-ocr-4-0 and mistral-ocr-latest, which now both resolve to
|
||||
OCR 4 at $4 / 1000 pages.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse, OCRUsageInfo
|
||||
|
||||
OCR4_COST_PER_PAGE = 0.004
|
||||
OCR4_ANNOTATION_COST_PER_PAGE = 0.005
|
||||
|
||||
REPO_ROOT = Path(__file__).parents[5]
|
||||
MAIN_COST_MAP = REPO_ROOT / "model_prices_and_context_window.json"
|
||||
BACKUP_COST_MAP = REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json"
|
||||
|
||||
OCR3_MODEL = "mistral/mistral-ocr-2512"
|
||||
OCR3_COST_PER_PAGE = 0.002
|
||||
OCR3_ANNOTATION_COST_PER_PAGE = 0.003
|
||||
|
||||
AZURE_DOC_AI_MODEL = "azure_ai/mistral-document-ai-2512"
|
||||
AZURE_DOC_AI_COST_PER_PAGE = 0.003
|
||||
|
||||
|
||||
def _ocr_response(model: str, pages_processed: int) -> OCRResponse:
|
||||
return OCRResponse(
|
||||
pages=[OCRPage(index=i, markdown=f"page {i}") for i in range(pages_processed)],
|
||||
model=model,
|
||||
usage_info=OCRUsageInfo(pages_processed=pages_processed),
|
||||
)
|
||||
|
||||
|
||||
def _annotated_ocr_response(model: str, pages_processed: int | None, annotation_pages: int) -> OCRResponse:
|
||||
return OCRResponse(
|
||||
pages=[],
|
||||
model=model,
|
||||
usage_info=OCRUsageInfo(pages_processed=pages_processed, pages_processed_annotation=annotation_pages),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", ["mistral-ocr-4-0", "mistral-ocr-latest"])
|
||||
@pytest.mark.parametrize("pages_processed", [1, 3, 10])
|
||||
def test_ocr4_cost_scales_with_pages(model: str, pages_processed: int) -> None:
|
||||
cost = completion_cost(
|
||||
completion_response=_ocr_response(model, pages_processed),
|
||||
model=f"mistral/{model}",
|
||||
custom_llm_provider="mistral",
|
||||
call_type="ocr",
|
||||
)
|
||||
assert cost == pytest.approx(OCR4_COST_PER_PAGE * pages_processed)
|
||||
|
||||
|
||||
def test_ocr3_model_info_price(local_model_cost_map) -> None:
|
||||
info = litellm.get_model_info(model=OCR3_MODEL, custom_llm_provider="mistral")
|
||||
assert info["ocr_cost_per_page"] == OCR3_COST_PER_PAGE
|
||||
|
||||
|
||||
@pytest.mark.parametrize("pages_processed", [1, 3, 10])
|
||||
def test_ocr3_cost_scales_with_pages(local_model_cost_map, pages_processed: int) -> None:
|
||||
cost = completion_cost(
|
||||
completion_response=_ocr_response("mistral-ocr-2512", pages_processed),
|
||||
model=OCR3_MODEL,
|
||||
custom_llm_provider="mistral",
|
||||
call_type="ocr",
|
||||
)
|
||||
assert cost == pytest.approx(OCR3_COST_PER_PAGE * pages_processed)
|
||||
|
||||
|
||||
def test_ocr3_bills_ocr_and_annotation_pages_at_their_own_rates(local_model_cost_map) -> None:
|
||||
cost = completion_cost(
|
||||
completion_response=_annotated_ocr_response("mistral-ocr-2512", 2, 3),
|
||||
model=OCR3_MODEL,
|
||||
custom_llm_provider="mistral",
|
||||
call_type="ocr",
|
||||
)
|
||||
assert cost == pytest.approx(2 * OCR3_COST_PER_PAGE + 3 * OCR3_ANNOTATION_COST_PER_PAGE)
|
||||
|
||||
|
||||
def test_ocr3_bills_annotation_only_response(local_model_cost_map) -> None:
|
||||
cost = completion_cost(
|
||||
completion_response=_annotated_ocr_response("mistral-ocr-2512", 0, 3),
|
||||
model=OCR3_MODEL,
|
||||
custom_llm_provider="mistral",
|
||||
call_type="ocr",
|
||||
)
|
||||
assert cost == pytest.approx(3 * OCR3_ANNOTATION_COST_PER_PAGE)
|
||||
|
||||
|
||||
def test_ocr3_bills_annotation_pages_when_pages_processed_missing(local_model_cost_map) -> None:
|
||||
cost = completion_cost(
|
||||
completion_response=_annotated_ocr_response("mistral-ocr-2512", None, 4),
|
||||
model=OCR3_MODEL,
|
||||
custom_llm_provider="mistral",
|
||||
call_type="ocr",
|
||||
)
|
||||
assert cost == pytest.approx(4 * OCR3_ANNOTATION_COST_PER_PAGE)
|
||||
|
||||
|
||||
def test_azure_doc_ai_annotation_pages_fall_back_to_ocr_rate(local_model_cost_map) -> None:
|
||||
info = litellm.get_model_info(model=AZURE_DOC_AI_MODEL, custom_llm_provider="azure_ai")
|
||||
assert info.get("annotation_cost_per_page") is None
|
||||
assert info["ocr_cost_per_page"] == AZURE_DOC_AI_COST_PER_PAGE
|
||||
cost = completion_cost(
|
||||
completion_response=_annotated_ocr_response("mistral-document-ai-2512", 0, 1),
|
||||
model=AZURE_DOC_AI_MODEL,
|
||||
custom_llm_provider="azure_ai",
|
||||
call_type="ocr",
|
||||
)
|
||||
assert cost == pytest.approx(AZURE_DOC_AI_COST_PER_PAGE)
|
||||
|
||||
|
||||
def test_azure_ocr4_bills_ocr_and_annotation_pages_at_their_own_rates(local_model_cost_map) -> None:
|
||||
info = litellm.get_model_info(model="azure_ai/mistral-ocr-4-0", custom_llm_provider="azure_ai")
|
||||
assert info["ocr_cost_per_page"] == OCR4_COST_PER_PAGE
|
||||
assert info["annotation_cost_per_page"] == OCR4_ANNOTATION_COST_PER_PAGE
|
||||
cost = completion_cost(
|
||||
completion_response=_annotated_ocr_response("mistral-ocr-4-0", 2, 3),
|
||||
model="azure_ai/mistral-ocr-4-0",
|
||||
custom_llm_provider="azure_ai",
|
||||
call_type="ocr",
|
||||
)
|
||||
assert cost == pytest.approx(2 * OCR4_COST_PER_PAGE + 3 * OCR4_ANNOTATION_COST_PER_PAGE)
|
||||
|
|
@ -0,0 +1,296 @@
|
|||
import json
|
||||
from types import MappingProxyType
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.llms.nvidia_nim.passthrough.transformation import (
|
||||
NvidiaNimPassthroughConfig,
|
||||
nvidia_nim_model_group_in_path,
|
||||
nvidia_nim_model_groups,
|
||||
nvidia_nim_router_model_in_endpoint,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
NIM_BASE = "http://nim.internal:8000"
|
||||
INFER_BODY = {
|
||||
"input": [
|
||||
{"type": "image_url", "url": "data:image/png;base64,AAAA"},
|
||||
{"type": "image_url", "url": "data:image/png;base64,BBBB"},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clear_nvidia_nim_env(monkeypatch):
|
||||
for env_var in ("NVIDIA_NIM_API_BASE", "NVIDIA_NIM_API_KEY"):
|
||||
monkeypatch.delenv(env_var, raising=False)
|
||||
monkeypatch.setattr(litellm, "api_base", None)
|
||||
monkeypatch.setattr(litellm, "api_key", None)
|
||||
|
||||
|
||||
def test_provider_config_manager_resolves_nvidia_nim_passthrough_config():
|
||||
config = ProviderConfigManager.get_provider_passthrough_config(
|
||||
model="nvidia/nemoretriever-page-elements-v2", provider=LlmProviders.NVIDIA_NIM
|
||||
)
|
||||
|
||||
assert isinstance(config, NvidiaNimPassthroughConfig)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base, endpoint, litellm_params, expected",
|
||||
[
|
||||
(NIM_BASE, "nim-page/v1/infer", {"litellm_metadata": {"model_group": "nim-page"}}, f"{NIM_BASE}/v1/infer"),
|
||||
(
|
||||
f"{NIM_BASE}/v1",
|
||||
"nim-page/v1/infer",
|
||||
{"litellm_metadata": {"model_group": "nim-page"}},
|
||||
f"{NIM_BASE}/v1/infer",
|
||||
),
|
||||
(f"{NIM_BASE}/v1/", "/v1/infer", {}, f"{NIM_BASE}/v1/infer"),
|
||||
(NIM_BASE, "v1/infer", {}, f"{NIM_BASE}/v1/infer"),
|
||||
(f"{NIM_BASE}/v2", "v1/infer", {}, f"{NIM_BASE}/v2/v1/infer"),
|
||||
(f"{NIM_BASE}/infer", "infer", {}, f"{NIM_BASE}/infer/infer"),
|
||||
(NIM_BASE, "nvidia/nemoretriever-page-elements-v2/v1/infer", {}, f"{NIM_BASE}/v1/infer"),
|
||||
(
|
||||
NIM_BASE,
|
||||
"nvidia/nemoretriever-page-elements-v2/v1/infer",
|
||||
{"litellm_metadata": {"model_group": "nvidia"}},
|
||||
f"{NIM_BASE}/v1/infer",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_relay_url_strips_the_model_group_and_never_doubles_the_api_version(
|
||||
api_base, endpoint, litellm_params, expected
|
||||
):
|
||||
url, base = NvidiaNimPassthroughConfig().get_complete_url(
|
||||
api_base=api_base,
|
||||
api_key=None,
|
||||
model="nvidia/nemoretriever-page-elements-v2",
|
||||
endpoint=endpoint,
|
||||
request_query_params=None,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert str(url) == expected
|
||||
assert base == expected.removesuffix("/v1/infer").removesuffix("/infer")
|
||||
|
||||
|
||||
def test_query_params_are_forwarded_on_the_relay_url():
|
||||
url, _ = NvidiaNimPassthroughConfig().get_complete_url(
|
||||
api_base=NIM_BASE,
|
||||
api_key=None,
|
||||
model="nvidia/nemoretriever-page-elements-v2",
|
||||
endpoint="v1/infer",
|
||||
request_query_params={"timeout": "30"},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert str(url) == f"{NIM_BASE}/v1/infer?timeout=30"
|
||||
|
||||
|
||||
def test_env_api_base_is_used_when_the_deployment_has_none(monkeypatch):
|
||||
monkeypatch.setenv("NVIDIA_NIM_API_BASE", f"{NIM_BASE}/v1")
|
||||
|
||||
url, _ = NvidiaNimPassthroughConfig().get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="nvidia/nemoretriever-page-elements-v2",
|
||||
endpoint="v1/infer",
|
||||
request_query_params=None,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert str(url) == f"{NIM_BASE}/v1/infer"
|
||||
|
||||
|
||||
def test_missing_api_base_raises_instead_of_building_a_relative_url():
|
||||
with pytest.raises(ValueError, match="NVIDIA_NIM_API_BASE"):
|
||||
NvidiaNimPassthroughConfig().get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="nvidia/nemoretriever-page-elements-v2",
|
||||
endpoint="v1/infer",
|
||||
request_query_params=None,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
|
||||
def test_deployment_key_becomes_a_bearer_token_and_caller_headers_are_kept():
|
||||
caller_headers = MappingProxyType({"x-request-id": "abc"})
|
||||
|
||||
headers = NvidiaNimPassthroughConfig().validate_environment(
|
||||
headers=caller_headers,
|
||||
model="nvidia/nemoretriever-page-elements-v2",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key="nvapi-secret",
|
||||
)
|
||||
|
||||
assert headers == {"x-request-id": "abc", "Authorization": "Bearer nvapi-secret"}
|
||||
|
||||
|
||||
def test_self_hosted_nim_without_a_key_sends_no_authorization_header():
|
||||
headers = NvidiaNimPassthroughConfig().validate_environment(
|
||||
headers={}, model="nvidia/x", messages=[], optional_params={}, litellm_params={}, api_key=None
|
||||
)
|
||||
|
||||
assert "Authorization" not in headers
|
||||
|
||||
|
||||
def test_env_api_key_fills_in_when_the_deployment_has_none(monkeypatch):
|
||||
monkeypatch.setenv("NVIDIA_NIM_API_KEY", "nvapi-from-env")
|
||||
|
||||
assert NvidiaNimPassthroughConfig.get_api_key(None) == "nvapi-from-env"
|
||||
assert NvidiaNimPassthroughConfig.get_api_key("nvapi-deployment") == "nvapi-deployment"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint, router_models, expected",
|
||||
[
|
||||
("nim-page/v1/infer", ("nim-page", "nim-table"), "nim-page"),
|
||||
("/nim-page/v1/infer", ("nim-page",), "nim-page"),
|
||||
(
|
||||
"nvidia/nemoretriever-page-elements-v2/v1/infer",
|
||||
("nvidia/nemoretriever-page-elements-v2",),
|
||||
"nvidia/nemoretriever-page-elements-v2",
|
||||
),
|
||||
("nim/v1/infer", ("nim", "nim/v1"), "nim/v1"),
|
||||
("v1/infer", ("nim-page",), None),
|
||||
("nim-page-elements/v1/infer", ("nim-page",), None),
|
||||
("", ("nim-page",), None),
|
||||
],
|
||||
)
|
||||
def test_router_model_in_endpoint_takes_the_longest_leading_model_group(endpoint, router_models, expected):
|
||||
assert nvidia_nim_router_model_in_endpoint(endpoint, frozenset(router_models)) == expected
|
||||
|
||||
|
||||
def _deployment(model_name: str, model: str, custom_llm_provider: str | None = None):
|
||||
litellm_params = (
|
||||
{"model": model}
|
||||
if custom_llm_provider is None
|
||||
else {"model": model, "custom_llm_provider": custom_llm_provider}
|
||||
)
|
||||
return {"model_name": model_name, "litellm_params": litellm_params}
|
||||
|
||||
|
||||
MIXED_DEPLOYMENTS = (
|
||||
_deployment("nim-page", "nvidia_nim/nvidia/nemoretriever-page-elements-v2"),
|
||||
_deployment("nim-table", "nvidia/nemoretriever-table-structure-v1", custom_llm_provider="nvidia_nim"),
|
||||
_deployment("mixed", "nvidia_nim/nvidia/nemoretriever-page-elements-v2"),
|
||||
_deployment("mixed", "openai/gpt-4o"),
|
||||
_deployment("gpt-4o", "openai/gpt-4o"),
|
||||
)
|
||||
|
||||
|
||||
def test_model_groups_only_admit_groups_whose_every_deployment_is_nim_backed():
|
||||
assert nvidia_nim_model_groups(MIXED_DEPLOYMENTS) == frozenset({"nim-page", "nim-table"})
|
||||
assert nvidia_nim_model_groups(None) == frozenset()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"path, expected",
|
||||
[
|
||||
("/nvidia_nim/nim-page/v1/infer", "nim-page"),
|
||||
("/NVIDIA_NIM/nim-table/v1/infer", "nim-table"),
|
||||
("nim-page/v1/infer", "nim-page"),
|
||||
("/nvidia_nim/mixed/v1/infer", None),
|
||||
("mixed/v1/infer", None),
|
||||
("/nvidia_nim/gpt-4o/v1/infer", None),
|
||||
("/nvidia_nim/v1/infer", None),
|
||||
],
|
||||
)
|
||||
def test_model_group_in_path_resolves_the_same_nim_only_groups_for_routes_and_endpoints(path, expected):
|
||||
assert nvidia_nim_model_group_in_path(path, MIXED_DEPLOYMENTS) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize("request_data, expected", [({"stream": True}, True), ({"stream": False}, False), ({}, False)])
|
||||
def test_is_streaming_request_reads_the_stream_flag(request_data, expected):
|
||||
assert NvidiaNimPassthroughConfig().is_streaming_request("v1/infer", request_data) is expected
|
||||
|
||||
|
||||
def test_non_streaming_relay_logs_the_upstream_json_body():
|
||||
response = httpx.Response(
|
||||
200,
|
||||
json={"data": [{"index": 0, "bounding_boxes": {}}]},
|
||||
request=httpx.Request("POST", f"{NIM_BASE}/v1/infer"),
|
||||
)
|
||||
|
||||
result = NvidiaNimPassthroughConfig().logging_non_streaming_response(
|
||||
model="nvidia/nemoretriever-page-elements-v2",
|
||||
custom_llm_provider="nvidia_nim",
|
||||
httpx_response=response,
|
||||
request_data=INFER_BODY,
|
||||
logging_obj=None, # pyright: ignore[reportArgumentType] # not read for a plain passthrough body
|
||||
endpoint="v1/infer",
|
||||
)
|
||||
|
||||
assert result == {"response": {"data": [{"index": 0, "bounding_boxes": {}}]}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_object_detection_relay_sends_the_native_body_unchanged_to_v1_infer():
|
||||
upstream_requests: list[httpx.Request] = []
|
||||
|
||||
def nim(request: httpx.Request) -> httpx.Response:
|
||||
upstream_requests.append(request)
|
||||
return httpx.Response(200, json={"data": [{"index": 0}, {"index": 1}]}, headers={"x-nim": "1"})
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=httpx.MockTransport(nim))
|
||||
|
||||
response = await litellm.allm_passthrough_route(
|
||||
model="nvidia_nim/nvidia/nemoretriever-page-elements-v2",
|
||||
endpoint="nim-page/v1/infer",
|
||||
method="POST",
|
||||
api_base=f"{NIM_BASE}/v1",
|
||||
api_key="nvapi-secret",
|
||||
json=dict(INFER_BODY),
|
||||
litellm_metadata={"model_group": "nim-page"},
|
||||
client=client,
|
||||
)
|
||||
|
||||
(sent,) = upstream_requests
|
||||
assert str(sent.url) == f"{NIM_BASE}/v1/infer"
|
||||
assert json.loads(sent.content) == INFER_BODY
|
||||
assert sent.headers["authorization"] == "Bearer nvapi-secret"
|
||||
assert response.status_code == 200
|
||||
assert response.headers["x-nim"] == "1"
|
||||
assert response.json() == {"data": [{"index": 0}, {"index": 1}]}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_relay_reaches_v1_infer_when_the_group_name_is_a_leading_segment_of_the_model_id():
|
||||
upstream_requests: list[httpx.Request] = []
|
||||
|
||||
def nim(request: httpx.Request) -> httpx.Response:
|
||||
upstream_requests.append(request)
|
||||
return httpx.Response(200, json={"data": [{"index": 0}]})
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
client.client = httpx.AsyncClient(transport=httpx.MockTransport(nim))
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "nvidia",
|
||||
"litellm_params": {
|
||||
"model": "nvidia_nim/nvidia/nemoretriever-page-elements-v2",
|
||||
"api_base": NIM_BASE,
|
||||
"api_key": "nvapi-secret",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
response = await router.allm_passthrough_route(
|
||||
model="nvidia", endpoint="nvidia/v1/infer", method="POST", json=dict(INFER_BODY), client=client
|
||||
)
|
||||
|
||||
(sent,) = upstream_requests
|
||||
assert str(sent.url) == f"{NIM_BASE}/v1/infer"
|
||||
assert json.loads(sent.content) == INFER_BODY
|
||||
assert response.status_code == 200
|
||||
|
|
@ -75,9 +75,3 @@ def test_shipped_per_second_models_bill_a_non_zero_cost(model, provider):
|
|||
prompt_cost, completion_cost = cost_per_second(model=model, custom_llm_provider=provider, duration=60.0)
|
||||
|
||||
assert prompt_cost + completion_cost > 0.0
|
||||
|
||||
|
||||
def test_whisper_bills_its_documented_rate_once():
|
||||
prompt_cost, completion_cost = cost_per_second(model="whisper-1", custom_llm_provider="openai", duration=30.0)
|
||||
|
||||
assert prompt_cost + completion_cost == pytest.approx(0.003)
|
||||
|
|
|
|||
|
|
@ -172,7 +172,6 @@ class TestSCXAIModelMetadata:
|
|||
assert info["supports_prompt_caching"] is True
|
||||
assert 0 < info["cache_read_input_token_cost"] < info["input_cost_per_token"]
|
||||
|
||||
assert info["max_output_tokens"] == 131072
|
||||
assert info["max_tokens"] == info["max_output_tokens"]
|
||||
assert info["max_input_tokens"] >= 1_000_000
|
||||
|
||||
|
|
|
|||
|
|
@ -14,17 +14,15 @@ from unittest.mock import patch
|
|||
import pytest
|
||||
|
||||
# Add the project root to Python path
|
||||
|
||||
import litellm
|
||||
from litellm.cost_calculator import completion_cost, cost_per_token
|
||||
from litellm.llms.perplexity.cost_calculator import (
|
||||
cost_per_token as perplexity_cost_per_token,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
CompletionTokensDetailsWrapper,
|
||||
OffPeakPricing,
|
||||
Usage,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -64,167 +62,6 @@ class TestPerplexityCostCalculator:
|
|||
}
|
||||
}
|
||||
|
||||
def test_basic_cost_calculation(self):
|
||||
"""Test basic cost calculation without additional fields."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
# Expected costs:
|
||||
# Input: 100 tokens * $2e-6 = $0.0002
|
||||
# Output: 50 tokens * $8e-6 = $0.0004
|
||||
expected_prompt_cost = 100 * 2e-6
|
||||
expected_completion_cost = 50 * 8e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_citation_tokens_cost_calculation(self):
|
||||
"""Test cost calculation with citation tokens."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
|
||||
# Add citation tokens
|
||||
usage.citation_tokens = 25
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
# Expected costs:
|
||||
# Input: 100 tokens * $2e-6 = $0.0002
|
||||
# Citation: 25 tokens * $2e-6 = $0.00005
|
||||
# Total prompt cost: $0.00025
|
||||
# Output: 50 tokens * $8e-6 = $0.0004
|
||||
expected_prompt_cost = (100 * 2e-6) + (25 * 2e-6)
|
||||
expected_completion_cost = 50 * 8e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_search_queries_cost_calculation(self):
|
||||
"""Test cost calculation with search queries."""
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
total_tokens=150,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(web_search_requests=3),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
# Expected costs:
|
||||
# Input: 100 tokens * $2e-6 = $0.0002
|
||||
# Output: 50 tokens * $8e-6 = $0.0004
|
||||
# Search: 3 queries * $0.005 per request = $0.015
|
||||
# Total completion cost: $0.0154
|
||||
expected_prompt_cost = 100 * 2e-6
|
||||
expected_completion_cost = (50 * 8e-6) + (3 * 0.005)
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_reasoning_tokens_from_direct_attribute(self):
|
||||
"""Test reasoning tokens cost calculation from direct attribute."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
|
||||
# Set reasoning tokens directly
|
||||
usage.reasoning_tokens = 20
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
# `completion_tokens` includes `reasoning_tokens` per the OpenAI/Perplexity
|
||||
# convention codified in PR #18607. Non-reasoning portion = 50 - 20 = 30.
|
||||
# Input: 100 tokens * $2e-6 = $0.0002
|
||||
# Output (text): 30 tokens * $8e-6 = $0.00024
|
||||
# Reasoning: 20 tokens * $3e-6 = $0.00006
|
||||
# Total completion cost = $0.0003
|
||||
expected_prompt_cost = 100 * 2e-6
|
||||
expected_completion_cost = ((50 - 20) * 8e-6) + (20 * 3e-6)
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_reasoning_tokens_from_completion_tokens_details(self):
|
||||
"""Test reasoning tokens cost calculation from completion_tokens_details."""
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
total_tokens=150,
|
||||
reasoning_tokens=20, # This should be stored in completion_tokens_details
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
# Same convention as the direct-attribute case above; reasoning is a subset of
|
||||
# completion_tokens, so non-reasoning portion = 50 - 20 = 30.
|
||||
expected_prompt_cost = 100 * 2e-6
|
||||
expected_completion_cost = ((50 - 20) * 8e-6) + (20 * 3e-6)
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_comprehensive_cost_calculation(self):
|
||||
"""Test cost calculation with all fields combined."""
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
total_tokens=150,
|
||||
reasoning_tokens=15,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(web_search_requests=2),
|
||||
)
|
||||
|
||||
# Add custom fields
|
||||
usage.citation_tokens = 30
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
# Expected costs (reasoning is a subset of completion_tokens):
|
||||
# Input: 100 tokens * $2e-6 = $0.0002
|
||||
# Citation: 30 tokens * $2e-6 = $0.00006
|
||||
# Total prompt cost = $0.00026
|
||||
# Output (text): (50 - 15) tokens * $8e-6 = $0.00028
|
||||
# Reasoning: 15 tokens * $3e-6 = $0.000045
|
||||
# Search: 2 queries * $0.005 per request = $0.01
|
||||
# Total completion cost = $0.010325
|
||||
expected_prompt_cost = (100 * 2e-6) + (30 * 2e-6)
|
||||
expected_completion_cost = ((50 - 15) * 8e-6) + (15 * 3e-6) + (2 * 0.005)
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_zero_values_handling(self):
|
||||
"""Test that zero or missing values are handled correctly."""
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
total_tokens=150,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(web_search_requests=0),
|
||||
)
|
||||
|
||||
# These should not raise errors and should not affect cost
|
||||
usage.citation_tokens = 0
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
# Should be same as basic calculation
|
||||
expected_prompt_cost = 100 * 2e-6
|
||||
expected_completion_cost = 50 * 8e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_missing_model_info_fields(self):
|
||||
"""Test behavior when model info is missing some fields."""
|
||||
usage = Usage(
|
||||
|
|
@ -237,18 +74,14 @@ class TestPerplexityCostCalculator:
|
|||
usage.citation_tokens = 25
|
||||
|
||||
# Mock get_model_info to return incomplete model info
|
||||
with patch(
|
||||
"litellm.llms.perplexity.cost_calculator.get_model_info"
|
||||
) as mock_get_model_info:
|
||||
with patch("litellm.llms.perplexity.cost_calculator.get_model_info") as mock_get_model_info:
|
||||
mock_get_model_info.return_value = {
|
||||
"input_cost_per_token": 2e-6,
|
||||
"output_cost_per_token": 8e-6,
|
||||
# Missing search_queries_cost_per_query
|
||||
}
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(model="sonar-deep-research", usage=usage)
|
||||
|
||||
# Should only calculate basic costs when fields are missing
|
||||
expected_prompt_cost = 100 * 2e-6
|
||||
|
|
@ -257,104 +90,6 @@ class TestPerplexityCostCalculator:
|
|||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_integration_with_main_cost_calculator(self):
|
||||
"""Test integration with the main LiteLLM cost calculator."""
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
total_tokens=150,
|
||||
reasoning_tokens=10,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(web_search_requests=1),
|
||||
)
|
||||
|
||||
usage.citation_tokens = 20
|
||||
|
||||
# Test main cost calculator
|
||||
prompt_cost, completion_cost_val = cost_per_token(
|
||||
model="sonar-deep-research",
|
||||
custom_llm_provider="perplexity",
|
||||
usage_object=usage,
|
||||
)
|
||||
|
||||
# Should match direct call to perplexity cost calculator
|
||||
expected_prompt, expected_completion = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost_val, expected_completion, rel_tol=1e-6)
|
||||
|
||||
def test_integration_with_completion_cost_function(self):
|
||||
"""Test integration with the completion_cost function."""
|
||||
from litellm import ModelResponse
|
||||
|
||||
# Create a mock ModelResponse
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
total_tokens=150,
|
||||
reasoning_tokens=10,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(web_search_requests=1),
|
||||
)
|
||||
usage.citation_tokens = 15
|
||||
|
||||
response = ModelResponse()
|
||||
response.usage = usage
|
||||
response.model = "sonar-deep-research"
|
||||
|
||||
# Test completion_cost function
|
||||
total_cost = completion_cost(
|
||||
completion_response=response, custom_llm_provider="perplexity"
|
||||
)
|
||||
|
||||
# Calculate expected total cost (reasoning is a subset of completion_tokens)
|
||||
expected_prompt_cost = (100 * 2e-6) + (15 * 2e-6) # Input + citation
|
||||
expected_completion_cost = (
|
||||
((50 - 10) * 8e-6) + (10 * 3e-6) + (1 * 0.005)
|
||||
) # Output (text) + reasoning + search
|
||||
expected_total = expected_prompt_cost + expected_completion_cost
|
||||
|
||||
assert math.isclose(total_cost, expected_total, rel_tol=1e-6)
|
||||
|
||||
@pytest.mark.parametrize("citation_tokens", [0, 10, 25, 100])
|
||||
@pytest.mark.parametrize("search_queries", [0, 1, 5, 10])
|
||||
@pytest.mark.parametrize("reasoning_tokens", [0, 15, 30])
|
||||
def test_cost_calculation_combinations(
|
||||
self, citation_tokens, search_queries, reasoning_tokens
|
||||
):
|
||||
"""Test various combinations of citation tokens, search queries, and reasoning tokens."""
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
total_tokens=150,
|
||||
reasoning_tokens=reasoning_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
web_search_requests=search_queries
|
||||
),
|
||||
)
|
||||
|
||||
usage.citation_tokens = citation_tokens
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
# Calculate expected costs. `completion_tokens` includes `reasoning_tokens`,
|
||||
# so non-reasoning portion = 50 - reasoning_tokens.
|
||||
expected_prompt_cost = (100 * 2e-6) + (citation_tokens * 2e-6)
|
||||
expected_completion_cost = (
|
||||
((50 - reasoning_tokens) * 8e-6)
|
||||
+ (reasoning_tokens * 3e-6)
|
||||
+ (search_queries * 0.005)
|
||||
)
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
# Ensure costs are non-negative
|
||||
assert prompt_cost >= 0
|
||||
assert completion_cost >= 0
|
||||
|
||||
def test_uses_perplexity_provided_cost_when_available(self):
|
||||
"""
|
||||
Test that when Perplexity provides pre-calculated cost in usage.cost.total_cost,
|
||||
|
|
@ -374,9 +109,7 @@ class TestPerplexityCostCalculator:
|
|||
"total_cost": 0.008,
|
||||
}
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-pro", usage=usage
|
||||
)
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(model="sonar-pro", usage=usage)
|
||||
|
||||
# When Perplexity provides total_cost, we use it directly
|
||||
# prompt_cost should be 0, completion_cost should be total_cost
|
||||
|
|
@ -402,9 +135,7 @@ class TestPerplexityCostCalculator:
|
|||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
usage.cost = 0.008
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-pro", usage=usage
|
||||
)
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(model="sonar-pro", usage=usage)
|
||||
|
||||
assert prompt_cost == 0.0
|
||||
assert completion_cost == 0.008
|
||||
|
|
@ -417,9 +148,7 @@ class TestPerplexityCostCalculator:
|
|||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
# No cost object - should use manual calculation
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(model="sonar-deep-research", usage=usage)
|
||||
|
||||
# Should calculate manually: 100 * 2e-6 + 50 * 8e-6
|
||||
expected_prompt = 100 * 2e-6
|
||||
|
|
@ -428,57 +157,6 @@ class TestPerplexityCostCalculator:
|
|||
assert math.isclose(prompt_cost, expected_prompt, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost, expected_completion, rel_tol=1e-6)
|
||||
|
||||
def test_reasoning_tokens_not_double_billed(self):
|
||||
"""
|
||||
Regression: `completion_tokens` includes `reasoning_tokens` per the
|
||||
OpenAI/Perplexity usage convention (codified for the central path in PR #18607).
|
||||
When `output_cost_per_reasoning_token` is configured the manual fallback must
|
||||
subtract reasoning from completion before applying the output rate so the
|
||||
reasoning tokens are not billed at BOTH the output rate and the reasoning rate.
|
||||
|
||||
Uses the exact usage shape produced by the live response fixture in
|
||||
`tests/llm_translation/test_perplexity_reasoning.py`.
|
||||
"""
|
||||
usage = Usage(
|
||||
prompt_tokens=9,
|
||||
completion_tokens=20,
|
||||
total_tokens=29,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
reasoning_tokens=15
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = perplexity_cost_per_token(
|
||||
model="sonar-deep-research", usage=usage
|
||||
)
|
||||
|
||||
# sonar-deep-research rates: input 2e-6, output 8e-6, reasoning 3e-6.
|
||||
# Non-reasoning portion of the 20 completion tokens = 20 - 15 = 5.
|
||||
# Pre-fix this asserted 20 * 8e-6 + 15 * 3e-6 = 2.05e-4 (a 2.16x overcharge).
|
||||
expected_prompt = 9 * 2e-6
|
||||
expected_completion = (20 - 15) * 8e-6 + 15 * 3e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt, rel_tol=1e-9)
|
||||
assert math.isclose(completion_cost, expected_completion, rel_tol=1e-9)
|
||||
|
||||
def test_agent_api_fallback_rates_price_a_response_without_metered_cost(self):
|
||||
"""Perplexity meters cost on the response, but when `usage.cost` is absent the
|
||||
calculator falls back to the mapped per-token rates. Regression: that fallback
|
||||
raised "This model isn't mapped yet" for every Agent API third-party model,
|
||||
because the doubled cost-map key was unreachable from the resolution ladder.
|
||||
"""
|
||||
from litellm import ModelResponse
|
||||
|
||||
response = ModelResponse()
|
||||
response.usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
|
||||
response.model = "perplexity/perplexity/glm-5.2"
|
||||
|
||||
total_cost = completion_cost(
|
||||
completion_response=response, custom_llm_provider="perplexity"
|
||||
)
|
||||
|
||||
assert math.isclose(total_cost, 1000 * 1.4e-06 + 500 * 4.4e-06, rel_tol=1e-9)
|
||||
|
||||
OFF_PEAK_MODEL = "sonar-off-peak-test"
|
||||
OFF_PEAK_WINDOW = "14:00-00:00"
|
||||
INSIDE_WINDOW = datetime(2026, 9, 3, 17, 25, tzinfo=timezone.utc)
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
"""
|
||||
Integration tests for Perplexity cost calculation and transformation.
|
||||
|
||||
Tests the end-to-end functionality of Perplexity cost calculation
|
||||
Tests the end-to-end functionality of Perplexity cost calculation
|
||||
including integration with the main LiteLLM cost calculator.
|
||||
"""
|
||||
|
||||
|
|
@ -12,10 +12,9 @@ import os
|
|||
import pytest
|
||||
|
||||
# Add the project root to Python path
|
||||
|
||||
import litellm
|
||||
from litellm import ModelResponse
|
||||
from litellm.cost_calculator import completion_cost, cost_per_token
|
||||
from litellm.cost_calculator import cost_per_token
|
||||
from litellm.llms.perplexity.chat.transformation import PerplexityChatConfig
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
|
||||
from litellm.utils import get_model_info
|
||||
|
|
@ -57,109 +56,9 @@ class TestPerplexityIntegration:
|
|||
}
|
||||
}
|
||||
|
||||
def test_end_to_end_cost_calculation_with_transformation(self):
|
||||
"""Test end-to-end cost calculation with response transformation."""
|
||||
# Create a Perplexity API response that includes citations and search queries
|
||||
config = PerplexityChatConfig()
|
||||
|
||||
# Create a ModelResponse with basic usage (before transformation)
|
||||
model_response = ModelResponse()
|
||||
model_response.model = "sonar-deep-research"
|
||||
model_response.usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=50,
|
||||
total_tokens=150,
|
||||
reasoning_tokens=10,
|
||||
)
|
||||
|
||||
# Simulate raw response from Perplexity API
|
||||
raw_response_dict = {
|
||||
"choices": [{"message": {"content": "Test response with citations"}}],
|
||||
"usage": {
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 50,
|
||||
"total_tokens": 150,
|
||||
"num_search_queries": 2,
|
||||
},
|
||||
"citations": [
|
||||
"This is the first citation with important information about the topic",
|
||||
"Another citation providing additional context for the response",
|
||||
],
|
||||
}
|
||||
|
||||
# Apply transformation to extract Perplexity-specific fields
|
||||
config._enhance_usage_with_perplexity_fields(model_response, raw_response_dict)
|
||||
|
||||
# Now calculate the cost with the enhanced usage
|
||||
total_cost = completion_cost(
|
||||
completion_response=model_response, custom_llm_provider="perplexity"
|
||||
)
|
||||
|
||||
# Calculate expected cost
|
||||
citation_chars = sum(
|
||||
len(citation) for citation in raw_response_dict["citations"]
|
||||
)
|
||||
citation_tokens = citation_chars // 4
|
||||
|
||||
expected_prompt_cost = (100 * 2e-6) + (citation_tokens * 2e-6)
|
||||
expected_completion_cost = (
|
||||
((50 - 10) * 8e-6) + (10 * 3e-6) + (2 * 0.005)
|
||||
) # Output (text) + reasoning + search
|
||||
expected_total = expected_prompt_cost + expected_completion_cost
|
||||
|
||||
assert math.isclose(total_cost, expected_total, rel_tol=1e-6)
|
||||
|
||||
def test_cost_calculation_without_custom_fields(self):
|
||||
"""Test that cost calculation works normally when custom fields are absent."""
|
||||
# Create a standard response without Perplexity-specific fields
|
||||
model_response = ModelResponse()
|
||||
model_response.model = "sonar-deep-research"
|
||||
model_response.usage = Usage(
|
||||
prompt_tokens=100, completion_tokens=50, total_tokens=150
|
||||
)
|
||||
|
||||
# Calculate cost without custom fields
|
||||
total_cost = completion_cost(
|
||||
completion_response=model_response, custom_llm_provider="perplexity"
|
||||
)
|
||||
|
||||
# Should only include basic input/output costs
|
||||
expected_cost = (100 * 2e-6) + (50 * 8e-6)
|
||||
|
||||
assert math.isclose(total_cost, expected_cost, rel_tol=1e-6)
|
||||
|
||||
def test_main_cost_calculator_integration(self):
|
||||
"""Test integration with the main LiteLLM cost calculator."""
|
||||
# Create usage with all Perplexity fields
|
||||
usage = Usage(
|
||||
prompt_tokens=200,
|
||||
completion_tokens=100,
|
||||
total_tokens=300,
|
||||
reasoning_tokens=25,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(web_search_requests=3),
|
||||
)
|
||||
usage.citation_tokens = 40
|
||||
|
||||
# Test main cost calculator
|
||||
prompt_cost, completion_cost_val = cost_per_token(
|
||||
model="sonar-deep-research",
|
||||
custom_llm_provider="perplexity",
|
||||
usage_object=usage,
|
||||
)
|
||||
|
||||
expected_prompt_cost = (200 * 2e-6) + (40 * 2e-6)
|
||||
expected_completion_cost = (
|
||||
((100 - 25) * 8e-6) + (25 * 3e-6) + (3 * 0.005)
|
||||
) # Output (text) + reasoning + search
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_model_info_includes_custom_fields(self):
|
||||
"""Test that get_model_info returns the custom Perplexity cost fields."""
|
||||
model_info = get_model_info(
|
||||
model="sonar-deep-research", custom_llm_provider="perplexity"
|
||||
)
|
||||
model_info = get_model_info(model="sonar-deep-research", custom_llm_provider="perplexity")
|
||||
|
||||
# Verify custom fields are included
|
||||
required_fields = [
|
||||
|
|
@ -192,9 +91,7 @@ class TestPerplexityIntegration:
|
|||
for citations, expected_approx_tokens in test_cases:
|
||||
model_response = ModelResponse()
|
||||
model_response.model = "sonar-deep-research"
|
||||
model_response.usage = Usage(
|
||||
prompt_tokens=100, completion_tokens=50, total_tokens=150
|
||||
)
|
||||
model_response.usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
|
||||
raw_response_dict = {
|
||||
"usage": {
|
||||
|
|
@ -205,9 +102,7 @@ class TestPerplexityIntegration:
|
|||
"citations": citations,
|
||||
}
|
||||
|
||||
config._enhance_usage_with_perplexity_fields(
|
||||
model_response, raw_response_dict
|
||||
)
|
||||
config._enhance_usage_with_perplexity_fields(model_response, raw_response_dict)
|
||||
|
||||
citation_tokens = getattr(model_response.usage, "citation_tokens", 0)
|
||||
|
||||
|
|
@ -217,55 +112,6 @@ class TestPerplexityIntegration:
|
|||
else:
|
||||
assert abs(citation_tokens - expected_approx_tokens) <= 5
|
||||
|
||||
def test_cost_calculation_with_zero_values(self):
|
||||
"""Test cost calculation handles zero values for custom fields correctly."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
|
||||
# Set custom fields to zero
|
||||
usage.citation_tokens = 0
|
||||
usage.prompt_tokens_details = PromptTokensDetailsWrapper(web_search_requests=0)
|
||||
|
||||
# Should not add any extra cost
|
||||
prompt_cost, completion_cost_val = cost_per_token(
|
||||
model="sonar-deep-research",
|
||||
custom_llm_provider="perplexity",
|
||||
usage_object=usage,
|
||||
)
|
||||
|
||||
expected_prompt_cost = 100 * 2e-6
|
||||
expected_completion_cost = 50 * 8e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6)
|
||||
assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6)
|
||||
|
||||
def test_high_volume_cost_calculation(self):
|
||||
"""Test cost calculation with high token and query counts."""
|
||||
usage = Usage(
|
||||
prompt_tokens=50000,
|
||||
completion_tokens=25000,
|
||||
total_tokens=75000,
|
||||
reasoning_tokens=10000,
|
||||
)
|
||||
|
||||
usage.citation_tokens = 5000
|
||||
usage.prompt_tokens_details = PromptTokensDetailsWrapper(
|
||||
web_search_requests=100
|
||||
)
|
||||
|
||||
total_cost = completion_cost(
|
||||
completion_response=ModelResponse(usage=usage, model="sonar-deep-research"),
|
||||
custom_llm_provider="perplexity",
|
||||
)
|
||||
|
||||
expected_prompt_cost = (50000 * 2e-6) + (5000 * 2e-6)
|
||||
expected_completion_cost = (
|
||||
((25000 - 10000) * 8e-6) + (10000 * 3e-6) + (100 * 0.005)
|
||||
) # $0.65
|
||||
expected_total = expected_prompt_cost + expected_completion_cost # $0.76
|
||||
|
||||
assert math.isclose(total_cost, expected_total, rel_tol=1e-6)
|
||||
assert total_cost > 0.25
|
||||
|
||||
def test_transformation_preserves_existing_usage_fields(self):
|
||||
"""Test that transformation doesn't overwrite existing standard usage fields."""
|
||||
config = PerplexityChatConfig()
|
||||
|
|
@ -305,9 +151,7 @@ class TestPerplexityIntegration:
|
|||
assert hasattr(model_response.usage, "citation_tokens")
|
||||
assert model_response.usage.prompt_tokens_details.web_search_requests == 3
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_name", ["perplexity", "PERPLEXITY", "Perplexity"]
|
||||
)
|
||||
@pytest.mark.parametrize("provider_name", ["perplexity", "PERPLEXITY", "Perplexity"])
|
||||
def test_case_insensitive_provider_matching(self, provider_name):
|
||||
"""Test that cost calculation works with different case variations of provider name."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
|
|
|
|||
|
|
@ -1,29 +0,0 @@
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.tencent.cost_calculator import cost_per_token
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
|
||||
|
||||
def test_cost_per_token_uses_tencent_model_pricing(local_model_cost_map):
|
||||
usage = Usage(prompt_tokens=1000, completion_tokens=2000, total_tokens=3000)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="tencent/deepseek-v4-pro", usage=usage)
|
||||
|
||||
assert prompt_cost == pytest.approx(1000 * 4.35e-07)
|
||||
assert completion_cost == pytest.approx(2000 * 8.7e-07)
|
||||
|
||||
|
||||
def test_top_level_dispatcher_routes_tencent_to_wrapper(local_model_cost_map):
|
||||
from litellm.cost_calculator import cost_per_token as dispatch_cost_per_token
|
||||
|
||||
prompt_cost, completion_cost = dispatch_cost_per_token(
|
||||
model="tencent/deepseek-v4-pro",
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=1000,
|
||||
custom_llm_provider="tencent",
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(1000 * 4.35e-07)
|
||||
assert completion_cost == pytest.approx(1000 * 8.7e-07)
|
||||
|
|
@ -9,7 +9,95 @@ import litellm
|
|||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import (
|
||||
VertexPassthroughLoggingHandler,
|
||||
)
|
||||
from litellm.types.utils import PassthroughCallTypes
|
||||
from litellm.types.utils import ModelResponse, PassthroughCallTypes
|
||||
|
||||
_OMNI_INTERACTIONS_USAGE: Final = {
|
||||
"total_tokens": 4041,
|
||||
"total_input_tokens": 12,
|
||||
"input_tokens_by_modality": [{"modality": "text", "tokens": 12}],
|
||||
"total_output_tokens": 4009,
|
||||
"output_tokens_by_modality": [
|
||||
{"modality": "text", "tokens": 9},
|
||||
{"modality": "video", "tokens": 4000},
|
||||
],
|
||||
"total_tool_use_tokens": 0,
|
||||
"total_thought_tokens": 20,
|
||||
}
|
||||
|
||||
|
||||
def test_interactions_create_response_logs_modality_usage_and_cost() -> None:
|
||||
"""
|
||||
Regression for LIT-6896: gemini-omni Interactions passthrough rows were logged
|
||||
with zero tokens and zero spend. Input, text-output and video-output tokens
|
||||
must land in usage, priced with the model's per-modality rates, and the
|
||||
response id must stay the litellm_call_id so SpendLogs keep their request_id.
|
||||
"""
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
logging_obj.optional_params = {}
|
||||
logging_obj.litellm_call_id = "call-6896"
|
||||
response = httpx.Response(
|
||||
status_code=200,
|
||||
json={
|
||||
"id": "interactions/abc",
|
||||
"model": "gemini-omni-flash-preview",
|
||||
"status": "completed",
|
||||
"outputs": [{"type": "text", "text": "hi"}],
|
||||
"usage": _OMNI_INTERACTIONS_USAGE,
|
||||
},
|
||||
)
|
||||
|
||||
result = VertexPassthroughLoggingHandler.vertex_passthrough_handler(
|
||||
httpx_response=response,
|
||||
logging_obj=logging_obj,
|
||||
url_route="https://aiplatform.googleapis.com/v1beta1/projects/p/locations/global/interactions",
|
||||
result=response.text,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
request_body={"model": "gemini-omni-flash-preview", "input": [{"type": "text", "text": "say hi"}]},
|
||||
)
|
||||
|
||||
model_response = result["result"]
|
||||
assert isinstance(model_response, ModelResponse)
|
||||
assert model_response.id == "call-6896"
|
||||
usage = model_response.usage
|
||||
assert usage.prompt_tokens == 12
|
||||
assert usage.completion_tokens == 4009 + 20
|
||||
assert usage.completion_tokens_details.text_tokens == 9
|
||||
assert usage.completion_tokens_details.video_tokens == 4000
|
||||
|
||||
model_info = litellm.get_model_info(model="gemini-omni-flash-preview", custom_llm_provider="vertex_ai")
|
||||
expected_cost = (
|
||||
12 * model_info["input_cost_per_token"]
|
||||
+ (9 + 20) * model_info["output_cost_per_token"]
|
||||
+ 4000 * model_info["output_cost_per_video_token"]
|
||||
)
|
||||
assert result["kwargs"]["response_cost"] == pytest.approx(expected_cost)
|
||||
assert result["kwargs"]["custom_llm_provider"] == "vertex_ai"
|
||||
assert logging_obj.model_call_details["model"] == "gemini-omni-flash-preview"
|
||||
assert logging_obj.model_call_details["custom_llm_provider"] == "vertex_ai"
|
||||
|
||||
|
||||
def test_interactions_response_without_usage_falls_back_to_generic_logging() -> None:
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
logging_obj.optional_params = {}
|
||||
response = httpx.Response(status_code=200, json={"id": "interactions/abc", "status": "in_progress"})
|
||||
|
||||
result = VertexPassthroughLoggingHandler.vertex_passthrough_handler(
|
||||
httpx_response=response,
|
||||
logging_obj=logging_obj,
|
||||
url_route="https://aiplatform.googleapis.com/v1beta1/projects/p/locations/global/interactions",
|
||||
result=response.text,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
request_body={"agent": "projects/p/locations/global/reasoningEngines/1"},
|
||||
)
|
||||
|
||||
assert result["result"] is None
|
||||
assert "response_cost" not in result["kwargs"]
|
||||
|
||||
|
||||
def test_lyria_predict_response_preserves_audio_response_and_logs_cost(
|
||||
|
|
|
|||
|
|
@ -3,8 +3,6 @@ import json
|
|||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.experimental_pass_through.transformation import (
|
||||
VertexAIPartnerModelsAnthropicMessagesConfig,
|
||||
)
|
||||
|
|
@ -23,12 +21,8 @@ def test_validate_environment_uses_vertex_ai_location():
|
|||
optional_params = {}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("token", "test-project")
|
||||
),
|
||||
patch.object(
|
||||
config, "get_complete_vertex_url", return_value="https://mock-url"
|
||||
) as mock_get_url,
|
||||
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
|
||||
patch.object(config, "get_complete_vertex_url", return_value="https://mock-url") as mock_get_url,
|
||||
):
|
||||
config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
|
|
@ -51,17 +45,11 @@ def test_web_search_header_added_for_messages_endpoint():
|
|||
"vertex_credentials": "{}",
|
||||
}
|
||||
# Include web search tool in optional_params
|
||||
optional_params = {
|
||||
"tools": [{"type": "web_search_20250305", "name": "web_search", "max_uses": 5}]
|
||||
}
|
||||
optional_params = {"tools": [{"type": "web_search_20250305", "name": "web_search", "max_uses": 5}]}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("token", "test-project")
|
||||
),
|
||||
patch.object(
|
||||
config, "get_complete_vertex_url", return_value="https://mock-url"
|
||||
),
|
||||
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
|
||||
patch.object(config, "get_complete_vertex_url", return_value="https://mock-url"),
|
||||
):
|
||||
updated_headers, api_base = config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
|
|
@ -73,12 +61,10 @@ def test_web_search_header_added_for_messages_endpoint():
|
|||
)
|
||||
|
||||
# Assert that the anthropic-beta header with web-search is present
|
||||
assert (
|
||||
"anthropic-beta" in updated_headers
|
||||
), "anthropic-beta header should be present"
|
||||
assert (
|
||||
updated_headers["anthropic-beta"] == "web-search-2025-03-05"
|
||||
), f"anthropic-beta should be 'web-search-2025-03-05', got: {updated_headers['anthropic-beta']}"
|
||||
assert "anthropic-beta" in updated_headers, "anthropic-beta header should be present"
|
||||
assert updated_headers["anthropic-beta"] == "web-search-2025-03-05", (
|
||||
f"anthropic-beta should be 'web-search-2025-03-05', got: {updated_headers['anthropic-beta']}"
|
||||
)
|
||||
|
||||
|
||||
def test_web_search_header_not_added_without_tool():
|
||||
|
|
@ -94,12 +80,8 @@ def test_web_search_header_not_added_without_tool():
|
|||
optional_params = {}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("token", "test-project")
|
||||
),
|
||||
patch.object(
|
||||
config, "get_complete_vertex_url", return_value="https://mock-url"
|
||||
),
|
||||
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
|
||||
patch.object(config, "get_complete_vertex_url", return_value="https://mock-url"),
|
||||
):
|
||||
updated_headers, api_base = config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
|
|
@ -111,9 +93,9 @@ def test_web_search_header_not_added_without_tool():
|
|||
)
|
||||
|
||||
# Assert that the anthropic-beta header is NOT present when no web search tool
|
||||
assert (
|
||||
"anthropic-beta" not in updated_headers
|
||||
), "anthropic-beta header should not be present without web search tool"
|
||||
assert "anthropic-beta" not in updated_headers, (
|
||||
"anthropic-beta header should not be present without web search tool"
|
||||
)
|
||||
|
||||
|
||||
def test_compact_context_management_header_added():
|
||||
|
|
@ -129,12 +111,8 @@ def test_compact_context_management_header_added():
|
|||
optional_params = {"context_management": {"edits": [{"type": "compact_20260112"}]}}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("token", "test-project")
|
||||
),
|
||||
patch.object(
|
||||
config, "get_complete_vertex_url", return_value="https://mock-url"
|
||||
),
|
||||
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
|
||||
patch.object(config, "get_complete_vertex_url", return_value="https://mock-url"),
|
||||
):
|
||||
updated_headers, api_base = config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
|
|
@ -146,12 +124,10 @@ def test_compact_context_management_header_added():
|
|||
)
|
||||
|
||||
# Assert that the anthropic-beta header with compact-2026-01-12 is present
|
||||
assert (
|
||||
"anthropic-beta" in updated_headers
|
||||
), "anthropic-beta header should be present"
|
||||
assert (
|
||||
"compact-2026-01-12" in updated_headers["anthropic-beta"]
|
||||
), f"anthropic-beta should contain 'compact-2026-01-12', got: {updated_headers['anthropic-beta']}"
|
||||
assert "anthropic-beta" in updated_headers, "anthropic-beta header should be present"
|
||||
assert "compact-2026-01-12" in updated_headers["anthropic-beta"], (
|
||||
f"anthropic-beta should contain 'compact-2026-01-12', got: {updated_headers['anthropic-beta']}"
|
||||
)
|
||||
|
||||
|
||||
def test_context_management_header_added_for_other_edits():
|
||||
|
|
@ -167,12 +143,8 @@ def test_context_management_header_added_for_other_edits():
|
|||
optional_params = {"context_management": {"edits": [{"type": "some_other_type"}]}}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("token", "test-project")
|
||||
),
|
||||
patch.object(
|
||||
config, "get_complete_vertex_url", return_value="https://mock-url"
|
||||
),
|
||||
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
|
||||
patch.object(config, "get_complete_vertex_url", return_value="https://mock-url"),
|
||||
):
|
||||
updated_headers, api_base = config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
|
|
@ -184,12 +156,10 @@ def test_context_management_header_added_for_other_edits():
|
|||
)
|
||||
|
||||
# Assert that the anthropic-beta header with context-management-2025-06-27 is present
|
||||
assert (
|
||||
"anthropic-beta" in updated_headers
|
||||
), "anthropic-beta header should be present"
|
||||
assert (
|
||||
"context-management-2025-06-27" in updated_headers["anthropic-beta"]
|
||||
), f"anthropic-beta should contain 'context-management-2025-06-27', got: {updated_headers['anthropic-beta']}"
|
||||
assert "anthropic-beta" in updated_headers, "anthropic-beta header should be present"
|
||||
assert "context-management-2025-06-27" in updated_headers["anthropic-beta"], (
|
||||
f"anthropic-beta should contain 'context-management-2025-06-27', got: {updated_headers['anthropic-beta']}"
|
||||
)
|
||||
|
||||
|
||||
def test_both_compact_and_context_management_headers_added():
|
||||
|
|
@ -202,19 +172,11 @@ def test_both_compact_and_context_management_headers_added():
|
|||
"vertex_credentials": "{}",
|
||||
}
|
||||
# Include context_management with both compact and other edit types
|
||||
optional_params = {
|
||||
"context_management": {
|
||||
"edits": [{"type": "compact_20260112"}, {"type": "some_other_type"}]
|
||||
}
|
||||
}
|
||||
optional_params = {"context_management": {"edits": [{"type": "compact_20260112"}, {"type": "some_other_type"}]}}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("token", "test-project")
|
||||
),
|
||||
patch.object(
|
||||
config, "get_complete_vertex_url", return_value="https://mock-url"
|
||||
),
|
||||
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
|
||||
patch.object(config, "get_complete_vertex_url", return_value="https://mock-url"),
|
||||
):
|
||||
updated_headers, api_base = config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
|
|
@ -226,15 +188,13 @@ def test_both_compact_and_context_management_headers_added():
|
|||
)
|
||||
|
||||
# Assert that both beta headers are present
|
||||
assert (
|
||||
"anthropic-beta" in updated_headers
|
||||
), "anthropic-beta header should be present"
|
||||
assert (
|
||||
"compact-2026-01-12" in updated_headers["anthropic-beta"]
|
||||
), f"anthropic-beta should contain 'compact-2026-01-12', got: {updated_headers['anthropic-beta']}"
|
||||
assert (
|
||||
"context-management-2025-06-27" in updated_headers["anthropic-beta"]
|
||||
), f"anthropic-beta should contain 'context-management-2025-06-27', got: {updated_headers['anthropic-beta']}"
|
||||
assert "anthropic-beta" in updated_headers, "anthropic-beta header should be present"
|
||||
assert "compact-2026-01-12" in updated_headers["anthropic-beta"], (
|
||||
f"anthropic-beta should contain 'compact-2026-01-12', got: {updated_headers['anthropic-beta']}"
|
||||
)
|
||||
assert "context-management-2025-06-27" in updated_headers["anthropic-beta"], (
|
||||
f"anthropic-beta should contain 'context-management-2025-06-27', got: {updated_headers['anthropic-beta']}"
|
||||
)
|
||||
|
||||
|
||||
def test_validate_environment_always_refreshes_token_ignoring_stale_bearer():
|
||||
|
|
@ -248,12 +208,8 @@ def test_validate_environment_always_refreshes_token_ignoring_stale_bearer():
|
|||
}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("fresh-token", "test-project")
|
||||
) as mock_ensure,
|
||||
patch.object(
|
||||
config, "get_complete_vertex_url", return_value="https://mock-vertex-url"
|
||||
),
|
||||
patch.object(config, "_ensure_access_token", return_value=("fresh-token", "test-project")) as mock_ensure,
|
||||
patch.object(config, "get_complete_vertex_url", return_value="https://mock-vertex-url"),
|
||||
):
|
||||
updated_headers, api_base = config.validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
|
|
@ -286,9 +242,7 @@ def test_validate_environment_appends_stream_raw_predict_with_custom_api_base():
|
|||
"get_complete_vertex_url",
|
||||
wraps=config.get_complete_vertex_url,
|
||||
) as spy_get_url,
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("token", "test-project")
|
||||
),
|
||||
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
|
||||
):
|
||||
_, api_base = config.validate_anthropic_messages_environment(
|
||||
headers={},
|
||||
|
|
@ -318,9 +272,7 @@ def test_validate_environment_appends_raw_predict_with_custom_api_base():
|
|||
"get_complete_vertex_url",
|
||||
wraps=config.get_complete_vertex_url,
|
||||
) as spy_get_url,
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("token", "test-project")
|
||||
),
|
||||
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
|
||||
):
|
||||
_, api_base = config.validate_anthropic_messages_environment(
|
||||
headers={},
|
||||
|
|
@ -447,20 +399,14 @@ def test_validate_environment_does_not_mutate_caller_headers():
|
|||
caller_headers: dict = {}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
config, "_ensure_access_token", return_value=("token", "test-project")
|
||||
),
|
||||
patch.object(
|
||||
config, "get_complete_vertex_url", return_value="https://mock-url"
|
||||
),
|
||||
patch.object(config, "_ensure_access_token", return_value=("token", "test-project")),
|
||||
patch.object(config, "get_complete_vertex_url", return_value="https://mock-url"),
|
||||
):
|
||||
config.validate_anthropic_messages_environment(
|
||||
headers=caller_headers,
|
||||
model="claude-sonnet-4",
|
||||
messages=[],
|
||||
optional_params={
|
||||
"tools": [{"type": "web_search_20250305", "name": "web_search"}]
|
||||
},
|
||||
optional_params={"tools": [{"type": "web_search_20250305", "name": "web_search"}]},
|
||||
litellm_params={
|
||||
"vertex_ai_project": "p",
|
||||
"vertex_ai_location": "us-central1",
|
||||
|
|
@ -468,9 +414,7 @@ def test_validate_environment_does_not_mutate_caller_headers():
|
|||
api_base=None,
|
||||
)
|
||||
|
||||
assert (
|
||||
caller_headers == {}
|
||||
), "validate_anthropic_messages_environment must not mutate the caller's headers dict"
|
||||
assert caller_headers == {}, "validate_anthropic_messages_environment must not mutate the caller's headers dict"
|
||||
|
||||
|
||||
def test_vertex_claude_completion_does_not_mutate_shared_extra_headers():
|
||||
|
|
@ -483,12 +427,8 @@ def test_vertex_claude_completion_does_not_mutate_shared_extra_headers():
|
|||
mock_response = MagicMock()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
handler, "_ensure_access_token", return_value=("ya29.fresh", "proj")
|
||||
),
|
||||
patch.object(
|
||||
handler, "get_complete_vertex_url", return_value="https://mock-url"
|
||||
),
|
||||
patch.object(handler, "_ensure_access_token", return_value=("ya29.fresh", "proj")),
|
||||
patch.object(handler, "get_complete_vertex_url", return_value="https://mock-url"),
|
||||
patch(
|
||||
"litellm.llms.anthropic.chat.AnthropicChatCompletion.completion",
|
||||
return_value=mock_response,
|
||||
|
|
@ -509,10 +449,7 @@ def test_vertex_claude_completion_does_not_mutate_shared_extra_headers():
|
|||
litellm_params={},
|
||||
)
|
||||
|
||||
assert (
|
||||
shared_extra_headers == {}
|
||||
), "extra_headers must not be mutated by completion()"
|
||||
|
||||
assert shared_extra_headers == {}, "extra_headers must not be mutated by completion()"
|
||||
|
||||
|
||||
def test_messages_thinking_shape_follows_exact_vertex_entry_flag(local_model_cost_map, monkeypatch):
|
||||
|
|
@ -541,9 +478,7 @@ def test_messages_thinking_shape_follows_exact_vertex_entry_flag(local_model_cos
|
|||
assert result.get("thinking") == {"type": "adaptive", "display": "summarized"}
|
||||
assert result.get("output_config") == {"effort": "medium"}
|
||||
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost["vertex_ai/claude-opus-4-8"], "supports_adaptive_thinking", False
|
||||
)
|
||||
monkeypatch.setitem(litellm.model_cost["vertex_ai/claude-opus-4-8"], "supports_adaptive_thinking", False)
|
||||
litellm.get_model_info.cache_clear()
|
||||
assert litellm.model_cost["claude-opus-4-8"]["supports_adaptive_thinking"] is True
|
||||
|
||||
|
|
@ -614,9 +549,7 @@ class TestVertexAnthropicMidConversationSystem:
|
|||
{"role": "assistant", "content": "reading"},
|
||||
{"role": "user", "content": "continue"},
|
||||
]
|
||||
result = _vertex_transform(
|
||||
"claude-sonnet-4-6", messages, system=[{"type": "text", "text": "Base."}]
|
||||
)
|
||||
result = _vertex_transform("claude-sonnet-4-6", messages, system=[{"type": "text", "text": "Base."}])
|
||||
assert result["messages"] == [
|
||||
{"role": "user", "content": "read the file"},
|
||||
{
|
||||
|
|
@ -660,9 +593,7 @@ def test_vertex_claude_4_8_plus_cost_map_entries_carry_mid_conversation_system_f
|
|||
|
||||
import litellm
|
||||
|
||||
cost_map_path = os.path.join(
|
||||
os.path.dirname(litellm.__file__), "model_prices_and_context_window_backup.json"
|
||||
)
|
||||
cost_map_path = os.path.join(os.path.dirname(litellm.__file__), "model_prices_and_context_window_backup.json")
|
||||
with open(cost_map_path) as f:
|
||||
cost_map = json.load(f)
|
||||
rules = cost_map["fallback_generalizations"]["rules"]
|
||||
|
|
|
|||
|
|
@ -717,6 +717,33 @@ class TestVertexAIVideoConfig:
|
|||
assert video_obj.usage["duration_seconds"] == 8.0
|
||||
assert video_obj.usage["video_resolution"] == "1080p"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"sample_count,expected_video_count",
|
||||
[(2, 2), (1, 1), (None, None), (0, None), ("2", None)],
|
||||
ids=["two", "one", "unset", "zero", "string"],
|
||||
)
|
||||
def test_transform_video_create_response_usage_includes_video_count(self, sample_count, expected_video_count):
|
||||
"""Regression for LIT-6896: sampleCount is the number of generated videos and must reach usage for billing."""
|
||||
mock_response = Mock(spec=httpx.Response)
|
||||
mock_response.json.return_value = {
|
||||
"name": "projects/p/locations/us-central1/publishers/google/models/veo-3.1-fast-generate-001/operations/op-1"
|
||||
}
|
||||
parameters = {"durationSeconds": 4, "resolution": "720p"}
|
||||
if sample_count is not None:
|
||||
parameters["sampleCount"] = sample_count
|
||||
|
||||
video_obj = self.config.transform_video_create_response(
|
||||
model="veo-3.1-fast-generate-001",
|
||||
raw_response=mock_response,
|
||||
logging_obj=self.mock_logging_obj,
|
||||
custom_llm_provider="vertex_ai",
|
||||
request_data={"instances": [{"prompt": "a red ball"}], "parameters": parameters},
|
||||
)
|
||||
|
||||
assert video_obj.usage is not None
|
||||
assert video_obj.usage["duration_seconds"] == 4.0
|
||||
assert video_obj.usage.get("video_count") == expected_video_count
|
||||
|
||||
def test_transform_video_remix_request_not_supported(self):
|
||||
"""Test that video remix raises NotImplementedError."""
|
||||
with pytest.raises(NotImplementedError, match="Video remix is not supported"):
|
||||
|
|
|
|||
|
|
@ -10,9 +10,7 @@ Source: litellm/llms/xai/responses/transformation.py
|
|||
from unittest.mock import MagicMock, Mock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.xai.cost_calculator import cost_per_token
|
||||
from litellm.llms.xai.responses.transformation import XAIResponsesAPIConfig
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils
|
||||
|
|
@ -366,12 +364,16 @@ class TestXAIResponsesWebSearchBilling:
|
|||
|
||||
def _raw_response_json(self, include_web_search: bool) -> dict:
|
||||
web_search_output = (
|
||||
[{
|
||||
"type": "web_search_call",
|
||||
"id": "ws_1",
|
||||
"status": "completed",
|
||||
"action": {"type": "search", "query": "grok"},
|
||||
}] if include_web_search else []
|
||||
[
|
||||
{
|
||||
"type": "web_search_call",
|
||||
"id": "ws_1",
|
||||
"status": "completed",
|
||||
"action": {"type": "search", "query": "grok"},
|
||||
}
|
||||
]
|
||||
if include_web_search
|
||||
else []
|
||||
)
|
||||
tool_usage = {"server_side_tool_usage_details": self._TOOL_DETAILS} if include_web_search else {}
|
||||
return {
|
||||
|
|
@ -431,20 +433,6 @@ class TestXAIResponsesWebSearchBilling:
|
|||
assert bridged.completion_tokens == 20
|
||||
assert getattr(bridged, "server_side_tool_usage_details") == self._TOOL_DETAILS
|
||||
|
||||
def test_completion_cost_bills_web_search_calls(self):
|
||||
with_search = litellm.completion_cost(
|
||||
completion_response=self._transform(include_web_search=True),
|
||||
model="xai/grok-4",
|
||||
custom_llm_provider="xai",
|
||||
)
|
||||
without_search = litellm.completion_cost(
|
||||
completion_response=self._transform(include_web_search=False),
|
||||
model="xai/grok-4",
|
||||
custom_llm_provider="xai",
|
||||
)
|
||||
|
||||
assert with_search - without_search == pytest.approx(2 * 5.0 / 1000.0)
|
||||
|
||||
def test_streaming_terminal_event_keeps_schema_and_details(self):
|
||||
parsed_chunk = {
|
||||
"type": "response.completed",
|
||||
|
|
@ -535,9 +523,7 @@ class TestXAIResponsesReportedCost:
|
|||
assert cost_per_token(model="grok-4-latest", usage=chat_usage) == (0.0, 0.0037756)
|
||||
|
||||
def test_usage_without_a_reported_cost_is_left_alone(self):
|
||||
usage = self._transformed_usage(
|
||||
{"input_tokens": 100, "output_tokens": 200, "total_tokens": 300}
|
||||
)
|
||||
usage = self._transformed_usage({"input_tokens": 100, "output_tokens": 200, "total_tokens": 300})
|
||||
|
||||
assert usage.cost is None
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
from unittest.mock import Mock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.xai.chat.transformation import (
|
||||
|
|
@ -26,11 +25,7 @@ class TestXAIReasoningTokenFolding:
|
|||
total_tokens: int,
|
||||
reasoning_tokens: int = 0,
|
||||
) -> ModelResponse:
|
||||
details = (
|
||||
CompletionTokensDetailsWrapper(reasoning_tokens=reasoning_tokens)
|
||||
if reasoning_tokens
|
||||
else None
|
||||
)
|
||||
details = CompletionTokensDetailsWrapper(reasoning_tokens=reasoning_tokens) if reasoning_tokens else None
|
||||
usage = Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
|
|
@ -194,31 +189,11 @@ class TestXAIChatWebSearchBilling:
|
|||
def test_enhance_noop_without_details(self):
|
||||
response = self._response_with_usage()
|
||||
|
||||
XAIChatConfig()._enhance_usage_with_xai_web_search_fields(
|
||||
response, {"usage": {"prompt_tokens": 100}}
|
||||
)
|
||||
XAIChatConfig()._enhance_usage_with_xai_web_search_fields(response, {"usage": {"prompt_tokens": 100}})
|
||||
|
||||
assert response.usage.prompt_tokens_details is None
|
||||
assert getattr(response.usage, "server_side_tool_usage_details", None) is None
|
||||
|
||||
def test_completion_cost_bills_chat_web_search_calls(self):
|
||||
billed = self._response_with_usage()
|
||||
XAIChatConfig()._enhance_usage_with_xai_web_search_fields(
|
||||
billed,
|
||||
{"usage": {"server_side_tool_usage_details": self._TOOL_DETAILS}},
|
||||
)
|
||||
|
||||
with_search = litellm.completion_cost(
|
||||
completion_response=billed, model="xai/grok-4", custom_llm_provider="xai"
|
||||
)
|
||||
without_search = litellm.completion_cost(
|
||||
completion_response=self._response_with_usage(),
|
||||
model="xai/grok-4",
|
||||
custom_llm_provider="xai",
|
||||
)
|
||||
|
||||
assert with_search - without_search == pytest.approx(3 * 5.0 / 1000.0)
|
||||
|
||||
|
||||
class TestXAIReportedCost:
|
||||
"""xAI reports what it charged; the transformation moves it to where litellm bills from.
|
||||
|
|
@ -275,9 +250,7 @@ class TestXAIReportedCost:
|
|||
assert cost_per_token(model="grok-4-latest", usage=usage) == (0.0, 0.0037756)
|
||||
|
||||
def test_usage_without_a_reported_cost_is_left_alone(self):
|
||||
usage = self._transformed_usage(
|
||||
{"prompt_tokens": 100, "completion_tokens": 200, "total_tokens": 300}
|
||||
)
|
||||
usage = self._transformed_usage({"prompt_tokens": 100, "completion_tokens": 200, "total_tokens": 300})
|
||||
|
||||
assert getattr(usage, "cost", None) is None
|
||||
|
||||
|
|
@ -300,9 +273,7 @@ class TestXAIReportedCost:
|
|||
Chunk aggregation rebuilds usage from the fields it models plus ``cost``, so a
|
||||
chunk still carrying only ``cost_in_usd_ticks`` loses the reported amount.
|
||||
"""
|
||||
handler = XAIChatCompletionStreamingHandler(
|
||||
streaming_response=iter([]), sync_stream=True
|
||||
)
|
||||
handler = XAIChatCompletionStreamingHandler(streaming_response=iter([]), sync_stream=True)
|
||||
|
||||
parsed = handler.chunk_parser(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -6,16 +6,6 @@ import math
|
|||
import os
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
CompletionTokensDetailsWrapper,
|
||||
Message,
|
||||
ModelResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
|
||||
StandardBuiltInToolCostTracking,
|
||||
)
|
||||
|
|
@ -26,6 +16,13 @@ from litellm.llms.xai.cost_calculator import (
|
|||
cost_per_token,
|
||||
cost_per_web_search_request,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
Message,
|
||||
ModelResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
|
||||
class TestXAICostCalculator:
|
||||
|
|
@ -45,241 +42,6 @@ class TestXAICostCalculator:
|
|||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
def test_basic_cost_calculation(self):
|
||||
"""Test basic cost calculation without reasoning tokens."""
|
||||
usage = Usage(prompt_tokens=12, completion_tokens=125, total_tokens=137)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-3-mini", usage=usage)
|
||||
|
||||
# Expected costs for grok-3-mini:
|
||||
# Input: 12 tokens * $3e-7 = $0.0000036
|
||||
# Output: 125 tokens * $5e-7 = $0.0000625
|
||||
expected_prompt_cost = 12 * 1.25e-6
|
||||
expected_completion_cost = 125 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_reasoning_tokens_cost_calculation(self):
|
||||
"""Test cost calculation with reasoning tokens from completion_tokens_details."""
|
||||
usage = Usage(
|
||||
prompt_tokens=12,
|
||||
completion_tokens=125,
|
||||
total_tokens=1086,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=949,
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=None, # Not set, but doesn't matter for XAI billing
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-3-mini", usage=usage)
|
||||
|
||||
# Expected costs for grok-3-mini:
|
||||
# Input: 12 tokens * $3e-7 = $0.0000036
|
||||
# Completion: (125 + 949) tokens * $5e-7 = $0.000537
|
||||
expected_prompt_cost = 12 * 1.25e-6
|
||||
expected_completion_cost = (125 + 949) * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_reasoning_and_text_tokens_cost_calculation(self):
|
||||
"""Test cost calculation with both reasoning and text tokens."""
|
||||
usage = Usage(
|
||||
prompt_tokens=12,
|
||||
completion_tokens=125,
|
||||
total_tokens=1086,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=949,
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=76, # Explicitly set (but ignored in XAI billing)
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-3-mini", usage=usage)
|
||||
|
||||
# Expected costs for grok-3-mini:
|
||||
# Input: 12 tokens * $3e-7 = $0.0000036
|
||||
# Completion: (125 + 949) tokens * $5e-7 = $0.000537
|
||||
# Note: text_tokens field is ignored, only completion_tokens + reasoning_tokens matters
|
||||
expected_prompt_cost = 12 * 1.25e-6
|
||||
expected_completion_cost = (125 + 949) * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_grok_4_cost_calculation(self):
|
||||
"""Test cost calculation for grok-4 model."""
|
||||
usage = Usage(
|
||||
prompt_tokens=10,
|
||||
completion_tokens=200,
|
||||
total_tokens=360,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=150,
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=50, # Ignored in XAI billing
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-4", usage=usage)
|
||||
|
||||
# grok-4 was retired on 2026-05-15 and now redirects to grok-4.3, so it bills
|
||||
# at grok-4.3's rates:
|
||||
# Input: 10 tokens * $1.25e-6
|
||||
# Completion: (200 + 150) tokens * $2.5e-6
|
||||
expected_prompt_cost = 10 * 1.25e-6
|
||||
expected_completion_cost = (200 + 150) * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_grok_3_fast_beta_cost_calculation(self):
|
||||
"""Test cost calculation for grok-3-fast-beta model."""
|
||||
usage = Usage(
|
||||
prompt_tokens=20,
|
||||
completion_tokens=300,
|
||||
total_tokens=520,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=200,
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=100, # Ignored in XAI billing
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="grok-3-fast-beta", usage=usage
|
||||
)
|
||||
|
||||
# Expected costs for grok-3-fast-beta:
|
||||
# Input: 20 tokens * $5e-6 = $0.0001
|
||||
# Completion: (300 + 200) tokens * $2.5e-5 = $0.0125
|
||||
expected_prompt_cost = 20 * 1.25e-6
|
||||
expected_completion_cost = (300 + 200) * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
|
||||
def test_edge_case_large_reasoning_tokens(self):
|
||||
"""Test cost calculation when reasoning_tokens is larger than completion_tokens."""
|
||||
usage = Usage(
|
||||
prompt_tokens=12,
|
||||
completion_tokens=50, # Less than reasoning_tokens
|
||||
total_tokens=162,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=100, # More than completion_tokens
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=None,
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-3-mini", usage=usage)
|
||||
|
||||
# Expected costs:
|
||||
# Input: 12 tokens * $3e-7 = $0.0000036
|
||||
# Completion: (50 + 100) tokens * $5e-7 = $0.000075
|
||||
expected_prompt_cost = 12 * 1.25e-6
|
||||
expected_completion_cost = (50 + 100) * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_tiered_pricing_above_200k_tokens(self):
|
||||
usage = Usage(
|
||||
prompt_tokens=250000,
|
||||
completion_tokens=100000,
|
||||
total_tokens=400000,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=50000,
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=None,
|
||||
),
|
||||
)
|
||||
prompt_cost, completion_cost = cost_per_token(model="xai/grok-4.3", usage=usage)
|
||||
expected_prompt_cost = 250000 * 2.5e-6
|
||||
expected_completion_cost = (100000 + 50000) * 5e-6
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_tiered_pricing_below_200k_tokens(self):
|
||||
usage = Usage(
|
||||
prompt_tokens=100000,
|
||||
completion_tokens=50000,
|
||||
total_tokens=160000,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=10000,
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=None,
|
||||
),
|
||||
)
|
||||
prompt_cost, completion_cost = cost_per_token(model="xai/grok-4.3", usage=usage)
|
||||
expected_prompt_cost = 100000 * 1.25e-6
|
||||
expected_completion_cost = (50000 + 10000) * 2.5e-6
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_tiered_pricing_grok_4_latest(self):
|
||||
"""Test tiered pricing for grok-4-latest model."""
|
||||
usage = Usage(
|
||||
prompt_tokens=250000, # Above the 200k threshold
|
||||
completion_tokens=100000,
|
||||
total_tokens=400000,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=50000,
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=None,
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="xai/grok-4-latest", usage=usage
|
||||
)
|
||||
|
||||
# grok-4-latest redirects to grok-4.3, which tiers at 200k rather than 128k:
|
||||
# Input: 250000 tokens * $2.5e-6 (ALL tokens at tiered rate since input > 200k)
|
||||
# Completion: (100000 + 50000) tokens * $5e-6 (tiered rate since input > 200k)
|
||||
expected_prompt_cost = 250000 * 2.5e-6
|
||||
expected_completion_cost = (100000 + 50000) * 5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_tiered_pricing_output_tokens_below_200k(self):
|
||||
usage = Usage(
|
||||
prompt_tokens=250000,
|
||||
completion_tokens=50000,
|
||||
total_tokens=310000,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=10000,
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=None,
|
||||
),
|
||||
)
|
||||
prompt_cost, completion_cost = cost_per_token(model="xai/grok-4.3", usage=usage)
|
||||
expected_prompt_cost = 250000 * 2.5e-6
|
||||
expected_completion_cost = (50000 + 10000) * 5e-6
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_tiered_pricing_model_without_tiered_pricing(self):
|
||||
litellm.model_cost["xai/flat-rate-fixture"] = {
|
||||
"input_cost_per_token": 3e-7,
|
||||
|
|
@ -294,29 +56,6 @@ class TestXAICostCalculator:
|
|||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_already_normalised_usage_does_not_double_count_reasoning(self):
|
||||
"""Cost calc must not double-bill when Usage is already OpenAI-normalised."""
|
||||
usage = Usage(
|
||||
prompt_tokens=12,
|
||||
completion_tokens=200,
|
||||
total_tokens=212,
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
accepted_prediction_tokens=0,
|
||||
audio_tokens=0,
|
||||
reasoning_tokens=100,
|
||||
rejected_prediction_tokens=0,
|
||||
text_tokens=None,
|
||||
),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="grok-3-mini", usage=usage)
|
||||
|
||||
expected_prompt_cost = 12 * 1.25e-6
|
||||
expected_completion_cost = 200 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_web_search_cost_via_server_side_tool_usage_details(self):
|
||||
"""usage.server_side_tool_usage_details.web_search_calls at default $5/1k."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
|
|
@ -344,9 +83,7 @@ class TestXAICostCalculator:
|
|||
"search_context_size_medium": 0.01,
|
||||
}
|
||||
}
|
||||
web_search_cost = cost_per_web_search_request(
|
||||
usage=usage, model_info=model_info
|
||||
)
|
||||
web_search_cost = cost_per_web_search_request(usage=usage, model_info=model_info)
|
||||
assert math.isclose(web_search_cost, 0.02, rel_tol=1e-10)
|
||||
|
||||
def test_web_search_cost_zero_without_details(self):
|
||||
|
|
@ -355,9 +92,7 @@ class TestXAICostCalculator:
|
|||
|
||||
def test_apply_details_sets_web_search_requests_for_cost_gate(self):
|
||||
usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
|
||||
apply_server_side_tool_usage_details_to_usage(
|
||||
usage, {"web_search_calls": 2, "x_search_calls": 0}
|
||||
)
|
||||
apply_server_side_tool_usage_details_to_usage(usage, {"web_search_calls": 2, "x_search_calls": 0})
|
||||
assert usage.prompt_tokens_details is not None
|
||||
assert usage.prompt_tokens_details.web_search_requests == 2
|
||||
assert StandardBuiltInToolCostTracking.response_object_includes_web_search_call(
|
||||
|
|
@ -413,9 +148,7 @@ class TestXAICostCalculator:
|
|||
|
||||
assert get_cost_for_web_search_request("xai", usage, {}) > 0.0
|
||||
|
||||
reported = Usage(
|
||||
prompt_tokens=100, completion_tokens=50, total_tokens=150, cost=0.0037756
|
||||
)
|
||||
reported = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150, cost=0.0037756)
|
||||
setattr(reported, "server_side_tool_usage_details", {"web_search_calls": 3})
|
||||
assert get_cost_for_web_search_request("xai", reported, {}) == 0.0
|
||||
|
||||
|
|
@ -503,82 +236,6 @@ class TestXAICostCalculator:
|
|||
|
||||
assert cost_per_token(model="grok-4-latest", usage=usage) == (0.0, 0.0)
|
||||
|
||||
def test_grok_4_20_beta_reasoning_cost_calculation(self):
|
||||
"""Test cost calculation for grok-4.20-beta-0309-reasoning model."""
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=200, total_tokens=300)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="grok-4.20-beta-0309-reasoning", usage=usage
|
||||
)
|
||||
|
||||
# Input: 100 tokens * $1.25e-6 = $0.000125
|
||||
# Output: 200 tokens * $2.5e-6 = $0.0005
|
||||
expected_prompt_cost = 100 * 1.25e-6
|
||||
expected_completion_cost = 200 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_grok_4_20_beta_non_reasoning_cost_calculation(self):
|
||||
"""Test cost calculation for grok-4.20-beta-0309-non-reasoning model."""
|
||||
usage = Usage(prompt_tokens=50, completion_tokens=100, total_tokens=150)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="grok-4.20-beta-0309-non-reasoning", usage=usage
|
||||
)
|
||||
|
||||
# Input: 50 tokens * $1.25e-6 = $0.0000625
|
||||
# Output: 100 tokens * $2.5e-6 = $0.00025
|
||||
expected_prompt_cost = 50 * 1.25e-6
|
||||
expected_completion_cost = 100 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_grok_4_20_at_exactly_200k_prompt_tokens_uses_higher_tier(self):
|
||||
"""xAI bills the >=200k tier once the prompt reaches 200k, so the boundary is inclusive."""
|
||||
usage = Usage(prompt_tokens=200_000, completion_tokens=1_000, total_tokens=201_000)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="grok-4.20-0309-reasoning", usage=usage
|
||||
)
|
||||
|
||||
expected_prompt_cost = 200_000 * 2.5e-6
|
||||
expected_completion_cost = 1_000 * 5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_grok_4_20_just_below_200k_prompt_tokens_uses_base_tier(self):
|
||||
"""One token under the boundary still bills at the base rates."""
|
||||
usage = Usage(prompt_tokens=199_999, completion_tokens=1_000, total_tokens=200_999)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="grok-4.20-0309-reasoning", usage=usage
|
||||
)
|
||||
|
||||
expected_prompt_cost = 199_999 * 1.25e-6
|
||||
expected_completion_cost = 1_000 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_grok_4_20_multi_agent_cost_calculation(self):
|
||||
"""Test cost calculation for grok-4.20-multi-agent-beta-0309 model."""
|
||||
usage = Usage(prompt_tokens=200, completion_tokens=300, total_tokens=500)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="grok-4.20-multi-agent-beta-0309", usage=usage
|
||||
)
|
||||
|
||||
# Input: 200 tokens * $1.25e-6 = $0.00025
|
||||
# Output: 300 tokens * $2.5e-6 = $0.00075
|
||||
expected_prompt_cost = 200 * 1.25e-6
|
||||
expected_completion_cost = 300 * 2.5e-6
|
||||
|
||||
assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-10)
|
||||
assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-10)
|
||||
|
||||
def test_custom_pricing_beats_the_reported_cost(self):
|
||||
response = ModelResponse(
|
||||
id="chatcmpl-xai",
|
||||
|
|
@ -635,10 +292,7 @@ class TestXAIWebSearchCostHelpers:
|
|||
details = {"web_search_calls": 0, "x_search_calls": 3}
|
||||
apply_server_side_tool_usage_details_to_usage(usage, details)
|
||||
assert getattr(usage, "server_side_tool_usage_details") == details
|
||||
assert (
|
||||
usage.prompt_tokens_details is None
|
||||
or usage.prompt_tokens_details.web_search_requests is None
|
||||
)
|
||||
assert usage.prompt_tokens_details is None or usage.prompt_tokens_details.web_search_requests is None
|
||||
|
||||
def test_apply_details_skips_mirror_when_web_search_calls_invalid(self):
|
||||
usage = Usage(prompt_tokens=1, completion_tokens=1, total_tokens=2)
|
||||
|
|
@ -660,10 +314,7 @@ class TestXAIWebSearchCostHelpers:
|
|||
assert usage.prompt_tokens_details.web_search_requests == 4
|
||||
|
||||
def test_web_search_cost_per_call_default_when_model_info_empty(self):
|
||||
assert (
|
||||
_web_search_cost_per_call_from_model_info({})
|
||||
== _DEFAULT_WEB_SEARCH_COST_PER_CALL
|
||||
)
|
||||
assert _web_search_cost_per_call_from_model_info({}) == _DEFAULT_WEB_SEARCH_COST_PER_CALL
|
||||
|
||||
def test_web_search_cost_per_call_prefers_medium_over_low(self):
|
||||
model_info = {
|
||||
|
|
|
|||
|
|
@ -13,19 +13,6 @@ REPO_ROOT = Path(__file__).parents[4]
|
|||
PRICES_PATH = REPO_ROOT / "model_prices_and_context_window.json"
|
||||
BACKUP_PRICES_PATH = REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json"
|
||||
|
||||
# Retired by xAI and no longer served: requests to these slugs 404 rather than
|
||||
# redirecting, and they are absent from https://docs.x.ai/docs/models
|
||||
RETIRED_MODELS = (
|
||||
"xai/grok-2",
|
||||
"xai/grok-2-1212",
|
||||
"xai/grok-2-latest",
|
||||
"xai/grok-2-vision",
|
||||
"xai/grok-2-vision-1212",
|
||||
"xai/grok-2-vision-latest",
|
||||
"xai/grok-beta",
|
||||
"xai/grok-vision-beta",
|
||||
)
|
||||
|
||||
# https://docs.x.ai/developers/model-capabilities/text/multi-agent
|
||||
# "The multi-agent model does not work with the OpenAI Chat Completions API."
|
||||
RESPONSES_ONLY_MODELS = (
|
||||
|
|
@ -42,17 +29,11 @@ def cost_map(request: pytest.FixtureRequest) -> dict:
|
|||
return json.loads(path.read_text(encoding="utf-8"))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", RETIRED_MODELS)
|
||||
def test_retired_xai_models_are_not_advertised(cost_map: dict, model: str):
|
||||
assert model not in cost_map
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", RESPONSES_ONLY_MODELS)
|
||||
def test_multi_agent_models_are_responses_only(cost_map: dict, model: str):
|
||||
entry = cost_map[model]
|
||||
assert entry["supported_endpoints"] == ["/v1/responses"]
|
||||
assert entry["mode"] == "responses"
|
||||
assert "/v1/chat/completions" not in entry["supported_endpoints"]
|
||||
|
||||
|
||||
def test_surviving_xai_chat_models_still_serve_chat_completions(cost_map: dict):
|
||||
|
|
@ -64,7 +45,6 @@ def test_surviving_xai_chat_models_still_serve_chat_completions(cost_map: dict):
|
|||
]
|
||||
assert "xai/grok-4.3" in chat_models
|
||||
assert "xai/grok-4.6" in chat_models
|
||||
assert not any(key.startswith("xai/grok-2") for key in chat_models)
|
||||
|
||||
|
||||
def test_both_cost_maps_agree_on_xai_entries():
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
get_key_mcp_rpm_limit,
|
||||
get_key_model_rpm_limit,
|
||||
get_key_model_tpm_limit,
|
||||
get_key_own_model_rate_limit,
|
||||
get_key_tag_rpm_limit,
|
||||
get_model_from_request,
|
||||
get_project_model_rpm_limit,
|
||||
|
|
@ -141,6 +142,35 @@ class TestLogOnceIfBudgetReservationDisabled:
|
|||
class TestGetKeyModelRpmLimit:
|
||||
"""Tests for get_key_model_rpm_limit function."""
|
||||
|
||||
def test_own_limit_excludes_team_metadata(self):
|
||||
"""A team-only limit is inherited, not owned: the key resolves it but does not override it."""
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-123",
|
||||
metadata={"some_other_key": "value"},
|
||||
team_metadata={"model_rpm_limit": {"gpt-4": 50}, "model_tpm_limit": {"gpt-4": 500}},
|
||||
)
|
||||
assert get_key_model_rpm_limit(user_api_key_dict) == {"gpt-4": 50}
|
||||
assert get_key_own_model_rate_limit(user_api_key_dict, "model_rpm_limit") is None
|
||||
assert get_key_own_model_rate_limit(user_api_key_dict, "model_tpm_limit") is None
|
||||
|
||||
def test_own_limit_resolves_metadata_then_model_max_budget(self):
|
||||
from_metadata = UserAPIKeyAuth(
|
||||
api_key="sk-123",
|
||||
metadata={"model_rpm_limit": {"gpt-4": 100}},
|
||||
model_max_budget={"gpt-4": {"rpm_limit": 10, "tpm_limit": 1000}},
|
||||
team_metadata={"model_rpm_limit": {"gpt-4": 50}},
|
||||
)
|
||||
assert get_key_own_model_rate_limit(from_metadata, "model_rpm_limit") == {"gpt-4": 100}
|
||||
assert get_key_own_model_rate_limit(from_metadata, "model_tpm_limit") == {"gpt-4": 1000}
|
||||
|
||||
from_budget = UserAPIKeyAuth(
|
||||
api_key="sk-123",
|
||||
model_max_budget={"gpt-4": {"rpm_limit": 10}, "gpt-3.5-turbo": {"tpm_limit": 1000}},
|
||||
team_metadata={"model_rpm_limit": {"gpt-4": 50}},
|
||||
)
|
||||
assert get_key_own_model_rate_limit(from_budget, "model_rpm_limit") == {"gpt-4": 10}
|
||||
assert get_key_own_model_rate_limit(from_budget, "model_tpm_limit") == {"gpt-3.5-turbo": 1000}
|
||||
|
||||
def test_returns_key_metadata_when_present(self):
|
||||
"""Key metadata takes priority over team metadata."""
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
|
|
@ -823,6 +853,82 @@ def test_get_model_from_request_azure_relay_routes_use_the_model_group_in_the_pa
|
|||
assert get_model_from_request(request_data=request_data, route=route, llm_router=_azure_relay_router()) == expected
|
||||
|
||||
|
||||
def _nvidia_nim_relay_router():
|
||||
from litellm.router import Router
|
||||
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "nim-page-elements",
|
||||
"litellm_params": {
|
||||
"model": "nvidia_nim/nvidia/nemoretriever-page-elements-v2",
|
||||
"api_base": "http://nim-a.internal:8000",
|
||||
"api_key": "k",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "nvidia/nemoretriever-table-structure-v1",
|
||||
"litellm_params": {
|
||||
"model": "nvidia_nim/nvidia/nemoretriever-table-structure-v1",
|
||||
"api_base": "http://nim-b.internal:8000",
|
||||
"api_key": "k",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "gpt-4o",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "k"},
|
||||
},
|
||||
{
|
||||
"model_name": "detect",
|
||||
"litellm_params": {
|
||||
"model": "nvidia_nim/nvidia/nemoretriever-page-elements-v2",
|
||||
"api_base": "http://nim-a.internal:8000",
|
||||
"api_key": "k",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "detect",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "k"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
NIM_INFER_BODY = {"input": [{"type": "image_url", "url": "data:image/png;base64,AAAA"}]}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route, request_data, expected",
|
||||
[
|
||||
("/nvidia_nim/nim-page-elements/v1/infer", NIM_INFER_BODY, "nim-page-elements"),
|
||||
(
|
||||
"/nvidia_nim/nim-page-elements/v1/infer",
|
||||
{"model": "nvidia/nemoretriever-table-structure-v1"},
|
||||
"nim-page-elements",
|
||||
),
|
||||
(
|
||||
"/nvidia_nim/nvidia/nemoretriever-table-structure-v1/v1/infer",
|
||||
NIM_INFER_BODY,
|
||||
"nvidia/nemoretriever-table-structure-v1",
|
||||
),
|
||||
("/nvidia_nim/v1/infer", NIM_INFER_BODY, None),
|
||||
("/nvidia_nim/unknown-group/v1/infer", NIM_INFER_BODY, None),
|
||||
("/nvidia_nim/nim-page-elements-v2/v1/infer", NIM_INFER_BODY, None),
|
||||
("/nvidia_nim/gpt-4o/v1/infer", NIM_INFER_BODY, None),
|
||||
("/nvidia_nim/detect/v1/infer", NIM_INFER_BODY, None),
|
||||
],
|
||||
)
|
||||
def test_get_model_from_request_nvidia_nim_relay_routes_use_the_model_group_in_the_path(route, request_data, expected):
|
||||
assert (
|
||||
get_model_from_request(request_data=request_data, route=route, llm_router=_nvidia_nim_relay_router())
|
||||
== expected
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_from_request_nvidia_nim_relay_without_a_router_has_no_model():
|
||||
assert get_model_from_request(request_data=NIM_INFER_BODY, route="/nvidia_nim/nim-page-elements/v1/infer") is None
|
||||
|
||||
|
||||
def test_get_model_from_request_includes_file_endpoint_header_model():
|
||||
assert (
|
||||
get_model_from_request(
|
||||
|
|
|
|||
|
|
@ -693,6 +693,7 @@ def test_virtual_key_allowed_routes_with_litellm_routes_member_name_denied():
|
|||
"/anthropic/v1/count_tokens",
|
||||
"/gemini/v1/models",
|
||||
"/gemini/countTokens",
|
||||
"/nvidia_nim/nim-page-elements/v1/infer",
|
||||
],
|
||||
)
|
||||
def test_virtual_key_llm_api_route_includes_passthrough_prefix(route):
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import json
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -266,6 +267,51 @@ async def test_get_all_transactions_from_redis_buffer_pipeline(redis_update_buff
|
|||
assert popped_keys[6] == REDIS_WINDOW_SPEND_UPDATE_BUFFER_KEY
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_org_member_spend_is_summed_across_pods_and_restored_on_rpush_failure(
|
||||
redis_update_buffer: RedisUpdateBuffer, mock_redis_cache: AsyncMock
|
||||
):
|
||||
from litellm.proxy._types import Litellm_EntityType
|
||||
from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
|
||||
DailySpendUpdateQueue,
|
||||
)
|
||||
from litellm.proxy.db.db_transaction_queue.spend_update_queue import (
|
||||
SpendUpdateQueue,
|
||||
)
|
||||
|
||||
member_key: Final = "organization_id::org-1::user_id::user-1"
|
||||
pod_json: Final = json.dumps({"org_member_list_transactions": {member_key: 0.25}})
|
||||
mock_redis_cache.async_lpop_pipeline = AsyncMock(
|
||||
return_value=[[pod_json, pod_json], None, None, None, None, None, None]
|
||||
)
|
||||
|
||||
(db_spend, *_rest) = await redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline()
|
||||
|
||||
assert db_spend is not None
|
||||
assert db_spend["org_member_list_transactions"] == {member_key: 0.5}
|
||||
|
||||
mock_redis_cache.async_rpush_pipeline = AsyncMock(side_effect=ConnectionError("redis went away"))
|
||||
spend_queue: Final = SpendUpdateQueue()
|
||||
await spend_queue.add_update(
|
||||
{
|
||||
"entity_type": Litellm_EntityType.ORGANIZATION_MEMBER,
|
||||
"entity_id": member_key,
|
||||
"response_cost": 1.5,
|
||||
}
|
||||
)
|
||||
await redis_update_buffer.store_in_memory_spend_updates_in_redis(
|
||||
spend_update_queue=spend_queue,
|
||||
daily_spend_update_queue=DailySpendUpdateQueue(),
|
||||
daily_team_spend_update_queue=DailySpendUpdateQueue(),
|
||||
daily_org_spend_update_queue=DailySpendUpdateQueue(),
|
||||
daily_end_user_spend_update_queue=DailySpendUpdateQueue(),
|
||||
daily_agent_spend_update_queue=DailySpendUpdateQueue(),
|
||||
)
|
||||
|
||||
restored_spend: Final = await spend_queue.flush_and_get_aggregated_db_spend_update_transactions()
|
||||
assert restored_spend["org_member_list_transactions"] == {member_key: 1.5}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_all_transactions_from_redis_buffer_pipeline_no_redis():
|
||||
"""When redis_cache is None, should return all Nones"""
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from collections.abc import Callable
|
|||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, call, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -944,6 +945,121 @@ async def test_commit_spend_updates_to_db_increments_team_member_spend_and_total
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_org_spend_increments_organization_membership_row_for_the_calling_user():
|
||||
"""A request made with a user_id inside an org must increment that user's
|
||||
LiteLLM_OrganizationMembership.spend, not only the org total, or the
|
||||
Organizations > Members UI renders '-' for every member."""
|
||||
db_writer: Final = DBSpendUpdateWriter()
|
||||
await db_writer._update_org_db(
|
||||
response_cost=0.75,
|
||||
org_id="org-abc",
|
||||
user_id="user-xyz",
|
||||
prisma_client=MagicMock(),
|
||||
)
|
||||
transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions()
|
||||
|
||||
mock_batcher: Final = MagicMock()
|
||||
mock_prisma_client: Final = MagicMock()
|
||||
mock_prisma_client.db.tx = MagicMock(return_value=_good_tx(mock_batcher))
|
||||
proxy_logging: Final = MagicMock()
|
||||
proxy_logging.call_details = {}
|
||||
|
||||
await db_writer._commit_spend_updates_to_db(
|
||||
prisma_client=mock_prisma_client,
|
||||
n_retry_times=0,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
db_spend_update_transactions=transactions,
|
||||
)
|
||||
|
||||
mock_batcher.litellm_organizationtable.update_many.assert_called_once_with(
|
||||
where={"organization_id": "org-abc"},
|
||||
data={"spend": {"increment": 0.75}},
|
||||
)
|
||||
mock_batcher.litellm_organizationmembership.update_many.assert_called_once_with(
|
||||
where={"organization_id": "org-abc", "user_id": "user-xyz"},
|
||||
data={"spend": {"increment": 0.75}},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_org_spend_without_user_id_leaves_organization_membership_untouched():
|
||||
db_writer: Final = DBSpendUpdateWriter()
|
||||
await db_writer._update_org_db(
|
||||
response_cost=0.75,
|
||||
org_id="org-abc",
|
||||
user_id=None,
|
||||
prisma_client=MagicMock(),
|
||||
)
|
||||
transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions()
|
||||
|
||||
mock_batcher: Final = MagicMock()
|
||||
mock_prisma_client: Final = MagicMock()
|
||||
mock_prisma_client.db.tx = MagicMock(return_value=_good_tx(mock_batcher))
|
||||
proxy_logging: Final = MagicMock()
|
||||
proxy_logging.call_details = {}
|
||||
|
||||
await db_writer._commit_spend_updates_to_db(
|
||||
prisma_client=mock_prisma_client,
|
||||
n_retry_times=0,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
db_spend_update_transactions=transactions,
|
||||
)
|
||||
|
||||
mock_batcher.litellm_organizationtable.update_many.assert_called_once()
|
||||
mock_batcher.litellm_organizationmembership.update_many.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_org_spend_keeps_member_attribution_when_ids_contain_the_key_delimiter():
|
||||
db_writer: Final = DBSpendUpdateWriter()
|
||||
await db_writer._update_org_db(
|
||||
response_cost=0.75,
|
||||
org_id="division::west",
|
||||
user_id="user::42",
|
||||
prisma_client=MagicMock(),
|
||||
)
|
||||
transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions()
|
||||
|
||||
mock_batcher: Final = MagicMock()
|
||||
mock_prisma_client: Final = MagicMock()
|
||||
mock_prisma_client.db.tx = MagicMock(return_value=_good_tx(mock_batcher))
|
||||
proxy_logging: Final = MagicMock()
|
||||
proxy_logging.call_details = {}
|
||||
|
||||
await db_writer._commit_spend_updates_to_db(
|
||||
prisma_client=mock_prisma_client,
|
||||
n_retry_times=0,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
db_spend_update_transactions=transactions,
|
||||
)
|
||||
|
||||
mock_batcher.litellm_organizationmembership.update_many.assert_called_once_with(
|
||||
where={"organization_id": "division::west", "user_id": "user::42"},
|
||||
data={"spend": {"increment": 0.75}},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_database_updates_queues_org_member_spend_for_the_request_user():
|
||||
db_writer: Final = DBSpendUpdateWriter()
|
||||
await db_writer._batch_database_updates(
|
||||
response_cost=0.1,
|
||||
user_id="u1",
|
||||
hashed_token="t1",
|
||||
team_id=None,
|
||||
org_id="org1",
|
||||
end_user_id=None,
|
||||
prisma_client=MagicMock(),
|
||||
litellm_proxy_budget_name=None,
|
||||
payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.1},
|
||||
)
|
||||
transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions()
|
||||
|
||||
assert transactions["org_list_transactions"] == {"org1": 0.1}
|
||||
assert transactions["org_member_list_transactions"] == {"organization_id::org1::user_id::u1": 0.1}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_spend_log_transaction_to_daily_tag_transaction_with_request_id():
|
||||
"""
|
||||
|
|
@ -2904,6 +3020,7 @@ async def test_update_daily_spend_retries_deadlock(monkeypatch):
|
|||
("team_list_transactions", "team-1"),
|
||||
("team_member_list_transactions", "team_id::team-1::user_id::user-1"),
|
||||
("org_list_transactions", "org-1"),
|
||||
("org_member_list_transactions", "organization_id::org-1::user_id::user-1"),
|
||||
("tag_list_transactions", "tag-1"),
|
||||
("agent_list_transactions", "agent-1"),
|
||||
],
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -250,6 +250,76 @@ async def test_custom_code_flag_default_reason_and_empty_metadata():
|
|||
}
|
||||
|
||||
|
||||
IDENTITY_ECHO_CODE = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
" return flag('identity', metadata={\n"
|
||||
" 'ids': [request_data['user_id'], request_data['team_id'], request_data['end_user_id']],\n"
|
||||
" 'metadata_keys': sorted(request_data['metadata'].keys()),\n"
|
||||
" })\n"
|
||||
)
|
||||
CALLER_IDENTITY = {
|
||||
"user_api_key_user_id": "someone@example.com",
|
||||
"user_api_key_team_id": "team-1",
|
||||
"user_api_key_end_user_id": "end-user-1",
|
||||
"user_api_key_alias": "guardrail-repro-key",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
|
||||
async def test_custom_code_sandbox_sees_caller_identity_from_proxy_metadata_bucket(metadata_key):
|
||||
"""LIT-6609: the proxy writes user_api_key_* into `metadata` (chat) or `litellm_metadata`
|
||||
(/v1/messages, responses, batches, files); the sandbox must resolve ids from either."""
|
||||
guardrail = _compile(IDENTITY_ECHO_CODE)
|
||||
request_data = {"model": "m", metadata_key: dict(CALLER_IDENTITY)}
|
||||
|
||||
await guardrail.apply_guardrail(inputs={"texts": ["x"]}, request_data=request_data, input_type="request")
|
||||
|
||||
entry = request_data[metadata_key]["standard_logging_guardrail_information"][0]
|
||||
assert entry["guardrail_response"]["metadata"] == {
|
||||
"ids": ["someone@example.com", "team-1", "end-user-1"],
|
||||
"metadata_keys": sorted(CALLER_IDENTITY),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_code_sandbox_merges_caller_metadata_with_litellm_metadata():
|
||||
"""On litellm_metadata routes the caller's own `metadata` field must stay visible next to
|
||||
the proxy identity block, and the proxy block wins on key collisions."""
|
||||
guardrail = _compile(IDENTITY_ECHO_CODE)
|
||||
request_data = {
|
||||
"model": "m",
|
||||
"metadata": {"trace_id": "abc", "user_api_key_user_id": "forged"},
|
||||
"litellm_metadata": dict(CALLER_IDENTITY),
|
||||
}
|
||||
|
||||
await guardrail.apply_guardrail(inputs={"texts": ["x"]}, request_data=request_data, input_type="request")
|
||||
|
||||
entry = request_data["litellm_metadata"]["standard_logging_guardrail_information"][0]
|
||||
assert entry["guardrail_response"]["metadata"] == {
|
||||
"ids": ["someone@example.com", "team-1", "end-user-1"],
|
||||
"metadata_keys": sorted([*CALLER_IDENTITY, "trace_id"]),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_code_sandbox_ignores_top_level_identity_fields():
|
||||
"""Only the proxy-owned metadata buckets carry identity; user_api_key_* keys at the top level
|
||||
of the request body are caller-controlled on ordinary routes and must never become ids."""
|
||||
code = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n"
|
||||
" ids = [request_data['user_id'], request_data['team_id'], request_data['end_user_id']]\n"
|
||||
" return flag('identity', metadata={'ids': str(ids)})\n"
|
||||
)
|
||||
guardrail = _compile(code)
|
||||
request_data = {"model": "m", **CALLER_IDENTITY, "metadata": {"headers": {}}}
|
||||
|
||||
await guardrail.apply_guardrail(inputs={"texts": ["x"]}, request_data=request_data, input_type="request")
|
||||
|
||||
entry = request_data["metadata"]["standard_logging_guardrail_information"][0]
|
||||
assert entry["guardrail_response"]["metadata"]["ids"] == "[None, None, None]"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_code_allow_still_records_success_not_flagged():
|
||||
code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()\n"
|
||||
|
|
|
|||
|
|
@ -6311,7 +6311,7 @@ async def test_an_open_circuit_breaker_reads_the_sliding_window_locally_without_
|
|||
(
|
||||
{
|
||||
"team_id": "t",
|
||||
"metadata": {"model_rpm_limit": {"test-model": 100}},
|
||||
"metadata": {"model_rpm_limit": {"other-model": 100}},
|
||||
"team_metadata": {"model_rpm_limit": {"test-model": 1}},
|
||||
},
|
||||
{},
|
||||
|
|
@ -6529,3 +6529,113 @@ async def test_request_capacity_rejection_keeps_existing_redis_mirror():
|
|||
pytest.fail("rejection released another request's mirrored slot")
|
||||
assert exc.value.status_code == 429
|
||||
assert await cache.async_get_cache(counter_key, local_only=True) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key_limits",
|
||||
[
|
||||
{"metadata": {"model_rpm_limit": {"test-model": 3}}},
|
||||
{"model_max_budget": {"test-model": {"rpm_limit": 3}}},
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_model_rpm_override_takes_precedence_over_team_model_rpm_limit(key_limits):
|
||||
cache = DualCache()
|
||||
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
auth = UserAPIKeyAuth(
|
||||
api_key=hash_token("sk-key-override"),
|
||||
team_id="t",
|
||||
team_metadata={"model_rpm_limit": {"test-model": 1}},
|
||||
**key_limits,
|
||||
)
|
||||
|
||||
async def request():
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=auth, cache=cache, data={"model": "test-model"}, call_type="acompletion"
|
||||
)
|
||||
|
||||
for _ in range(3):
|
||||
await request()
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await request()
|
||||
assert exc.value.status_code == 429
|
||||
assert "model_per_key" in str(exc.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key_limits, override_key_gets_through",
|
||||
[
|
||||
({"model_rpm_limit": {"test-model": 10}}, False),
|
||||
({"model_rpm_limit": {"test-model": 10}, "model_tpm_limit": {"test-model": 5000}}, True),
|
||||
],
|
||||
ids=["rpm_only_override_still_shares_team_tpm", "rpm_and_tpm_override_leaves_team_tpm"],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_model_rpm_override_keeps_team_model_tpm_limit(key_limits, override_key_gets_through):
|
||||
cache = DualCache()
|
||||
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(cache))
|
||||
team_metadata = {"model_rpm_limit": {"test-model": 5}, "model_tpm_limit": {"test-model": 500}}
|
||||
sibling_key = UserAPIKeyAuth(api_key=hash_token("sk-sibling"), team_id="t", team_metadata=team_metadata)
|
||||
override_key = UserAPIKeyAuth(
|
||||
api_key=hash_token("sk-key-override"), team_id="t", metadata=key_limits, team_metadata=team_metadata
|
||||
)
|
||||
|
||||
async def request(auth):
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=auth,
|
||||
cache=cache,
|
||||
data={"model": "test-model", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 300},
|
||||
call_type="acompletion",
|
||||
)
|
||||
|
||||
await request(sibling_key)
|
||||
if override_key_gets_through:
|
||||
await request(override_key)
|
||||
return
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await request(override_key)
|
||||
assert exc.value.status_code == 429
|
||||
assert "model_per_team" in str(exc.value.detail)
|
||||
assert exc.value.headers["rate_limit_type"] == "tokens"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"key_metadata, charges_team_model_pool",
|
||||
[
|
||||
({}, True),
|
||||
({"model_rpm_limit": {"test-model": 10}}, True),
|
||||
({"model_tpm_limit": {"test-model": 5000}}, False),
|
||||
({"model_tpm_limit": {"other-model": 5000}}, True),
|
||||
],
|
||||
ids=["no_override", "rpm_only_override", "tpm_override", "tpm_override_on_other_model"],
|
||||
)
|
||||
def test_success_tpm_accounting_skips_team_model_pool_when_key_owns_model_tpm_limit(
|
||||
key_metadata, charges_team_model_pool
|
||||
):
|
||||
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
response = ModelResponse(
|
||||
id="team-pool-tpm",
|
||||
object="chat.completion",
|
||||
created=int(datetime.now().timestamp()),
|
||||
model="test-model",
|
||||
usage=Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150),
|
||||
choices=[],
|
||||
)
|
||||
kwargs = {
|
||||
"standard_logging_object": {"metadata": {"user_api_key_hash": hash_token("sk-pool"), "user_api_key_team_id": "t"}},
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"model_group": "test-model",
|
||||
"user_api_key_metadata": key_metadata,
|
||||
"user_api_key_team_metadata": {"model_tpm_limit": {"test-model": 500}},
|
||||
}
|
||||
},
|
||||
"model": "test-model",
|
||||
}
|
||||
|
||||
ops = handler._build_success_event_pipeline_operations(kwargs=kwargs, response_obj=response, rate_limit_type="output")
|
||||
|
||||
charged_keys = {op["key"] for op in ops}
|
||||
assert handler.create_rate_limit_keys("model_per_key", f"{hash_token('sk-pool')}:test-model", "tokens") in charged_keys
|
||||
team_pool_key = handler.create_rate_limit_keys("model_per_team", "t:test-model", "tokens")
|
||||
assert (team_pool_key in charged_keys) is charges_team_model_pool
|
||||
|
|
|
|||
|
|
@ -4982,6 +4982,26 @@ class TestStrategyRouterWriteValidation:
|
|||
|
||||
_V2 = {"classifier_type": "heuristic_v2", "tiers": {"SIMPLE": "gpt-4o-mini"}}
|
||||
_V1 = {"classifier_type": "heuristic", "tiers": {"SIMPLE": "gpt-4o-mini"}}
|
||||
_FORECAST_BASE = {
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini"},
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini", "REASONING": "gpt-4o"},
|
||||
}
|
||||
_CAPABILITY = {
|
||||
**_FORECAST_BASE,
|
||||
"classifier_type": "capability",
|
||||
"capability_classifier_config": {
|
||||
"efficient_tier": "SIMPLE", "capable_tier": "REASONING", "base_threshold": 0.7,
|
||||
},
|
||||
}
|
||||
_FUSE = {
|
||||
**_FORECAST_BASE,
|
||||
"classifier_type": "llm_v2",
|
||||
"adaptive": False,
|
||||
"llm_v2_config": {
|
||||
"efficient_profile": "Small solver", "capable_profile": "Large solver",
|
||||
"harness": "One attempt", "max_quality_gap": 0.05,
|
||||
},
|
||||
}
|
||||
_CUSTOM_TIERS = {
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini"},
|
||||
|
|
@ -5049,6 +5069,16 @@ class TestStrategyRouterWriteValidation:
|
|||
@pytest.mark.parametrize(
|
||||
"limit,effective_params,db_models,config_config,model_id,expected",
|
||||
[
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _CAPABILITY}, ["auto_router/complexity_router"], None, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _CAPABILITY}, [], _CAPABILITY, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _CAPABILITY}, [], _FUSE, None, "reserved"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _CAPABILITY}, [], None, "held-id", "reserved"),
|
||||
(None, {"model": "auto_router/complexity_router", "complexity_router_config": _CAPABILITY}, ["auto_router/complexity_router"], _CAPABILITY, None, "plain"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _FUSE}, ["auto_router/complexity_router"], None, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _FUSE}, [], _FUSE, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _FUSE}, [], _CAPABILITY, None, "reserved"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _FUSE}, [], None, "held-id", "reserved"),
|
||||
(None, {"model": "auto_router/complexity_router", "complexity_router_config": _FUSE}, ["auto_router/complexity_router"], _FUSE, None, "plain"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _V2}, ["auto_router/complexity_router"], None, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _V2}, [], _V2, None, "refused"),
|
||||
(1, {"model": "auto_router/complexity_router", "complexity_router_config": _V2}, [], None, None, "reserved"),
|
||||
|
|
@ -5338,7 +5368,8 @@ class TestStrategyRouterWriteValidation:
|
|||
assert events == ["slot-enter", "slot-exit", "team_model_add"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_new_model_refuses_a_second_heuristic_v2_router_before_the_db_write(self) -> None:
|
||||
@pytest.mark.parametrize("config", [_V2, _CAPABILITY, _FUSE])
|
||||
async def test_add_new_model_refuses_a_second_gated_classifier_router_before_the_db_write(self, config: Mapping[str, object]) -> None:
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
add_new_model,
|
||||
|
|
@ -5366,7 +5397,7 @@ class TestStrategyRouterWriteValidation:
|
|||
await add_new_model(
|
||||
model_params=Deployment(
|
||||
model_name="second-v2",
|
||||
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=self._V2),
|
||||
litellm_params=LiteLLM_Params(model="auto_router/complexity_router", complexity_router_config=config),
|
||||
),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
|
|
@ -5463,7 +5494,8 @@ class TestStrategyRouterWriteValidation:
|
|||
assert fake.litellm_proxymodeltable.update.await_count == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_patch_model_refuses_switching_another_router_to_heuristic_v2(self) -> None:
|
||||
@pytest.mark.parametrize("config", [_V2, _CAPABILITY, _FUSE])
|
||||
async def test_patch_model_refuses_switching_another_router_to_gated_classifier(self, config: Mapping[str, object]) -> None:
|
||||
"""patch_model relays HTTPException as-is, so the license refusal reaches the client as a plain 403."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -5498,7 +5530,7 @@ class TestStrategyRouterWriteValidation:
|
|||
with pytest.raises(HTTPException) as exc_info:
|
||||
await patch_model(
|
||||
model_id=model_id,
|
||||
patch_data=updateDeployment(litellm_params=updateLiteLLMParams(complexity_router_config=self._V2)),
|
||||
patch_data=updateDeployment(litellm_params=updateLiteLLMParams(complexity_router_config=config)),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
|
@ -5506,7 +5538,8 @@ class TestStrategyRouterWriteValidation:
|
|||
fake.litellm_proxymodeltable.update.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_model_refuses_switching_another_router_to_heuristic_v2(self) -> None:
|
||||
@pytest.mark.parametrize("config", [_V2, _CAPABILITY, _FUSE])
|
||||
async def test_update_model_refuses_switching_another_router_to_gated_classifier(self, config: Mapping[str, object]) -> None:
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
update_model,
|
||||
|
|
@ -5542,7 +5575,7 @@ class TestStrategyRouterWriteValidation:
|
|||
with pytest.raises(ProxyException) as exc_info:
|
||||
await update_model(
|
||||
model_params=updateDeployment(
|
||||
litellm_params=updateLiteLLMParams(complexity_router_config=self._V2),
|
||||
litellm_params=updateLiteLLMParams(complexity_router_config=config),
|
||||
model_info=ModelInfo(id=model_id),
|
||||
),
|
||||
user_api_key_dict=admin,
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.gemini_passthrough_logging_handler import (
|
||||
GeminiPassthroughLoggingHandler,
|
||||
|
|
@ -397,3 +397,52 @@ class TestGeminiPassthroughLoggingHandler:
|
|||
assert mock_logging_obj.model_call_details["response_cost"] == expected_cost
|
||||
assert mock_logging_obj.model_call_details["model"] == "veo-2.0-generate-001"
|
||||
assert mock_logging_obj.model_call_details["custom_llm_provider"] == "gemini"
|
||||
|
||||
def test_interactions_create_response_is_priced_as_gemini(self):
|
||||
"""Regression for LIT-6896: Gemini API Interactions passthrough must not log zero usage."""
|
||||
usage = {
|
||||
"total_tokens": 1030,
|
||||
"total_input_tokens": 10,
|
||||
"input_tokens_by_modality": [{"modality": "text", "tokens": 10}],
|
||||
"total_output_tokens": 1020,
|
||||
"output_tokens_by_modality": [
|
||||
{"modality": "text", "tokens": 20},
|
||||
{"modality": "video", "tokens": 1000},
|
||||
],
|
||||
"total_tool_use_tokens": 0,
|
||||
"total_thought_tokens": 0,
|
||||
}
|
||||
mock_httpx_response = MagicMock(spec=httpx.Response)
|
||||
mock_httpx_response.json.return_value = {
|
||||
"id": "interactions/abc",
|
||||
"model": "gemini-omni-flash-preview",
|
||||
"status": "completed",
|
||||
"usage": usage,
|
||||
}
|
||||
mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj)
|
||||
mock_logging_obj.model_call_details = {}
|
||||
mock_logging_obj.litellm_call_id = "call-6896"
|
||||
|
||||
result = GeminiPassthroughLoggingHandler.gemini_passthrough_handler(
|
||||
httpx_response=mock_httpx_response,
|
||||
response_body=mock_httpx_response.json.return_value,
|
||||
logging_obj=mock_logging_obj,
|
||||
url_route="https://generativelanguage.googleapis.com/v1beta/interactions",
|
||||
result="",
|
||||
start_time=self.start_time,
|
||||
end_time=self.end_time,
|
||||
cache_hit=False,
|
||||
request_body={"model": "gemini-omni-flash-preview", "input": "make a clip"},
|
||||
)
|
||||
|
||||
model_info = litellm.get_model_info(model="gemini-omni-flash-preview", custom_llm_provider="gemini")
|
||||
expected_cost = (
|
||||
10 * model_info["input_cost_per_token"]
|
||||
+ 20 * model_info["output_cost_per_token"]
|
||||
+ 1000 * model_info["output_cost_per_video_token"]
|
||||
)
|
||||
assert result["result"].id == "call-6896"
|
||||
assert result["result"].usage.completion_tokens_details.video_tokens == 1000
|
||||
assert result["kwargs"]["response_cost"] == pytest.approx(expected_cost)
|
||||
assert result["kwargs"]["custom_llm_provider"] == "gemini"
|
||||
assert mock_logging_obj.model_call_details["custom_llm_provider"] == "gemini"
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
|||
llm_passthrough_factory_proxy_route,
|
||||
milvus_proxy_route,
|
||||
mistral_proxy_route,
|
||||
relay_nvidia_nim_request,
|
||||
openai_proxy_route,
|
||||
vertex_discovery_proxy_route,
|
||||
vertex_proxy_route,
|
||||
|
|
@ -5375,6 +5376,186 @@ class TestRouterModelRelayUpstreamContract:
|
|||
assert result.headers["x-ms-request-id"] == "req-1"
|
||||
|
||||
|
||||
NIM_INFER_BODY = {
|
||||
"input": [
|
||||
{"type": "image_url", "url": "data:image/png;base64,AAAA"},
|
||||
{"type": "image_url", "url": "data:image/png;base64,BBBB"},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
class TestNvidiaNimProxyRoute:
|
||||
def _request(self) -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = "POST"
|
||||
request.headers = {"content-type": "application/json"}
|
||||
request.query_params = {}
|
||||
return request
|
||||
|
||||
def _recording_router(self, captured: list[dict], deployments: dict[str, str]):
|
||||
class RecordingRouter:
|
||||
def get_model_list(self):
|
||||
return [{"model_name": name, "litellm_params": {"model": model}} for name, model in deployments.items()]
|
||||
|
||||
async def allm_passthrough_route(self, **kwargs):
|
||||
captured.append(kwargs)
|
||||
return httpx.Response(
|
||||
200, json={"data": [{"index": 0, "bounding_boxes": {}}]}, headers={"x-nim-request": "r1"}
|
||||
)
|
||||
|
||||
return RecordingRouter()
|
||||
|
||||
async def _relay(self, llm_router, endpoint: str, body: dict, user_api_key_dict=None) -> Response:
|
||||
return await relay_nvidia_nim_request(
|
||||
llm_router=llm_router,
|
||||
endpoint=endpoint,
|
||||
request=self._request(),
|
||||
request_body=dict(body),
|
||||
user_api_key_dict=user_api_key_dict or UserAPIKeyAuth(api_key="hashed-token"),
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_group_in_the_path_selects_the_deployment_and_the_body_stays_model_free(self):
|
||||
captured: list[dict] = []
|
||||
router = self._recording_router(
|
||||
captured,
|
||||
{
|
||||
"nim-page-elements": "nvidia_nim/nvidia/nemoretriever-page-elements-v2",
|
||||
"nim-table": "nvidia_nim/nvidia/nemoretriever-table-structure-v1",
|
||||
},
|
||||
)
|
||||
|
||||
result = await self._relay(
|
||||
router,
|
||||
"nim-page-elements/v1/infer",
|
||||
NIM_INFER_BODY,
|
||||
UserAPIKeyAuth(api_key="hashed-token", team_id="team-1"),
|
||||
)
|
||||
|
||||
(relay,) = captured
|
||||
assert relay["model"] == "nim-page-elements"
|
||||
assert relay["endpoint"] == "nim-page-elements/v1/infer"
|
||||
assert relay["method"] == "POST"
|
||||
assert relay["json"] == NIM_INFER_BODY
|
||||
assert "model" not in relay["json"]
|
||||
assert relay["litellm_metadata"]["user_api_key_team_id"] == "team-1"
|
||||
assert result.status_code == 200
|
||||
assert json.loads(result.body) == {"data": [{"index": 0, "bounding_boxes": {}}]}
|
||||
assert result.headers["x-nim-request"] == "r1"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_group_with_a_slash_is_matched_as_the_longest_leading_path(self):
|
||||
captured: list[dict] = []
|
||||
router = self._recording_router(
|
||||
captured, {"nvidia/nemoretriever-page-elements-v2": "nvidia_nim/nvidia/nemoretriever-page-elements-v2"}
|
||||
)
|
||||
|
||||
await self._relay(router, "nvidia/nemoretriever-page-elements-v2/v1/infer", NIM_INFER_BODY)
|
||||
|
||||
assert captured[0]["model"] == "nvidia/nemoretriever-page-elements-v2"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_llm_provider_marks_a_deployment_as_nim_without_the_model_prefix(self):
|
||||
captured: list[dict] = []
|
||||
|
||||
class ProviderRouter:
|
||||
def get_model_list(self):
|
||||
return [
|
||||
{
|
||||
"model_name": "page-elements",
|
||||
"litellm_params": {
|
||||
"model": "nvidia/nemoretriever-page-elements-v2",
|
||||
"custom_llm_provider": "nvidia_nim",
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
async def allm_passthrough_route(self, **kwargs):
|
||||
captured.append(kwargs)
|
||||
return httpx.Response(200, json={"data": []})
|
||||
|
||||
await self._relay(ProviderRouter(), "page-elements/v1/infer", NIM_INFER_BODY)
|
||||
|
||||
assert captured[0]["model"] == "page-elements"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint",
|
||||
["v1/infer", "unknown-group/v1/infer", "nim-page-elements-v2/v1/infer", "gpt-4o/v1/infer"],
|
||||
)
|
||||
async def test_path_without_a_nim_model_group_is_rejected_before_any_upstream_call(self, endpoint):
|
||||
captured: list[dict] = []
|
||||
router = self._recording_router(
|
||||
captured,
|
||||
{"nim-page-elements": "nvidia_nim/nvidia/nemoretriever-page-elements-v2", "gpt-4o": "openai/gpt-4o"},
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await self._relay(router, endpoint, NIM_INFER_BODY)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert captured == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_group_mixing_nim_and_other_deployments_is_rejected_before_any_upstream_call(self):
|
||||
captured: list[dict] = []
|
||||
|
||||
class MixedRouter:
|
||||
def get_model_list(self):
|
||||
return [
|
||||
{
|
||||
"model_name": "detect",
|
||||
"litellm_params": {"model": "nvidia_nim/nvidia/nemoretriever-page-elements-v2"},
|
||||
},
|
||||
{"model_name": "detect", "litellm_params": {"model": "openai/gpt-4o"}},
|
||||
]
|
||||
|
||||
async def allm_passthrough_route(self, **kwargs):
|
||||
captured.append(kwargs)
|
||||
return httpx.Response(200, json={"data": []})
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await self._relay(MixedRouter(), "detect/v1/infer", NIM_INFER_BODY)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert captured == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_router_is_rejected_before_any_upstream_call(self):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await self._relay(None, "nim-page-elements/v1/infer", NIM_INFER_BODY)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_upstream_rejection_is_relayed_with_its_status_body_and_headers(self):
|
||||
upstream_body = {"detail": "input[0].url must be a data URL"}
|
||||
|
||||
class RejectingRouter:
|
||||
def get_model_list(self):
|
||||
return [
|
||||
{
|
||||
"model_name": "nim-page-elements",
|
||||
"litellm_params": {"model": "nvidia_nim/nvidia/nemoretriever-page-elements-v2"},
|
||||
}
|
||||
]
|
||||
|
||||
async def allm_passthrough_route(self, **kwargs):
|
||||
upstream_request = httpx.Request("POST", "http://nim.internal:8000/v1/infer")
|
||||
upstream = httpx.Response(
|
||||
422, json=upstream_body, headers={"x-nim-request": "r2"}, request=upstream_request
|
||||
)
|
||||
raise httpx.HTTPStatusError("422", request=upstream_request, response=upstream)
|
||||
|
||||
result = await self._relay(
|
||||
RejectingRouter(), "nim-page-elements/v1/infer", {"input": [{"type": "image_url", "url": "x"}]}
|
||||
)
|
||||
|
||||
assert result.status_code == 422
|
||||
assert json.loads(result.body) == upstream_body
|
||||
assert result.headers["x-nim-request"] == "r2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_count_tokens_error_forwards_provider_headers():
|
||||
"""The count tokens route converts BedrockError into an HTTPException, and dropping the
|
||||
|
|
|
|||
|
|
@ -497,6 +497,33 @@ def test_is_vertex_route_ignores_plain_predict_path_segment():
|
|||
)
|
||||
|
||||
|
||||
def test_interactions_create_routes_are_tracked_for_vertex_and_gemini():
|
||||
"""
|
||||
Regression for LIT-6896: Interactions API (gemini-omni) passthrough responses
|
||||
were never handed to the Vertex/Gemini logging handlers, so SpendLogs rows
|
||||
landed with zero tokens and zero spend. Only the create URL is billable;
|
||||
GET/DELETE on an interaction id and non-Google `/interactions` URLs stay generic.
|
||||
"""
|
||||
handler = PassThroughEndpointLogging()
|
||||
vertex_create = "https://aiplatform.googleapis.com/v1beta1/projects/p/locations/global/interactions"
|
||||
gemini_create = "https://generativelanguage.googleapis.com/v1beta/interactions"
|
||||
|
||||
assert handler.is_vertex_route(vertex_create) is True
|
||||
assert handler.is_vertex_route(f"{vertex_create}/abc123") is False
|
||||
assert handler.is_vertex_route("https://upstream.example.com/api/interactions") is False
|
||||
assert handler.is_vertex_route("https://upstream.example.com/locations/eu/interactions") is False
|
||||
assert (
|
||||
handler.is_vertex_route(
|
||||
"https://us-central1-aiplatform.googleapis.com/v1/projects/p/locations/us-central1/interactions"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
assert handler.is_gemini_route(gemini_create, custom_llm_provider="gemini") is True
|
||||
assert handler.is_gemini_route(f"{gemini_create}/abc123", custom_llm_provider="gemini") is False
|
||||
assert handler.is_gemini_route(gemini_create, custom_llm_provider=None) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_passthrough_predict_path_logs_via_generic_handler():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -317,13 +317,28 @@ _TWO_HEURISTIC_V2_ROUTERS_YAML = (
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("license_limit", [1, None])
|
||||
async def test_ProxyConfig_load_config_takes_the_heuristic_v2_limit_from_the_license_only(
|
||||
tmp_path, monkeypatch, license_limit: int | None
|
||||
@pytest.mark.parametrize("classifier_type", ["heuristic_v2", "capability", "llm_v2"])
|
||||
async def test_ProxyConfig_load_config_takes_the_classifier_limit_from_the_license_only(
|
||||
tmp_path, monkeypatch, license_limit: int | None, classifier_type: str
|
||||
) -> None:
|
||||
"""`router_settings.auto_router_capability_limit` is managed outside config.yaml: an operator
|
||||
cannot grant the entitlement by editing the config, and a licensed proxy boots both routers."""
|
||||
f = tmp_path / "c.yaml"
|
||||
f.write_text(_TWO_HEURISTIC_V2_ROUTERS_YAML)
|
||||
forecast_settings = {
|
||||
"capability": (
|
||||
" classifier_llm_config: {model: gpt-4o-mini}\n"
|
||||
" capability_classifier_config: {efficient_tier: SIMPLE, capable_tier: REASONING, base_threshold: 0.7}\n"
|
||||
),
|
||||
"llm_v2": (
|
||||
" classifier_llm_config: {model: gpt-4o-mini}\n"
|
||||
" adaptive: false\n"
|
||||
" llm_v2_config: {efficient_profile: Small solver, capable_profile: Large solver, harness: One attempt, max_quality_gap: 0.05}\n"
|
||||
),
|
||||
}
|
||||
config_yaml = _TWO_HEURISTIC_V2_ROUTERS_YAML.replace(
|
||||
"classifier_type: heuristic_v2\n", f"classifier_type: {classifier_type}\n{forecast_settings.get(classifier_type, '')}"
|
||||
).replace("tiers: {SIMPLE: gpt-4o-mini}", "tiers: {SIMPLE: gpt-4o-mini, REASONING: gpt-4o}")
|
||||
f.write_text(config_yaml)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
|
||||
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
|
||||
|
|
|
|||
|
|
@ -235,33 +235,6 @@ def test_negative_ttl_counts_do_not_become_cache_write_credits() -> None:
|
|||
assert results[0].prompt_caching < 0
|
||||
|
||||
|
||||
def test_unpublished_one_hour_price_uses_the_ordinary_write_price() -> None:
|
||||
model: Final = "claude-4-opus-20250514"
|
||||
pricing: Final = litellm.get_model_info(model=model, custom_llm_provider="anthropic")
|
||||
assert pricing.get("cache_creation_input_token_cost_above_1hr") is None
|
||||
assert pricing["cache_creation_input_token_cost"] > pricing["input_cost_per_token"]
|
||||
results: Final = tuple(
|
||||
compute_savings_spend(
|
||||
model=model,
|
||||
custom_llm_provider="anthropic",
|
||||
compression_saved_tokens=0,
|
||||
gateway_injected_cache=True,
|
||||
usage_object={
|
||||
"prompt_tokens": 6000,
|
||||
"completion_tokens": 100,
|
||||
"prompt_tokens_details": {
|
||||
"text_tokens": 1000,
|
||||
"cache_creation_tokens": 5000,
|
||||
"cache_creation_token_details": ttl,
|
||||
},
|
||||
},
|
||||
)
|
||||
for ttl in (None, {"ephemeral_1h_input_tokens": 5000})
|
||||
)
|
||||
assert results[0] == results[1]
|
||||
assert results[0].prompt_caching < 0
|
||||
|
||||
|
||||
def test_prompt_caching_savings_nets_out_the_cache_write_premium():
|
||||
"""A cache-writing request is only credited the read discount minus the write premium."""
|
||||
input_cost, cache_read_cost = _anthropic_costs("claude-sonnet-5")
|
||||
|
|
@ -354,82 +327,6 @@ def test_openai_style_cache_write_tokens_are_netted_out():
|
|||
)
|
||||
|
||||
|
||||
def test_model_without_a_cache_write_price_takes_no_premium():
|
||||
"""An absent write price must mean zero premium, never a bonus.
|
||||
|
||||
``_get_cost_per_unit`` in the cost calculator defaults a missing price to 0.0. Were
|
||||
that default copied here the premium would be ``0 - input_cost``, and a model with no
|
||||
write pricing would report cache writes as free money. This is the common case: most
|
||||
of the pricing map publishes a cache-read price and no cache-write price.
|
||||
"""
|
||||
model = "amazon.nova-2-lite-v1:0"
|
||||
info = litellm.get_model_info(model=model)
|
||||
input_cost = info["input_cost_per_token"]
|
||||
cache_read_cost = info["cache_read_input_token_cost"]
|
||||
assert info.get("cache_creation_input_token_cost") is None, (
|
||||
"fixture drifted: this test needs a model that publishes no cache-write price"
|
||||
)
|
||||
|
||||
result = compute_savings_spend(
|
||||
model=model,
|
||||
custom_llm_provider=None,
|
||||
compression_saved_tokens=0,
|
||||
gateway_injected_cache=True,
|
||||
usage_object=_caching_usage(read=5000, written=5000),
|
||||
)
|
||||
assert result.prompt_caching == pytest.approx(5000 * (input_cost - cache_read_cost))
|
||||
assert result.prompt_caching > 0
|
||||
|
||||
|
||||
def test_zero_cache_write_price_is_read_as_unpublished():
|
||||
"""A ``0.0`` write price means "no separate price", not "writes are free".
|
||||
|
||||
``deepseek-chat`` carries an explicit zero in the pricing map. Taken literally the
|
||||
premium would be ``0 - input_cost``, paying out a saving of ``writes * input_cost``
|
||||
on traffic that cached nothing. No provider gives cache writes away, so a falsy
|
||||
price falls open to the input cost like an absent one does.
|
||||
"""
|
||||
info = litellm.get_model_info(model="deepseek-chat", custom_llm_provider="deepseek")
|
||||
assert info.get("cache_creation_input_token_cost") == 0.0, (
|
||||
"fixture drifted: this test exists because deepseek-chat publishes a literal 0.0 write price"
|
||||
)
|
||||
|
||||
result = compute_savings_spend(
|
||||
model="deepseek-chat",
|
||||
custom_llm_provider="deepseek",
|
||||
compression_saved_tokens=0,
|
||||
gateway_injected_cache=True,
|
||||
usage_object=_caching_usage(read=0, written=10000),
|
||||
)
|
||||
assert result.prompt_caching == pytest.approx(0.0)
|
||||
|
||||
|
||||
def test_zero_cache_read_price_stays_literal():
|
||||
"""The read leg must NOT copy the write leg's falsy fall-open.
|
||||
|
||||
The two zeros mean opposite things. A free cache *write* is unpublished pricing, so
|
||||
it falls open to input. A free cache *read* is real and is the largest discount
|
||||
available -- 15 models charge for input and serve reads for nothing. Falling that
|
||||
open to the input cost would zero out their savings entirely.
|
||||
"""
|
||||
model = "gemini-robotics-er-1.5-preview"
|
||||
info = litellm.get_model_info(model=model)
|
||||
input_cost = info["input_cost_per_token"]
|
||||
assert info.get("cache_read_input_token_cost") == 0.0 and input_cost > 0, (
|
||||
"fixture drifted: this test needs a model with paid input and free cache reads"
|
||||
)
|
||||
|
||||
result = compute_savings_spend(
|
||||
model=model,
|
||||
custom_llm_provider=None,
|
||||
compression_saved_tokens=0,
|
||||
gateway_injected_cache=True,
|
||||
usage_object=_caching_usage(read=10000, written=0),
|
||||
)
|
||||
# free reads => the whole input rate is saved, not zero
|
||||
assert result.prompt_caching == pytest.approx(10000 * input_cost)
|
||||
|
||||
|
||||
def test_sub_input_cache_write_price_is_an_extra_saving():
|
||||
"""A few models price writes below input; there the premium is a real credit.
|
||||
|
||||
|
|
@ -441,9 +338,6 @@ def test_sub_input_cache_write_price_is_an_extra_saving():
|
|||
input_cost = info["input_cost_per_token"]
|
||||
cheap_write = info["cache_creation_input_token_cost"]
|
||||
assert 0 < cheap_write < input_cost, "fixture drifted: this test needs a model pricing cache writes below input"
|
||||
# no published read price, so the read leg mirrors input and contributes nothing;
|
||||
# the whole result is the negative premium, i.e. a credit.
|
||||
assert info.get("cache_read_input_token_cost") is None
|
||||
|
||||
result = compute_savings_spend(
|
||||
model=model,
|
||||
|
|
@ -728,21 +622,6 @@ def test_malformed_usage_object_does_not_fail_the_spend_write():
|
|||
assert result.compression > 0
|
||||
|
||||
|
||||
def test_model_without_cache_read_pricing_yields_no_caching_savings():
|
||||
"""A model with no discounted cache-read rate cannot have saved anything by
|
||||
reading from cache, so the driver must report zero rather than the full input rate."""
|
||||
model = "azure/gpt-3.5-turbo"
|
||||
assert litellm.get_model_info(model=model).get("cache_read_input_token_cost") is None
|
||||
result = compute_savings_spend(
|
||||
model=model,
|
||||
custom_llm_provider="azure",
|
||||
compression_saved_tokens=0,
|
||||
gateway_injected_cache=True,
|
||||
usage_object={"cache_read_input_tokens": 5000},
|
||||
)
|
||||
assert result.prompt_caching == 0.0
|
||||
|
||||
|
||||
def test_the_same_deployment_spelled_two_ways_is_not_a_switch():
|
||||
"""The spend log records a normalized model name while the baseline arrives as the
|
||||
operator wrote it in config. Comparing the raw strings makes a request that never
|
||||
|
|
|
|||
|
|
@ -5353,6 +5353,87 @@ async def test_build_ui_spend_logs_response_sums_multi_round_session_tokens():
|
|||
assert all(key not in rows[2] for key in token_keys)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_ui_spend_logs_response_sums_multi_round_session_duration():
|
||||
"""
|
||||
Regression test: a multi-round session collapses into a single UI row, so that row
|
||||
must carry the duration of every round summed, not just the representative call's.
|
||||
Rows written before request_duration_ms existed are NULL, so the aggregate falls back
|
||||
to endTime - startTime for them.
|
||||
"""
|
||||
from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
||||
_build_ui_spend_logs_response,
|
||||
)
|
||||
|
||||
session_id = "sess-multi-round-duration"
|
||||
api_key = "hashed-key-xyz"
|
||||
dict_rows = [
|
||||
{
|
||||
"request_id": "req-1",
|
||||
"session_id": session_id,
|
||||
"call_type": "completion",
|
||||
"api_key": api_key,
|
||||
"spend": 0.01,
|
||||
"request_duration_ms": 1200,
|
||||
},
|
||||
{
|
||||
"request_id": "req-2",
|
||||
"session_id": session_id,
|
||||
"call_type": "completion",
|
||||
"api_key": api_key,
|
||||
"spend": 0.02,
|
||||
"request_duration_ms": 4200,
|
||||
},
|
||||
{
|
||||
"request_id": "req-3",
|
||||
"session_id": None,
|
||||
"call_type": "completion",
|
||||
"api_key": api_key,
|
||||
"spend": 0.03,
|
||||
"request_duration_ms": 900,
|
||||
},
|
||||
]
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.query_raw = AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"session_id": session_id,
|
||||
"api_key": api_key,
|
||||
"session_total_count": 2,
|
||||
"session_total_spend": 0.03,
|
||||
"session_total_duration_ms": 5400,
|
||||
"mcp_tool_call_count": 0,
|
||||
"mcp_tool_call_spend": 0.0,
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
result = await _build_ui_spend_logs_response(
|
||||
prisma_client=mock_prisma,
|
||||
data=dict_rows,
|
||||
total_records=3,
|
||||
page=1,
|
||||
page_size=50,
|
||||
total_pages=1,
|
||||
enrich_session_counts=True,
|
||||
)
|
||||
|
||||
rows = result["data"]
|
||||
session_rows = rows[:2]
|
||||
assert [row["session_total_duration_ms"] for row in session_rows] == [5400, 5400]
|
||||
assert all(isinstance(row["session_total_duration_ms"], int) for row in session_rows)
|
||||
assert [row["request_duration_ms"] for row in rows] == [1200, 4200, 900]
|
||||
assert "session_total_duration_ms" not in rows[2]
|
||||
|
||||
_, call_args, _ = mock_prisma.db.query_raw.mock_calls[0]
|
||||
sql = " ".join(call_args[0].split())
|
||||
assert (
|
||||
'SUM( COALESCE( request_duration_ms, (EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER ) )'
|
||||
in sql
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_ui_spend_logs_response_session_cache_hit_count():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -139,6 +139,35 @@ class TestProxyInitializationHelpers:
|
|||
)
|
||||
assert args["timeout_worker_healthcheck"] == 15
|
||||
|
||||
@staticmethod
|
||||
def _uvicorn_access_info_enabled(args: dict) -> bool:
|
||||
import logging
|
||||
|
||||
loggers = tuple(logging.getLogger(n) for n in ("uvicorn", "uvicorn.error", "uvicorn.access", "uvicorn.asgi"))
|
||||
saved = tuple((lg, lg.handlers[:], lg.level, lg.propagate) for lg in loggers)
|
||||
try:
|
||||
uvicorn.Config(**args).configure_logging()
|
||||
return logging.getLogger("uvicorn.access").isEnabledFor(logging.INFO)
|
||||
finally:
|
||||
for lg, handlers, level, propagate in saved:
|
||||
lg.handlers[:] = handlers
|
||||
lg.setLevel(level)
|
||||
lg.propagate = propagate
|
||||
|
||||
def test_litellm_log_error_silences_uvicorn_info_lines(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOG", "ERROR")
|
||||
args = ProxyInitializationHelpers._get_default_unvicorn_init_args("localhost", 8000)
|
||||
|
||||
assert "log_config" not in args
|
||||
assert self._uvicorn_access_info_enabled(args) is False
|
||||
|
||||
def test_unset_litellm_log_keeps_uvicorn_default_info_lines(self, monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_LOG", raising=False)
|
||||
args = ProxyInitializationHelpers._get_default_unvicorn_init_args("localhost", 8000)
|
||||
|
||||
assert "log_level" not in args
|
||||
assert self._uvicorn_access_info_enabled(args) is True
|
||||
|
||||
def test_installed_uvicorn_supports_worker_flags(self):
|
||||
params = inspect.signature(uvicorn.Config.__init__).parameters
|
||||
assert "timeout_worker_healthcheck" in params
|
||||
|
|
|
|||
|
|
@ -87,6 +87,27 @@ def test_convert_mcp_to_llm_format_exposes_headers_on_metadata(proxy_logging, ma
|
|||
assert out["metadata"]["headers"] == {"x-nuid": "nuid-1"}
|
||||
|
||||
|
||||
def test_convert_mcp_to_llm_format_exposes_caller_identity_on_metadata(proxy_logging, make_mcp_request_obj):
|
||||
"""Custom code guardrails resolve user_id/team_id/end_user_id from the proxy-owned metadata
|
||||
bucket on every route, so the MCP bridge has to write the authenticated ids there too."""
|
||||
req = make_mcp_request_obj()
|
||||
out = proxy_logging._convert_mcp_to_llm_format(
|
||||
request_obj=req,
|
||||
kwargs={
|
||||
"user_api_key_user_id": "u-1",
|
||||
"user_api_key_team_id": "t-1",
|
||||
"user_api_key_end_user_id": "eu-1",
|
||||
"headers": {"x-nuid": "nuid-1"},
|
||||
},
|
||||
)
|
||||
assert out["metadata"] == {
|
||||
"headers": {"x-nuid": "nuid-1"},
|
||||
"user_api_key_user_id": "u-1",
|
||||
"user_api_key_team_id": "t-1",
|
||||
"user_api_key_end_user_id": "eu-1",
|
||||
}
|
||||
|
||||
|
||||
def test_convert_mcp_to_llm_format_defaults_headers_to_empty(proxy_logging, make_mcp_request_obj):
|
||||
req = make_mcp_request_obj()
|
||||
out = proxy_logging._convert_mcp_to_llm_format(request_obj=req, kwargs={})
|
||||
|
|
|
|||
|
|
@ -1417,6 +1417,66 @@ class TestRouterComplexityDeploymentMethods:
|
|||
router.init_complexity_router_deployment(deployment)
|
||||
assert "auto_router/complexity_router/test-router" in router.complexity_routers
|
||||
|
||||
@staticmethod
|
||||
def _forecast_row(model_name: str, model_id: str, classifier_type: str) -> dict[str, object]:
|
||||
settings: Final = (
|
||||
{"capability_classifier_config": {
|
||||
"efficient_tier": "SIMPLE", "capable_tier": "REASONING", "base_threshold": 0.7,
|
||||
}} if classifier_type == "capability" else {
|
||||
"adaptive": False,
|
||||
"llm_v2_config": {
|
||||
"efficient_profile": "Small solver", "capable_profile": "Large solver",
|
||||
"harness": "One attempt", "max_quality_gap": 0.05,
|
||||
},
|
||||
}
|
||||
)
|
||||
return {
|
||||
"model_name": model_name,
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": {
|
||||
"classifier_type": classifier_type,
|
||||
"classifier_llm_config": {"model": "gpt-4o-mini"},
|
||||
"tiers": {"SIMPLE": "gpt-4o-mini", "REASONING": "gpt-4o"},
|
||||
**settings,
|
||||
},
|
||||
},
|
||||
"model_info": {"id": model_id},
|
||||
}
|
||||
|
||||
@pytest.mark.parametrize("classifier_type,sibling", [("capability", "llm_v2"), ("llm_v2", "capability")])
|
||||
def test_forecast_cap_keeps_edits_and_refuses_extra_routers_and_type_switches(self, classifier_type: str, sibling: str) -> None:
|
||||
router: Final = Router(
|
||||
model_list=[
|
||||
self._POOL,
|
||||
self._forecast_row("held", "held-id", classifier_type),
|
||||
self._forecast_row("sibling", "sibling-id", sibling),
|
||||
self._router_row("other", "other-id", "heuristic_v2"),
|
||||
self._custom_tier_row("custom", "custom-id"),
|
||||
],
|
||||
auto_router_capability_limit=lambda: 1,
|
||||
ignore_invalid_deployments=True,
|
||||
)
|
||||
assert sorted(router.complexity_routers) == ["custom", "held", "other", "sibling"]
|
||||
assert router.upsert_deployment(Deployment(**self._forecast_row("edited", "held-id", classifier_type))) is not None
|
||||
assert router.upsert_deployment(Deployment(**self._forecast_row("second", "new-id", classifier_type))) is None
|
||||
assert router.upsert_deployment(Deployment(**self._forecast_row("switched", "other-id", classifier_type))) is None
|
||||
assert sorted(router.complexity_routers) == ["custom", "edited", "other", "sibling"]
|
||||
assert router.upsert_deployment(Deployment(**self._router_row("released", "held-id", "heuristic"))) is not None
|
||||
assert router.upsert_deployment(Deployment(**self._forecast_row("switched", "other-id", classifier_type))) is not None
|
||||
assert sorted(router.complexity_routers) == ["custom", "released", "sibling", "switched"]
|
||||
|
||||
@pytest.mark.parametrize("classifier_type", ["capability", "llm_v2"])
|
||||
@pytest.mark.parametrize("limit", [1, None])
|
||||
def test_forecast_registration_applies_the_resolved_license_limit(self, classifier_type: str, limit: int | None) -> None:
|
||||
rows: Final = [self._POOL, self._forecast_row("a", "id-a", classifier_type), self._forecast_row("b", "id-b", classifier_type)]
|
||||
if limit is not None:
|
||||
with pytest.raises(ValueError, match="At most 1 auto-router"):
|
||||
Router(model_list=rows, auto_router_capability_limit=lambda: limit)
|
||||
return
|
||||
router: Final = Router(model_list=rows, auto_router_capability_limit=lambda: limit)
|
||||
assert sorted(router.complexity_routers) == ["a", "b"]
|
||||
|
||||
@staticmethod
|
||||
def _router_row(model_name: str, model_id: str, classifier_type: str) -> dict[str, object]:
|
||||
return {
|
||||
|
|
|
|||
|
|
@ -395,6 +395,8 @@ def test_placement_is_scoped_to_complexity_router_deployments(model, present_fie
|
|||
|
||||
|
||||
_HV2_CONFIG: Mapping[str, object] = {"classifier_type": "heuristic_v2"}
|
||||
_CAPABILITY_CONFIG: Mapping[str, object] = {"classifier_type": "capability"}
|
||||
_FUSE_CONFIG: Mapping[str, object] = {"classifier_type": "llm_v2"}
|
||||
_CUSTOM_TIER_CONFIG: Mapping[str, object] = {
|
||||
"classifier_type": "llm",
|
||||
"tier_definitions": [{"name": "routine", "description": "easy"}, {"name": "hard", "description": "hard"}],
|
||||
|
|
@ -457,6 +459,10 @@ def test_is_complexity_router_model(model: str | None, expected: bool) -> None:
|
|||
@pytest.mark.parametrize(
|
||||
"litellm_params,expected_key",
|
||||
[
|
||||
({"model": "auto_router/complexity_router", "complexity_router_config": _CAPABILITY_CONFIG}, "capability"),
|
||||
({"model": "auto_router/complexity_router-eu", "complexity_router_config": _FUSE_CONFIG}, "llm_v2"),
|
||||
({"model": "openai/solver", "complexity_router_config": _CAPABILITY_CONFIG}, None),
|
||||
({"model": "auto_router/quality_router", "complexity_router_config": _FUSE_CONFIG}, None),
|
||||
({"model": "auto_router/complexity_router", "complexity_router_config": _HV2_CONFIG}, "heuristic_v2"),
|
||||
({"model": "auto_router/complexity_router-eu", "complexity_router_config": _HV2_CONFIG}, "heuristic_v2"),
|
||||
({"model": "auto_router/complexity_router", "complexity_router_config": _CUSTOM_TIER_CONFIG}, "tier_or_classifier_prompt"),
|
||||
|
|
@ -493,6 +499,8 @@ def test_count_capability_routers_counts_only_its_own_capability(capability) ->
|
|||
|
||||
by_key = {
|
||||
"heuristic_v2": (_HV2_CONFIG, _HV2_CONFIG),
|
||||
"capability": (_CAPABILITY_CONFIG, _CAPABILITY_CONFIG),
|
||||
"llm_v2": (_FUSE_CONFIG, _FUSE_CONFIG),
|
||||
"tier_or_classifier_prompt": (_CUSTOM_TIER_CONFIG, _CUSTOM_PROMPT_CONFIG),
|
||||
}
|
||||
mine_first, mine_second = by_key[capability.key]
|
||||
|
|
@ -545,6 +553,8 @@ def test_every_gated_capability_has_a_distinct_predicate_and_sql_spelling() -> N
|
|||
"config",
|
||||
[
|
||||
_HV2_CONFIG,
|
||||
_CAPABILITY_CONFIG,
|
||||
_FUSE_CONFIG,
|
||||
_CUSTOM_TIER_CONFIG,
|
||||
_CUSTOM_PROMPT_CONFIG,
|
||||
{"classifier_type": "heuristic"},
|
||||
|
|
|
|||
|
|
@ -34,6 +34,4 @@ def test_azure_ai_grok_4_3_backup_matches_main():
|
|||
main_cost = _load_model_cost(main_path)
|
||||
backup_cost = _load_model_cost(backup_path)
|
||||
|
||||
assert backup_cost.get(AZURE_AI_GROK_4_3_MODEL) == main_cost.get(
|
||||
AZURE_AI_GROK_4_3_MODEL
|
||||
)
|
||||
assert backup_cost.get(AZURE_AI_GROK_4_3_MODEL) == main_cost.get(AZURE_AI_GROK_4_3_MODEL)
|
||||
|
|
|
|||
|
|
@ -24,12 +24,6 @@ def test_azure_ai_grok_4_6_is_priced_and_routed() -> None:
|
|||
info = get_model_info(model=routed_model, custom_llm_provider=provider)
|
||||
assert info["litellm_provider"] == "azure_ai"
|
||||
assert info["mode"] == "chat"
|
||||
assert info["input_cost_per_token"] == 2e-06
|
||||
assert info["output_cost_per_token"] == 6e-06
|
||||
assert info["cache_read_input_token_cost"] == 5e-07
|
||||
assert info["max_input_tokens"] == 200000
|
||||
assert info["max_output_tokens"] == 128000
|
||||
assert info["max_tokens"] == 128000
|
||||
assert info["supports_function_calling"] is True
|
||||
assert info["supports_prompt_caching"] is True
|
||||
assert info["supports_reasoning"] is True
|
||||
|
|
@ -39,8 +33,8 @@ def test_azure_ai_grok_4_6_is_priced_and_routed() -> None:
|
|||
assert info["supports_web_search"] is True
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model=MODEL, prompt_tokens=1_000_000, completion_tokens=1_000_000)
|
||||
assert prompt_cost == pytest.approx(2.0)
|
||||
assert completion_cost == pytest.approx(6.0)
|
||||
assert prompt_cost > 0
|
||||
assert completion_cost > 0
|
||||
|
||||
|
||||
def test_azure_ai_grok_4_6_entry_source_and_backup_match() -> None:
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ from pathlib import Path
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
|
||||
from litellm.utils import supports_function_calling, supports_prompt_caching
|
||||
|
||||
REPO_ROOT = Path(__file__).parents[2]
|
||||
|
|
@ -41,26 +40,8 @@ def test_baseten_glm_5_3_capabilities_are_visible_to_callers(local_model_cost_ma
|
|||
assert supports_function_calling(model=MODEL) is True
|
||||
|
||||
info = litellm.get_model_info(model="zai-org/GLM-5.3", custom_llm_provider="baseten")
|
||||
assert info["max_input_tokens"] == 1048576
|
||||
assert info["max_output_tokens"] == 262144
|
||||
|
||||
|
||||
def test_cached_prompt_tokens_bill_at_the_cached_rate(local_model_cost_map):
|
||||
"""A cache hit reports its reused tokens under prompt_tokens_details, and those
|
||||
tokens cost a tenth of the input rate, not the full rate and not nothing."""
|
||||
usage = Usage(
|
||||
prompt_tokens=21010,
|
||||
completion_tokens=100,
|
||||
total_tokens=21110,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=20992),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = litellm.cost_per_token(
|
||||
model=MODEL, usage_object=usage, custom_llm_provider="baseten"
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(18 * INPUT_COST + 20992 * CACHED_INPUT_COST)
|
||||
assert completion_cost == pytest.approx(100 * OUTPUT_COST)
|
||||
assert info["max_input_tokens"] > 0
|
||||
assert info["max_output_tokens"] > 0
|
||||
|
||||
|
||||
def test_backup_matches_main():
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.constants import bedrock_embedding_models
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
|
||||
|
||||
REPO_ROOT = Path(__file__).parents[2]
|
||||
MAIN_PATH = REPO_ROOT / "model_prices_and_context_window.json"
|
||||
|
|
@ -37,38 +36,6 @@ def test_marengo_embed_3_is_visible_to_callers(model, local_model_cost_map):
|
|||
info = litellm.get_model_info(model=model, custom_llm_provider="bedrock")
|
||||
assert info["mode"] == "embedding"
|
||||
assert info["output_vector_size"] == 512
|
||||
assert info["max_input_tokens"] == 500
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", PER_REQUEST_MODELS)
|
||||
@pytest.mark.parametrize(
|
||||
"details,expected_cost",
|
||||
[
|
||||
(PromptTokensDetailsWrapper(query_count=1), TEXT_REQUEST_COST),
|
||||
(PromptTokensDetailsWrapper(image_count=1), IMAGE_REQUEST_COST),
|
||||
(PromptTokensDetailsWrapper(query_count=1, image_count=1), TEXT_REQUEST_COST + IMAGE_REQUEST_COST),
|
||||
(PromptTokensDetailsWrapper(query_count=1, image_count=2), TEXT_REQUEST_COST + 2 * IMAGE_REQUEST_COST),
|
||||
(PromptTokensDetailsWrapper(video_length_seconds=10), 10 * VIDEO_COST_PER_SECOND),
|
||||
(PromptTokensDetailsWrapper(audio_length_seconds=10), 10 * AUDIO_COST_PER_SECOND),
|
||||
],
|
||||
)
|
||||
def test_marengo_requests_are_billed_per_request(model, details, expected_cost, local_model_cost_map):
|
||||
usage = Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0, prompt_tokens_details=details)
|
||||
prompt_cost, completion_cost = litellm.cost_per_token(
|
||||
model=model, usage_object=usage, custom_llm_provider="bedrock"
|
||||
)
|
||||
assert prompt_cost == pytest.approx(expected_cost)
|
||||
assert completion_cost == 0.0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model", PER_REQUEST_MODELS)
|
||||
def test_marengo_token_counts_bill_nothing(model, local_model_cost_map):
|
||||
usage = Usage(prompt_tokens=128, completion_tokens=0, total_tokens=128)
|
||||
prompt_cost, completion_cost = litellm.cost_per_token(
|
||||
model=model, usage_object=usage, custom_llm_provider="bedrock"
|
||||
)
|
||||
assert prompt_cost == 0.0
|
||||
assert completion_cost == 0.0
|
||||
|
||||
|
||||
def test_marengo_embed_3_is_a_known_bedrock_embedding_model():
|
||||
|
|
|
|||
|
|
@ -26,15 +26,6 @@ def _load_root_cost_map() -> dict:
|
|||
return json.load(f)
|
||||
|
||||
|
||||
def test_fable_5_geo_multiplier_without_fast_mode():
|
||||
"""First-party ``inference_geo='us'`` carries the 1.1x premium, but unlike
|
||||
the Opus line there is no fast-mode variant for Fable 5; a ``fast`` key
|
||||
here would silently misprice ``speed='fast'`` requests."""
|
||||
model_data = _load_root_cost_map()
|
||||
entry = model_data["claude-fable-5"]["provider_specific_entry"]
|
||||
assert entry == {"us": 1.1}
|
||||
|
||||
|
||||
def test_fable_5_present_in_bundled_backup():
|
||||
"""The bundled backup is the runtime fallback (and what tests load with
|
||||
``LITELLM_LOCAL_MODEL_COST_MAP=True``) — it must carry the same entries as
|
||||
|
|
@ -75,9 +66,7 @@ def test_fable_5_all_variants_carry_adaptive_thinking_flag(cost_map):
|
|||
so adaptive is the only valid thinking shape LiteLLM can emit for it."""
|
||||
variants = [k for k in cost_map if "claude-fable-5" in k]
|
||||
assert variants, "no claude-fable-5 entries found in cost map"
|
||||
missing = [
|
||||
k for k in variants if cost_map[k].get("supports_adaptive_thinking") is not True
|
||||
]
|
||||
missing = [k for k in variants if cost_map[k].get("supports_adaptive_thinking") is not True]
|
||||
assert not missing, f"missing supports_adaptive_thinking: {missing}"
|
||||
|
||||
|
||||
|
|
@ -131,24 +120,6 @@ FABLE_5_1_VARIANTS = (
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cost_map",
|
||||
[_load_root_cost_map(), GetModelCostMap.load_local_model_cost_map()],
|
||||
ids=["root", "bundled_backup"],
|
||||
)
|
||||
def test_fable_5_1_cache_reads_cost_a_quarter_of_fable_5(cost_map):
|
||||
"""Fable 5.1 prices cache hits at 0.025x base input instead of the usual
|
||||
0.1x, so copying Fable 5's cache-read price overcharges every cache hit 4x."""
|
||||
for model_name in FABLE_5_1_VARIANTS:
|
||||
info = cost_map[model_name]
|
||||
geo_premium = model_name.startswith(("us.", "eu."))
|
||||
expected = 2.75e-07 if geo_premium else 2.5e-07
|
||||
assert info["cache_read_input_token_cost"] == expected, model_name
|
||||
assert info["cache_read_input_token_cost"] == pytest.approx(
|
||||
info["input_cost_per_token"] * 0.025
|
||||
), model_name
|
||||
|
||||
|
||||
def test_fable_5_1_present_in_bundled_backup():
|
||||
backup = GetModelCostMap.load_local_model_cost_map()
|
||||
root = _load_root_cost_map()
|
||||
|
|
@ -197,7 +168,5 @@ def test_sampling_params_flag_on_all_models_that_removed_them(cost_map):
|
|||
and not k.startswith("perplexity/")
|
||||
]
|
||||
assert variants, "no matching entries found in cost map"
|
||||
missing = [
|
||||
k for k in variants if cost_map[k].get("supports_sampling_params") is not False
|
||||
]
|
||||
missing = [k for k in variants if cost_map[k].get("supports_sampling_params") is not False]
|
||||
assert not missing, f"missing supports_sampling_params=false: {missing}"
|
||||
|
|
|
|||
|
|
@ -13,9 +13,7 @@ def test_bedrock_haiku_4_5_matches_sonnet_capabilities():
|
|||
(including computer_use, vision, tools, etc.)
|
||||
"""
|
||||
# Load model configuration
|
||||
json_path = os.path.join(
|
||||
os.path.dirname(__file__), "../../model_prices_and_context_window.json"
|
||||
)
|
||||
json_path = os.path.join(os.path.dirname(__file__), "../../model_prices_and_context_window.json")
|
||||
with open(json_path) as f:
|
||||
model_data = json.load(f)
|
||||
|
||||
|
|
@ -43,6 +41,6 @@ def test_bedrock_haiku_4_5_matches_sonnet_capabilities():
|
|||
]
|
||||
|
||||
for capability in shared_capabilities:
|
||||
assert haiku_info.get(capability) == sonnet_info.get(
|
||||
capability
|
||||
), f"Capability {capability} mismatch: Haiku={haiku_info.get(capability)}, Sonnet={sonnet_info.get(capability)}"
|
||||
assert haiku_info.get(capability) == sonnet_info.get(capability), (
|
||||
f"Capability {capability} mismatch: Haiku={haiku_info.get(capability)}, Sonnet={sonnet_info.get(capability)}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -88,7 +88,5 @@ def test_opus_5_all_variants_carry_adaptive_thinking_flag(cost_map):
|
|||
Opus 5 rejects with a 400."""
|
||||
variants = [k for k in cost_map if "claude-opus-5" in k]
|
||||
assert variants, "no claude-opus-5 entries found in cost map"
|
||||
missing = [
|
||||
k for k in variants if cost_map[k].get("supports_adaptive_thinking") is not True
|
||||
]
|
||||
missing = [k for k in variants if cost_map[k].get("supports_adaptive_thinking") is not True]
|
||||
assert not missing, f"missing supports_adaptive_thinking: {missing}"
|
||||
|
|
|
|||
|
|
@ -1,67 +0,0 @@
|
|||
"""
|
||||
Regression test: ``command-r7b-12-2024`` had its input/output per-token
|
||||
costs transposed in the model-cost maps (input=1.5e-07 / output=3.75e-08),
|
||||
even though Cohere publishes $0.0375/1M input and $0.15/1M output, i.e.
|
||||
output is ~4x input like every other ``command-r`` entry.
|
||||
|
||||
These tests pin the corrected values in both the primary price map and the
|
||||
``litellm/`` backup, and verify ``get_model_info`` surfaces them, so the
|
||||
swap cannot silently regress.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
|
||||
import litellm
|
||||
|
||||
MODEL = "command-r7b-12-2024"
|
||||
EXPECTED_INPUT_COST = 3.75e-08
|
||||
EXPECTED_OUTPUT_COST = 1.5e-07
|
||||
|
||||
|
||||
def _load_json(path: str) -> dict:
|
||||
with open(path, encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def _backup_path() -> str:
|
||||
return os.path.join(
|
||||
os.path.dirname(litellm.__file__),
|
||||
"model_prices_and_context_window_backup.json",
|
||||
)
|
||||
|
||||
|
||||
def _main_path() -> str:
|
||||
# This test lives at ``tests/test_litellm/``; the primary price map sits at
|
||||
# the repo root, two directories up. Resolve it relative to this file so the
|
||||
# test works regardless of where ``litellm`` itself is installed (e.g. a pip
|
||||
# install into site-packages).
|
||||
return os.path.join(
|
||||
os.path.dirname(__file__),
|
||||
"..",
|
||||
"..",
|
||||
"model_prices_and_context_window.json",
|
||||
)
|
||||
|
||||
|
||||
class TestCommandR7bPricingData:
|
||||
"""The JSON price maps must carry Cohere's published costs, with output
|
||||
more expensive than input."""
|
||||
|
||||
|
||||
class TestCommandR7bPricingModelInfo:
|
||||
"""``get_model_info`` must report the corrected, un-swapped costs."""
|
||||
|
||||
def test_get_model_info_costs(self):
|
||||
# Patch litellm.model_cost with the local backup so the test is not
|
||||
# dependent on the remote fetch hitting a not-yet-merged main branch.
|
||||
original = litellm.model_cost
|
||||
try:
|
||||
litellm.model_cost = _load_json(_backup_path())
|
||||
info = litellm.get_model_info(MODEL)
|
||||
assert info["input_cost_per_token"] == EXPECTED_INPUT_COST
|
||||
assert info["output_cost_per_token"] == EXPECTED_OUTPUT_COST
|
||||
assert info["output_cost_per_token"] > info["input_cost_per_token"]
|
||||
finally:
|
||||
litellm.model_cost = original
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -12,14 +12,12 @@ field set to ``True``.
|
|||
import json
|
||||
import os
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.utils import (
|
||||
_supports_factory,
|
||||
supports_response_schema,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Data-level tests – verify the JSON files are in sync
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -65,23 +63,13 @@ class TestSupportsResponseSchemaDeepSeek:
|
|||
assert supports_response_schema(model="deepseek/deepseek-chat") is True
|
||||
|
||||
def test_explicit_provider(self):
|
||||
assert (
|
||||
supports_response_schema(
|
||||
model="deepseek-chat", custom_llm_provider="deepseek"
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert supports_response_schema(model="deepseek-chat", custom_llm_provider="deepseek") is True
|
||||
|
||||
def test_reasoner_provider_slash_model(self):
|
||||
assert supports_response_schema(model="deepseek/deepseek-reasoner") is True
|
||||
|
||||
def test_reasoner_explicit_provider(self):
|
||||
assert (
|
||||
supports_response_schema(
|
||||
model="deepseek-reasoner", custom_llm_provider="deepseek"
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert supports_response_schema(model="deepseek-reasoner", custom_llm_provider="deepseek") is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -14,27 +14,12 @@ import os
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm import completion_cost
|
||||
from litellm.types.utils import Choices, Message, ModelResponse, Usage
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
|
||||
NEW_ENTRIES = {
|
||||
"fireworks_ai/accounts/fireworks/models/deepseek-v4-pro-0813": {
|
||||
"input_cost_per_token": 1.32e-06,
|
||||
"cache_read_input_token_cost": 4.4e-08,
|
||||
"output_cost_per_token": 3.96e-06,
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 131072,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def model_data():
|
||||
json_path = os.path.join(
|
||||
os.path.dirname(__file__), "../../model_prices_and_context_window.json"
|
||||
)
|
||||
json_path = os.path.join(os.path.dirname(__file__), "../../model_prices_and_context_window.json")
|
||||
with open(json_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
|
@ -48,44 +33,8 @@ def test_bare_fireworks_ids_resolve_through_prefixed_entries():
|
|||
),
|
||||
]:
|
||||
info = get_model_info(model=bare_id, custom_llm_provider="fireworks_ai")
|
||||
expected = NEW_ENTRIES[prefixed_key]
|
||||
assert info.get("key") == prefixed_key
|
||||
assert info["litellm_provider"] == "fireworks_ai"
|
||||
assert info["input_cost_per_token"] == pytest.approx(expected["input_cost_per_token"])
|
||||
assert info["cache_read_input_token_cost"] == pytest.approx(expected["cache_read_input_token_cost"])
|
||||
assert info["output_cost_per_token"] == pytest.approx(expected["output_cost_per_token"])
|
||||
assert info["max_input_tokens"] == expected["max_input_tokens"]
|
||||
assert info["max_output_tokens"] == expected["max_output_tokens"]
|
||||
|
||||
|
||||
def test_deepseek_v4p1_flash_twin_costs(local_model_cost_map):
|
||||
for model in (
|
||||
"fireworks_ai/deepseek-v4p1-flash",
|
||||
"fireworks_ai/accounts/fireworks/models/deepseek-v4p1-flash",
|
||||
):
|
||||
response = ModelResponse(
|
||||
model=model,
|
||||
choices=[Choices(index=0, message=Message(role="assistant", content="ok"))],
|
||||
usage=Usage(prompt_tokens=1000, completion_tokens=1000, total_tokens=2000),
|
||||
)
|
||||
cost = completion_cost(completion_response=response, model=model)
|
||||
assert cost == pytest.approx(8.8e-04)
|
||||
|
||||
|
||||
TWIN_PINNED_PRICES = {
|
||||
"deepseek-v4-flash-0731": {
|
||||
"input_cost_per_token": 2.2e-07,
|
||||
"cache_read_input_token_cost": 7e-09,
|
||||
"output_cost_per_token": 6.6e-07,
|
||||
},
|
||||
"deepseek-v4p1-flash": {
|
||||
"input_cost_per_token": 2.2e-07,
|
||||
"cache_read_input_token_cost": 7e-09,
|
||||
"output_cost_per_token": 6.6e-07,
|
||||
"supports_vision": True,
|
||||
"max_output_tokens": 393216,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_fireworks_account_prefixed_twins_agree_on_price(model_data):
|
||||
|
|
@ -95,7 +44,7 @@ def test_fireworks_account_prefixed_twins_agree_on_price(model_data):
|
|||
for key, entry in model_data.items():
|
||||
if not key.startswith(prefix):
|
||||
continue
|
||||
bare_key = f"fireworks_ai/{key[len(prefix):]}"
|
||||
bare_key = f"fireworks_ai/{key[len(prefix) :]}"
|
||||
bare_entry = model_data.get(bare_key)
|
||||
if bare_entry is None:
|
||||
continue
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue