mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge litellm_internal_staging into rolling registry PR
This commit is contained in:
commit
23ed208339
56 changed files with 5071 additions and 1156 deletions
|
|
@ -105,7 +105,7 @@
|
|||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38271
|
||||
"limit": 38269
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19584
|
||||
|
|
|
|||
|
|
@ -546,7 +546,7 @@ _key_management_system: Optional["KeyManagementSystem"] = None
|
|||
#### PII MASKING ####
|
||||
output_parse_pii: bool = False
|
||||
#############################################
|
||||
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
|
||||
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map, mark_litellm_import_complete
|
||||
|
||||
model_cost = get_model_cost_map(url=model_cost_map_url)
|
||||
cost_discount_config: Dict[str, float] = {} # Provider-specific cost discounts {"vertex_ai": 0.05} = 5% discount
|
||||
|
|
@ -2405,3 +2405,5 @@ def __getattr__(name: str) -> Any:
|
|||
|
||||
|
||||
# ALL_LITELLM_RESPONSE_TYPES is lazy-loaded via __getattr__ to avoid loading utils at import time
|
||||
|
||||
mark_litellm_import_complete()
|
||||
|
|
|
|||
|
|
@ -6,9 +6,33 @@ be settable from user input. Context variables are scoped to the current
|
|||
asyncio task and cannot be injected via HTTP request bodies.
|
||||
"""
|
||||
|
||||
from collections.abc import Generator
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final
|
||||
|
||||
# When True, suppresses async logging and billing for internal sub-calls
|
||||
# (e.g., emulated file-search steps that make nested LLM calls).
|
||||
is_internal_call: Final[ContextVar[bool]] = ContextVar("is_internal_call", default=False)
|
||||
|
||||
# One request prices its totals, its per-token-type lines and the rates it reports on
|
||||
# separate code paths. Each reads the clock for off-peak pricing, so without a pinned
|
||||
# moment they can land on either side of a window boundary and disagree with each other.
|
||||
_billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", default=None)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def pinned_billing_time(moment: datetime) -> Generator[None]:
|
||||
"""Price every rate lookup inside this block at ``moment`` rather than at each one's own clock read."""
|
||||
token: Final = _billing_time.set(moment)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_billing_time.reset(token)
|
||||
|
||||
|
||||
def current_billing_time() -> datetime:
|
||||
"""The pinned billing moment, or now in UTC outside a pinned block."""
|
||||
pinned: Final = _billing_time.get()
|
||||
return pinned if pinned is not None else datetime.now(timezone.utc)
|
||||
|
|
|
|||
|
|
@ -335,6 +335,7 @@ DEFAULT_SSL_CIPHERS: Final = os.getenv(
|
|||
|
||||
########### v2 Architecture constants for managing writing updates to the database ###########
|
||||
REDIS_UPDATE_BUFFER_KEY: Final = "litellm_spend_update_buffer"
|
||||
REDIS_GATEWAY_REQUESTS_BUFFER_KEY: Final = "litellm_gateway_requests_buffer"
|
||||
REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_spend_update_buffer"
|
||||
REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_team_spend_update_buffer"
|
||||
REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY: Final = "litellm_daily_org_spend_update_buffer"
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import
|
|||
TranscriptionUsageObjectTransformation,
|
||||
)
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
||||
BilledTokenRates,
|
||||
CostCalculatorUtils,
|
||||
_generic_cost_per_character,
|
||||
_get_regional_uplift_multiplier,
|
||||
|
|
@ -1125,6 +1126,7 @@ def _store_cost_breakdown_in_logging_obj(
|
|||
service_tier: str | None = None,
|
||||
data_residency: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
billed_token_rates: BilledTokenRates | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Helper function to store cost breakdown in the logging object.
|
||||
|
|
@ -1169,6 +1171,7 @@ def _store_cost_breakdown_in_logging_obj(
|
|||
service_tier=service_tier,
|
||||
data_residency=data_residency,
|
||||
vertex_location=vertex_location,
|
||||
billed_token_rates=billed_token_rates,
|
||||
)
|
||||
|
||||
except Exception as breakdown_error:
|
||||
|
|
@ -1737,6 +1740,7 @@ def completion_cost(
|
|||
_reasoning_cost: float | None = None
|
||||
_cache_read_cost: float | None = None
|
||||
_cache_creation_cost: float | None = None
|
||||
_billed_token_rates: BilledTokenRates | None = None
|
||||
if cost_per_token_usage_object is not None and model:
|
||||
_breakdown_provider: str | None = (
|
||||
custom_llm_provider if isinstance(custom_llm_provider, str) else None
|
||||
|
|
@ -1748,10 +1752,12 @@ def completion_cost(
|
|||
service_tier=service_tier,
|
||||
data_residency=data_residency,
|
||||
vertex_location=vertex_location,
|
||||
custom_cost_per_token=custom_cost_per_token,
|
||||
)
|
||||
_reasoning_cost = _token_type_breakdown.reasoning_cost
|
||||
_cache_read_cost = _token_type_breakdown.cache_read_cost
|
||||
_cache_creation_cost = _token_type_breakdown.cache_creation_cost
|
||||
_billed_token_rates = _token_type_breakdown.rates
|
||||
_store_cost_breakdown_in_logging_obj(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
prompt_tokens_cost_usd_dollar=prompt_tokens_cost_usd_dollar,
|
||||
|
|
@ -1771,6 +1777,7 @@ def completion_cost(
|
|||
service_tier=service_tier,
|
||||
data_residency=data_residency,
|
||||
vertex_location=vertex_location,
|
||||
billed_token_rates=_billed_token_rates,
|
||||
)
|
||||
|
||||
return _final_cost
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from mcp import ClientSession, McpError, ReadResourceResult, Resource, StdioServ
|
|||
from mcp.client.sse import sse_client
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.shared.session import RequestResponder
|
||||
from typing_extensions import Unpack
|
||||
|
||||
_TransportStreams: TypeAlias = tuple[
|
||||
|
|
@ -56,10 +57,13 @@ def missing_streamable_http_client_error() -> ImportError:
|
|||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
from mcp.types import CallToolResult as MCPCallToolResult
|
||||
from mcp.types import (
|
||||
ClientResult,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
Prompt,
|
||||
ResourceTemplate,
|
||||
ServerNotification,
|
||||
ServerRequest,
|
||||
TextContent,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
|
@ -146,8 +150,8 @@ _SDK_READ_TIMEOUT_CODE: Final = int(httpx.codes.REQUEST_TIMEOUT)
|
|||
otherwise carries JSON-RPC error codes."""
|
||||
|
||||
|
||||
def _as_read_timeout(exc: BaseException) -> TimeoutError | None:
|
||||
"""The session read timeout elapsing, re-expressed as a ``TimeoutError``, or ``None``.
|
||||
def as_mcp_read_timeout(exc: BaseException) -> TimeoutError | None:
|
||||
"""Normalize an MCP SDK read timeout for client and gateway diagnostics, or return ``None``.
|
||||
|
||||
The SDK reports its own elapsed read timeout as ``McpError`` carrying an HTTP status code in a
|
||||
field that otherwise holds JSON-RPC error codes, and it relays an upstream's JSON-RPC error
|
||||
|
|
@ -442,6 +446,18 @@ class MCPClient:
|
|||
in_flight_error: BaseException | None = None
|
||||
try:
|
||||
read_stream, write_stream = transport[0], transport[1]
|
||||
stream_error: Final[asyncio.Future[Exception]] = asyncio.get_running_loop().create_future()
|
||||
|
||||
async def receive_message(
|
||||
message: RequestResponder[ServerRequest, ClientResult] | ServerNotification | Exception,
|
||||
) -> None:
|
||||
if not isinstance(message, (ValueError, httpx.RequestError, OSError)):
|
||||
return
|
||||
if not stream_error.done():
|
||||
stream_error.set_result(message)
|
||||
# The SDK closes pending requests when its message handler raises.
|
||||
raise RuntimeError("MCP response stream failed")
|
||||
|
||||
# Build session kwargs with optional callbacks
|
||||
session_kwargs: Final[dict[str, Any]] = {}
|
||||
if self._sampling_callback is not None:
|
||||
|
|
@ -456,6 +472,7 @@ class MCPClient:
|
|||
read_stream,
|
||||
write_stream,
|
||||
read_timeout_seconds=timedelta(seconds=self.timeout),
|
||||
message_handler=receive_message,
|
||||
**session_kwargs,
|
||||
)
|
||||
session: Final = await session_ctx.__aenter__()
|
||||
|
|
@ -467,6 +484,10 @@ class MCPClient:
|
|||
if isinstance(ins, str) and ins.strip():
|
||||
self._last_initialize_instructions = ins.strip()
|
||||
return await operation(session)
|
||||
except McpError:
|
||||
if stream_error.done():
|
||||
raise stream_error.result()
|
||||
raise
|
||||
finally:
|
||||
try:
|
||||
await session_ctx.__aexit__(None, None, None)
|
||||
|
|
@ -501,11 +522,10 @@ class MCPClient:
|
|||
transport_ctx, http_client = self._create_transport_context()
|
||||
return await self._execute_session_operation(transport_ctx, operation)
|
||||
except Exception as e:
|
||||
read_timeout: Final = _as_read_timeout(e)
|
||||
read_timeout: Final = as_mcp_read_timeout(e)
|
||||
if read_timeout is not None:
|
||||
verbose_logger.warning(
|
||||
"MCP client timed out after %ss waiting for %s to answer; the server accepted the "
|
||||
"request and ended its response stream without a JSON-RPC reply",
|
||||
"MCP client timed out after %ss waiting for a valid MCP response from %s",
|
||||
self.timeout,
|
||||
self.server_url or "stdio",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -850,20 +850,24 @@ class CustomGuardrail(CustomLogger):
|
|||
if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True:
|
||||
return None
|
||||
|
||||
# CHECK IF GUARDRAIL REJECTS THE REQUEST
|
||||
target: Final = self._deployment_hook_target()
|
||||
hook_request_data: Final = {**request_data, "guardrail_to_apply": self} if target is not self else request_data
|
||||
result: Final = await target.async_post_call_success_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
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"),
|
||||
api_key=request_data.get("user_api_key_hash"),
|
||||
request_route=request_data.get("user_api_key_request_route"),
|
||||
),
|
||||
data=hook_request_data,
|
||||
response=response,
|
||||
)
|
||||
try:
|
||||
if target is not self:
|
||||
request_data["guardrail_to_apply"] = self # rebind-ok: dispatch consumes this key
|
||||
result: Final = await target.async_post_call_success_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
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"),
|
||||
api_key=request_data.get("user_api_key_hash"),
|
||||
request_route=request_data.get("user_api_key_request_route"),
|
||||
),
|
||||
data=request_data,
|
||||
response=response,
|
||||
)
|
||||
finally:
|
||||
if target is not self:
|
||||
request_data.pop("guardrail_to_apply", None)
|
||||
|
||||
if not self._is_valid_response_type(result):
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import hashlib
|
|||
import json
|
||||
import os
|
||||
import random
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass, replace
|
||||
|
|
@ -176,6 +177,11 @@ class GetModelCostMap:
|
|||
RETRYABLE_FETCH_STATUS_CODES: Final = frozenset({429, 500, 502, 503, 504})
|
||||
MODEL_COST_MAP_FETCH_MAX_ATTEMPTS: Final = 3
|
||||
MODEL_COST_MAP_FETCH_MAX_WAIT_SECONDS: Final = 30.0
|
||||
_litellm_import_complete = threading.Event()
|
||||
|
||||
|
||||
def mark_litellm_import_complete() -> None:
|
||||
_litellm_import_complete.set()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -314,12 +320,13 @@ async def _fetch_remote_model_cost_map_with_retry(
|
|||
def _fetch_remote_model_cost_map_with_retry_sync(
|
||||
url: str,
|
||||
timeout: int,
|
||||
max_attempts: int,
|
||||
attempts: range,
|
||||
sleep: Callable[[float], None],
|
||||
rng: random.Random,
|
||||
client: _SyncGetClient,
|
||||
) -> ModelCostMapReloadResult:
|
||||
for attempt in range(1, max_attempts + 1):
|
||||
max_attempts: Final = attempts.stop - 1
|
||||
for attempt in attempts:
|
||||
outcome = _attempt_fetch_sync(client=client, url=url, timeout=timeout)
|
||||
if not isinstance(outcome, _FetchAttemptRetryable):
|
||||
return outcome
|
||||
|
|
@ -520,6 +527,68 @@ def _finalize_loaded_model_cost_map(loaded: ModelCostMapReloaded) -> ModelCostMa
|
|||
return replace(loaded, model_cost_map=_finalize_model_cost_map(loaded.model_cost_map))
|
||||
|
||||
|
||||
def adopt_model_cost_map(
|
||||
new_model_cost_map: dict, # mutable-ok: public API preserves the mutable cost-map contract
|
||||
) -> int:
|
||||
import litellm
|
||||
from litellm import utils
|
||||
|
||||
litellm.model_cost = new_model_cost_map
|
||||
utils._invalidate_model_cost_lowercase_map() # pyright: ignore[reportPrivateUsage] # required cache invalidation
|
||||
litellm.add_known_models(model_cost_map=new_model_cost_map)
|
||||
fetched_model_count: Final = len(new_model_cost_map) if new_model_cost_map else 0
|
||||
utils.reapply_runtime_model_cost_registrations()
|
||||
return fetched_model_count
|
||||
|
||||
|
||||
def _retry_remote_fetch_in_background(
|
||||
url: str,
|
||||
timeout: int,
|
||||
max_attempts: int,
|
||||
sleep: Callable[[float], None],
|
||||
rng: random.Random,
|
||||
client: _SyncGetClient,
|
||||
first_outcome: _FetchAttemptRetryable,
|
||||
) -> None:
|
||||
try:
|
||||
first_wait: Final = _next_retry_wait(outcome=first_outcome, attempt=1, max_attempts=max_attempts, rng=rng)
|
||||
if isinstance(first_wait, ModelCostMapReloadUnavailable):
|
||||
return
|
||||
sleep(first_wait)
|
||||
result: Final = _fetch_remote_model_cost_map_with_retry_sync(
|
||||
url=url,
|
||||
timeout=timeout,
|
||||
attempts=range(2, max_attempts + 1),
|
||||
sleep=sleep,
|
||||
rng=rng,
|
||||
client=client,
|
||||
)
|
||||
if isinstance(result, ModelCostMapReloadUnavailable):
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Failed to fetch remote model cost map from %s after %d attempts; keeping local backup",
|
||||
url,
|
||||
max_attempts,
|
||||
)
|
||||
return
|
||||
_litellm_import_complete.wait()
|
||||
if not GetModelCostMap.validate_model_cost_map(
|
||||
fetched_map=result.model_cost_map,
|
||||
backup_model_count=GetModelCostMap._get_backup_model_count(), # pyright: ignore[reportPrivateUsage] # integrity cache
|
||||
):
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Fetched model cost map failed integrity check. Using local backup instead. url=%s",
|
||||
url,
|
||||
)
|
||||
return
|
||||
finalized: Final = _finalize_loaded_model_cost_map(result).model_cost_map
|
||||
_cost_map_source_info.source = "remote"
|
||||
_cost_map_source_info.fallback_reason = None
|
||||
_cost_map_source_info.loaded_at = datetime.now(timezone.utc)
|
||||
adopt_model_cost_map(finalized)
|
||||
except Exception as e:
|
||||
verbose_logger.warning("LiteLLM: Background model cost map retry failed: %s", e)
|
||||
|
||||
|
||||
def get_model_cost_map(
|
||||
url: str,
|
||||
timeout: int = 5,
|
||||
|
|
@ -532,9 +601,7 @@ def get_model_cost_map(
|
|||
Public entry point — returns the model cost map dict.
|
||||
|
||||
1. If ``LITELLM_LOCAL_MODEL_COST_MAP`` is set, uses the local backup only.
|
||||
2. Otherwise fetches from ``url``, retrying transient HTTP errors
|
||||
(429/5xx/transport) with Retry-After-aware backoff, validates
|
||||
integrity, and falls back to the local backup on any failure.
|
||||
2. Otherwise fetches from ``url``, retrying transient errors in a background thread.
|
||||
|
||||
Only the backup model count is cached (a single int) for validation.
|
||||
The full backup dict is only parsed when it must be *returned* as a
|
||||
|
|
@ -553,24 +620,34 @@ def get_model_cost_map(
|
|||
_cost_map_source_info.url = url
|
||||
_cost_map_source_info.is_env_forced = False
|
||||
|
||||
result: Final = _fetch_remote_model_cost_map_with_retry_sync(
|
||||
url=url,
|
||||
timeout=timeout,
|
||||
max_attempts=max_attempts,
|
||||
sleep=sleep,
|
||||
rng=rng if rng is not None else random.Random(),
|
||||
client=client if client is not None else httpx,
|
||||
)
|
||||
if isinstance(result, ModelCostMapReloadUnavailable):
|
||||
fetch_client: Final = client if client is not None else httpx
|
||||
fetch_rng: Final = rng if rng is not None else random.Random()
|
||||
outcome: Final = _attempt_fetch_sync(client=fetch_client, url=url, timeout=timeout)
|
||||
if isinstance(outcome, _FetchAttemptRetryable) and max_attempts > 1:
|
||||
threading.Thread(
|
||||
target=_retry_remote_fetch_in_background,
|
||||
kwargs={ # mutable-ok: threading requires a mutable keyword-arguments mapping
|
||||
"url": url,
|
||||
"timeout": timeout,
|
||||
"max_attempts": max_attempts,
|
||||
"sleep": sleep,
|
||||
"rng": fetch_rng,
|
||||
"client": fetch_client,
|
||||
"first_outcome": outcome,
|
||||
},
|
||||
name="litellm-model-cost-map-retry",
|
||||
daemon=True,
|
||||
).start()
|
||||
if not isinstance(outcome, ModelCostMapReloaded):
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: Failed to fetch remote model cost map from %s: %s. Falling back to local backup.",
|
||||
url,
|
||||
result.reason,
|
||||
outcome.reason,
|
||||
)
|
||||
_cost_map_source_info.source = "local"
|
||||
_cost_map_source_info.fallback_reason = f"Remote fetch failed: {result.reason}"
|
||||
_cost_map_source_info.fallback_reason = f"Remote fetch failed: {outcome.reason}"
|
||||
return _finalize_loaded_model_cost_map(GetModelCostMap.load_local_model_cost_map_with_revision()).model_cost_map
|
||||
content: Final = result.model_cost_map
|
||||
content: Final = outcome.model_cost_map
|
||||
|
||||
# Validate using cached count (cheap int comparison, no file I/O)
|
||||
if not GetModelCostMap.validate_model_cost_map(
|
||||
|
|
@ -587,4 +664,4 @@ def get_model_cost_map(
|
|||
|
||||
_cost_map_source_info.source = "remote"
|
||||
_cost_map_source_info.fallback_reason = None
|
||||
return _finalize_loaded_model_cost_map(result).model_cost_map
|
||||
return _finalize_loaded_model_cost_map(outcome).model_cost_map
|
||||
|
|
|
|||
|
|
@ -203,6 +203,7 @@ if TYPE_CHECKING:
|
|||
|
||||
from litellm.integrations.otel.logger import OpenTelemetryV2
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates
|
||||
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
|
||||
try:
|
||||
from litellm_enterprise.enterprise_callbacks.callback_controls import (
|
||||
|
|
@ -590,6 +591,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
# Initialize cost breakdown field
|
||||
self.cost_breakdown: CostBreakdown | None = None
|
||||
self.billed_token_rates: BilledTokenRates | None = None
|
||||
|
||||
# Init Caching related details
|
||||
self.caching_details: CachingDetails | None = None
|
||||
|
|
@ -1587,6 +1589,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
service_tier: str | None = None,
|
||||
data_residency: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
billed_token_rates: "BilledTokenRates | None" = None,
|
||||
) -> None:
|
||||
"""
|
||||
Helper method to store cost breakdown in the logging object.
|
||||
|
|
@ -1606,8 +1609,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
service_tier: Tier the costs above were priced on, already resolved
|
||||
data_residency: Region uplift the costs above were priced on, already resolved
|
||||
vertex_location: Vertex AI location the costs above were priced on, already resolved
|
||||
billed_token_rates: Per-token rates the costs above were billed at, already resolved
|
||||
"""
|
||||
|
||||
self.billed_token_rates = billed_token_rates
|
||||
self.cost_breakdown = CostBreakdown(
|
||||
input_cost=input_cost,
|
||||
output_cost=output_cost,
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from typing import Any, Final, Literal, TypedDict, cast
|
|||
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import current_billing_time
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import (
|
||||
select_tier_for_input,
|
||||
|
|
@ -19,6 +20,7 @@ from litellm.types.utils import (
|
|||
CacheCreationTokenDetails,
|
||||
CallTypes,
|
||||
CompletionTokensDetailsWrapper,
|
||||
CostPerToken,
|
||||
DataResidency,
|
||||
ImageResponse,
|
||||
ModelInfo,
|
||||
|
|
@ -305,7 +307,7 @@ def _is_within_off_peak_window(off_peak_hours_utc: str | Sequence[str], current_
|
|||
than being localised, so callers must pass datetime.now(timezone.utc), never datetime.now(),
|
||||
or every window shifts by the host's offset.
|
||||
"""
|
||||
reference: Final = current_time if current_time is not None else datetime.now(timezone.utc)
|
||||
reference: Final = current_time if current_time is not None else current_billing_time()
|
||||
now: Final = (reference.astimezone(timezone.utc) if reference.tzinfo is not None else reference).time()
|
||||
windows: Final = (off_peak_hours_utc,) if isinstance(off_peak_hours_utc, str) else off_peak_hours_utc
|
||||
for window in windows:
|
||||
|
|
@ -392,7 +394,7 @@ def _is_off_peak(off_peak: Mapping[str, object], current_time: datetime | None =
|
|||
rules: the flat hours_utc windows, which apply every day, or any entry in windows, whose
|
||||
hours apply only on its weekdays.
|
||||
"""
|
||||
reference: Final = current_time if current_time is not None else datetime.now(timezone.utc)
|
||||
reference: Final = current_time if current_time is not None else current_billing_time()
|
||||
reference_utc: Final = (
|
||||
reference.astimezone(timezone.utc) if reference.tzinfo is not None else reference.replace(tzinfo=timezone.utc)
|
||||
)
|
||||
|
|
@ -1195,7 +1197,7 @@ def generic_cost_per_token(
|
|||
usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens - video_tokens, 0
|
||||
)
|
||||
|
||||
billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
|
||||
billing_time: Final = current_time if current_time is not None else current_billing_time()
|
||||
(
|
||||
prompt_base_cost,
|
||||
completion_base_cost,
|
||||
|
|
@ -1309,42 +1311,90 @@ def _coerce_token_count(value: object) -> int:
|
|||
return value if isinstance(value, int) and value > 0 else 0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BilledTokenRates:
|
||||
"""Per-token rates one request's usage bills at, after token tiers, off-peak windows and the
|
||||
regional multipliers the totals apply, so each cost line equals its token count times its rate."""
|
||||
|
||||
input_cost_per_token: float
|
||||
output_cost_per_token: float
|
||||
cache_read_input_token_cost: float
|
||||
cache_creation_input_token_cost: float
|
||||
cache_creation_input_token_cost_above_1hr: float
|
||||
output_cost_per_reasoning_token: float
|
||||
|
||||
def scaled(self, multiplier: float) -> "BilledTokenRates":
|
||||
if multiplier == 1.0:
|
||||
return self
|
||||
return BilledTokenRates(
|
||||
input_cost_per_token=self.input_cost_per_token * multiplier,
|
||||
output_cost_per_token=self.output_cost_per_token * multiplier,
|
||||
cache_read_input_token_cost=self.cache_read_input_token_cost * multiplier,
|
||||
cache_creation_input_token_cost=self.cache_creation_input_token_cost * multiplier,
|
||||
cache_creation_input_token_cost_above_1hr=self.cache_creation_input_token_cost_above_1hr * multiplier,
|
||||
output_cost_per_reasoning_token=self.output_cost_per_reasoning_token * multiplier,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TokenTypeCostBreakdown:
|
||||
reasoning_cost: float
|
||||
cache_read_cost: float
|
||||
cache_creation_cost: float
|
||||
rates: BilledTokenRates | None = None
|
||||
"""Rates these lines were billed at, so a caller reporting both cannot resolve them a second,
|
||||
differently-argued way. None when the model's pricing could not be resolved."""
|
||||
|
||||
|
||||
def get_token_type_cost_breakdown(
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
def _reasoning_token_count(usage: Usage) -> int:
|
||||
parsed: Final = (
|
||||
parse_completion_tokens_details(usage)["reasoning_tokens"] if usage.completion_tokens_details is not None else 0
|
||||
)
|
||||
return parsed or _coerce_token_count(getattr(usage, "reasoning_tokens", 0))
|
||||
|
||||
|
||||
def _cache_token_counts(usage: Usage) -> tuple[int, int, CacheCreationTokenDetails | None]:
|
||||
"""(cache read tokens, cache creation tokens, cache creation details): read from prompt_tokens_details
|
||||
first, then the private top-level counters the Usage constructor mirrors cache tokens onto for
|
||||
providers/callers that bypass the details."""
|
||||
parsed: Final = parse_prompt_tokens_details(usage) if usage.prompt_tokens_details is not None else None
|
||||
parsed_read: Final = parsed["cache_hit_tokens"] if parsed is not None else 0
|
||||
parsed_creation: Final = parsed["cache_creation_tokens"] if parsed is not None else 0
|
||||
return (
|
||||
parsed_read or _coerce_token_count(getattr(usage, "_cache_read_input_tokens", 0)),
|
||||
parsed_creation or _coerce_token_count(getattr(usage, "_cache_creation_input_tokens", 0)),
|
||||
parsed["cache_creation_token_details"] if parsed is not None else None,
|
||||
)
|
||||
|
||||
|
||||
def _custom_pricing_rates(custom_cost_per_token: CostPerToken) -> BilledTokenRates:
|
||||
"""Flat custom pricing has no tiers, uplifts or reasoning rate: cache tokens bill at the configured
|
||||
cache rates (else the input rate) and reasoning at the output rate, as _cost_per_token_custom_pricing_helper does."""
|
||||
input_rate: Final = custom_cost_per_token["input_cost_per_token"]
|
||||
output_rate: Final = custom_cost_per_token["output_cost_per_token"]
|
||||
cache_creation_rate: Final = custom_cost_per_token.get("cache_creation_input_token_cost", input_rate)
|
||||
return BilledTokenRates(
|
||||
input_cost_per_token=input_rate,
|
||||
output_cost_per_token=output_rate,
|
||||
cache_read_input_token_cost=custom_cost_per_token.get("cache_read_input_token_cost", input_rate),
|
||||
cache_creation_input_token_cost=cache_creation_rate,
|
||||
cache_creation_input_token_cost_above_1hr=cache_creation_rate,
|
||||
output_cost_per_reasoning_token=output_rate,
|
||||
)
|
||||
|
||||
|
||||
def _cost_map_billed_rates(
|
||||
model_info: ModelInfo,
|
||||
usage: Usage,
|
||||
service_tier: str | None = None,
|
||||
data_residency: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
current_time: datetime | None = None,
|
||||
) -> TokenTypeCostBreakdown:
|
||||
"""
|
||||
Provider-agnostic cost of reasoning and cache tokens, derived from the usage
|
||||
object and model pricing alone.
|
||||
|
||||
This works for every provider, including Perplexity/Cerebras/Dashscope whose
|
||||
cost calculators bypass ``generic_cost_per_token``, because cache tokens always
|
||||
land on ``prompt_tokens_details`` (via the Usage constructor and provider
|
||||
transformations) and reasoning tokens on ``completion_tokens_details``. It reuses
|
||||
the same rate-resolution primitives as the total-cost path so the breakdown can
|
||||
never drift from the totals. Returns zeros (never raises) when the model or its
|
||||
pricing cannot be resolved.
|
||||
"""
|
||||
try:
|
||||
model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
except Exception:
|
||||
return TokenTypeCostBreakdown(0.0, 0.0, 0.0)
|
||||
|
||||
billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
|
||||
custom_llm_provider: str | None,
|
||||
service_tier: str | None,
|
||||
data_residency: str | None,
|
||||
vertex_location: str | None,
|
||||
current_time: datetime | None,
|
||||
) -> BilledTokenRates:
|
||||
billing_time: Final = current_time if current_time is not None else current_billing_time()
|
||||
(
|
||||
_prompt_base_cost,
|
||||
prompt_base_cost,
|
||||
completion_base_cost,
|
||||
cache_creation_cost_rate,
|
||||
cache_creation_cost_above_1hr_rate,
|
||||
|
|
@ -1356,13 +1406,6 @@ def get_token_type_cost_breakdown(
|
|||
current_time=billing_time,
|
||||
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
|
||||
)
|
||||
|
||||
reasoning_tokens = (
|
||||
parse_completion_tokens_details(usage)["reasoning_tokens"] if usage.completion_tokens_details is not None else 0
|
||||
)
|
||||
if not reasoning_tokens:
|
||||
reasoning_tokens = _coerce_token_count(getattr(usage, "reasoning_tokens", 0))
|
||||
|
||||
reasoning_rate: Final = _resolve_billed_reasoning_rate(
|
||||
model_info=model_info,
|
||||
usage=usage,
|
||||
|
|
@ -1370,57 +1413,103 @@ def get_token_type_cost_breakdown(
|
|||
completion_base_cost=completion_base_cost,
|
||||
current_time=billing_time,
|
||||
)
|
||||
reasoning_cost = float(reasoning_tokens) * reasoning_rate
|
||||
multiplier: Final = (
|
||||
_get_regional_uplift_multiplier(model_info, data_residency)
|
||||
* get_vertex_regional_endpoint_uplift(model_info, vertex_location)
|
||||
* get_provider_specific_geo_multiplier(model_info=model_info, usage=usage)
|
||||
)
|
||||
return BilledTokenRates(
|
||||
input_cost_per_token=prompt_base_cost,
|
||||
output_cost_per_token=completion_base_cost,
|
||||
cache_read_input_token_cost=cache_read_cost_rate,
|
||||
cache_creation_input_token_cost=cache_creation_cost_rate,
|
||||
cache_creation_input_token_cost_above_1hr=cache_creation_cost_above_1hr_rate,
|
||||
output_cost_per_reasoning_token=reasoning_rate,
|
||||
).scaled(multiplier)
|
||||
|
||||
cache_read_tokens = 0
|
||||
cache_creation_tokens = 0
|
||||
cache_creation_token_details: CacheCreationTokenDetails | None = None
|
||||
if usage.prompt_tokens_details is not None:
|
||||
prompt_tokens_details: Final = parse_prompt_tokens_details(usage)
|
||||
cache_read_tokens = prompt_tokens_details["cache_hit_tokens"]
|
||||
cache_creation_tokens = prompt_tokens_details["cache_creation_tokens"]
|
||||
cache_creation_token_details = prompt_tokens_details["cache_creation_token_details"]
|
||||
# Fall back to the private top-level counters the Usage constructor mirrors cache
|
||||
# tokens onto, so providers/callers that bypass prompt_tokens_details are covered.
|
||||
if not cache_read_tokens:
|
||||
cache_read_tokens = _coerce_token_count(getattr(usage, "_cache_read_input_tokens", 0))
|
||||
if not cache_creation_tokens:
|
||||
cache_creation_tokens = _coerce_token_count(getattr(usage, "_cache_creation_input_tokens", 0))
|
||||
|
||||
cache_read_cost = float(cache_read_tokens) * cache_read_cost_rate
|
||||
cache_creation_cost = calculate_cache_writing_cost(
|
||||
cache_creation_tokens=cache_creation_tokens,
|
||||
cache_creation_token_details=cache_creation_token_details,
|
||||
cache_creation_cost_above_1hr=cache_creation_cost_above_1hr_rate,
|
||||
cache_creation_cost=cache_creation_cost_rate,
|
||||
def get_billed_token_rates(
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
usage: Usage,
|
||||
service_tier: str | None = None,
|
||||
data_residency: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
current_time: datetime | None = None,
|
||||
custom_cost_per_token: CostPerToken | None = None,
|
||||
) -> BilledTokenRates | None:
|
||||
"""Rates the cost calculator bills ``usage`` at, resolved exactly as the totals and the token-type
|
||||
breakdown resolve them. None when the model's pricing cannot be resolved."""
|
||||
if custom_cost_per_token is not None:
|
||||
return _custom_pricing_rates(custom_cost_per_token)
|
||||
try:
|
||||
model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
except Exception:
|
||||
return None
|
||||
return _cost_map_billed_rates(
|
||||
model_info=model_info,
|
||||
usage=usage,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
service_tier=service_tier,
|
||||
data_residency=data_residency,
|
||||
vertex_location=vertex_location,
|
||||
current_time=current_time,
|
||||
)
|
||||
|
||||
# Apply the same flat regional-processing uplift the totals get, so per-type
|
||||
# costs stay reconciled with input_cost/output_cost for regionalized OpenAI hosts.
|
||||
uplift: Final = _get_regional_uplift_multiplier(model_info, data_residency)
|
||||
if uplift != 1.0:
|
||||
reasoning_cost *= uplift
|
||||
cache_read_cost *= uplift
|
||||
cache_creation_cost *= uplift
|
||||
|
||||
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
|
||||
if vertex_uplift != 1.0:
|
||||
reasoning_cost *= vertex_uplift
|
||||
cache_read_cost *= vertex_uplift
|
||||
cache_creation_cost *= vertex_uplift
|
||||
def get_token_type_cost_breakdown(
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
usage: Usage,
|
||||
service_tier: str | None = None,
|
||||
data_residency: str | None = None,
|
||||
vertex_location: str | None = None,
|
||||
current_time: datetime | None = None,
|
||||
custom_cost_per_token: CostPerToken | None = None,
|
||||
) -> TokenTypeCostBreakdown:
|
||||
"""
|
||||
Provider-agnostic cost of reasoning and cache tokens, derived from the usage
|
||||
object and model pricing alone.
|
||||
|
||||
# Mirror the provider-specific geo uplift (e.g. Anthropic us: 1.1) the totals
|
||||
# apply, so cache and reasoning line items stay reconciled with them.
|
||||
geo_multiplier: Final = get_provider_specific_geo_multiplier(model_info=model_info, usage=usage)
|
||||
if geo_multiplier != 1.0:
|
||||
reasoning_cost *= geo_multiplier
|
||||
cache_read_cost *= geo_multiplier
|
||||
cache_creation_cost *= geo_multiplier
|
||||
This works for every provider, including Perplexity/Cerebras/Dashscope whose
|
||||
cost calculators bypass ``generic_cost_per_token``, because cache tokens always
|
||||
land on ``prompt_tokens_details`` (via the Usage constructor and provider
|
||||
transformations) and reasoning tokens on ``completion_tokens_details``. It reuses
|
||||
the same rate resolution as the total-cost path (``get_billed_token_rates``) so the
|
||||
breakdown can never drift from the totals. A deployment billed by
|
||||
``custom_cost_per_token`` is priced from those flat rates instead of the cost map and,
|
||||
like its totals, bills cache writes flat rather than by their 5m/1h split.
|
||||
Returns zeros (never raises) when the model or its pricing cannot be resolved.
|
||||
"""
|
||||
rates: Final = get_billed_token_rates(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
usage=usage,
|
||||
service_tier=service_tier,
|
||||
data_residency=data_residency,
|
||||
vertex_location=vertex_location,
|
||||
current_time=current_time,
|
||||
custom_cost_per_token=custom_cost_per_token,
|
||||
)
|
||||
if rates is None:
|
||||
return TokenTypeCostBreakdown(0.0, 0.0, 0.0)
|
||||
|
||||
cache_read_tokens, cache_creation_tokens, cache_creation_token_details = _cache_token_counts(usage)
|
||||
cache_creation_cost: Final = (
|
||||
float(cache_creation_tokens) * rates.cache_creation_input_token_cost
|
||||
if custom_cost_per_token is not None
|
||||
else calculate_cache_writing_cost(
|
||||
cache_creation_tokens=cache_creation_tokens,
|
||||
cache_creation_token_details=cache_creation_token_details,
|
||||
cache_creation_cost_above_1hr=rates.cache_creation_input_token_cost_above_1hr,
|
||||
cache_creation_cost=rates.cache_creation_input_token_cost,
|
||||
)
|
||||
)
|
||||
return TokenTypeCostBreakdown(
|
||||
reasoning_cost=reasoning_cost,
|
||||
cache_read_cost=cache_read_cost,
|
||||
reasoning_cost=float(_reasoning_token_count(usage)) * rates.output_cost_per_reasoning_token,
|
||||
cache_read_cost=float(cache_read_tokens) * rates.cache_read_input_token_cost,
|
||||
cache_creation_cost=cache_creation_cost,
|
||||
rates=rates,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -3,12 +3,15 @@ import importlib
|
|||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from traceback import walk_tb
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from uuid import uuid4
|
||||
|
||||
import anyio
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from pydantic import ValidationError
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -30,6 +33,8 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import (
|
|||
list_fault_http_status,
|
||||
outcome_wire_value,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import _redact_mcp_resource_url
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
|
||||
acting_user_auth,
|
||||
build_effective_auth_contexts,
|
||||
|
|
@ -78,11 +83,39 @@ _MCP_GUARDRAIL_REJECTIONS: Final = (
|
|||
|
||||
|
||||
def _connection_error_message(exc: BaseException, url: str | None, timeout_seconds: float) -> str:
|
||||
reference: Final = uuid4().hex
|
||||
verbose_logger.error(
|
||||
"MCP connection test failed (reference=%s): %s",
|
||||
reference,
|
||||
tuple(
|
||||
(
|
||||
type(cause).__name__,
|
||||
tuple(
|
||||
(frame.f_code.co_filename, lineno, frame.f_code.co_name)
|
||||
for frame, lineno in walk_tb(cause.__traceback__)
|
||||
),
|
||||
)
|
||||
for cause in iter_exception_tree(exc)
|
||||
),
|
||||
)
|
||||
return next(
|
||||
(
|
||||
message
|
||||
for cause in iter_exception_tree(exc)
|
||||
if (message := _known_connection_error_message(cause, url, timeout_seconds)) is not None
|
||||
),
|
||||
"An unexpected error occurred while testing the MCP connection. "
|
||||
f"Retry; if it persists, share reference {reference} with your gateway administrator.",
|
||||
)
|
||||
|
||||
|
||||
def _known_connection_error_message(exc: BaseException, url: str | None, timeout_seconds: float) -> str | None:
|
||||
if isinstance(exc, MCPServerURLCredentialsError):
|
||||
return str(exc.detail)
|
||||
if isinstance(exc, TimeoutError):
|
||||
return (
|
||||
f"Failed to connect to MCP server: no response from {url or 'the server'} "
|
||||
"Failed to connect to MCP server: no valid MCP response received from "
|
||||
f"{_redact_mcp_resource_url(url) or 'the server'} "
|
||||
f"within {timeout_seconds:.0f}s. Check that the LiteLLM proxy can reach this URL "
|
||||
"from its network (DNS, egress rules, firewalls) and that the server answers MCP requests."
|
||||
)
|
||||
|
|
@ -99,13 +132,45 @@ def _connection_error_message(exc: BaseException, url: str | None, timeout_secon
|
|||
return "Failed to connect to MCP server: the connection timed out."
|
||||
if isinstance(exc, httpx.HTTPStatusError):
|
||||
return f"Failed to connect to MCP server: it returned HTTP {exc.response.status_code}."
|
||||
return "Failed to connect to MCP server. Check proxy logs for details."
|
||||
if isinstance(exc, (httpx.NetworkError, httpx.RemoteProtocolError, ConnectionError)):
|
||||
return (
|
||||
"Failed to connect to MCP server: the connection was interrupted. "
|
||||
"Check the server and network connection, then retry."
|
||||
)
|
||||
if isinstance(exc, ValueError) and str(exc).startswith("Unexpected content type:"):
|
||||
return (
|
||||
"Failed to connect to MCP server: the endpoint returned an unsupported content type. "
|
||||
"Check that the URL is an MCP endpoint, not a web page, and matches the selected transport."
|
||||
)
|
||||
if isinstance(exc, ValidationError) and exc.title in ("JSONRPCMessage", "InitializeResult", "ListToolsResult"):
|
||||
return (
|
||||
"Failed to connect to MCP server: the endpoint returned invalid JSON or an invalid MCP response. "
|
||||
"Check the MCP endpoint URL and the server's protocol implementation."
|
||||
)
|
||||
if MCP_AVAILABLE and isinstance(exc, McpError):
|
||||
if exc.error.code == -32000 and exc.error.message == "Connection closed":
|
||||
return (
|
||||
"Failed to connect to MCP server: the connection was closed before the request completed. "
|
||||
"Check that the server stays running and returns a complete MCP response, then retry."
|
||||
)
|
||||
if exc.error.code == 32600 and exc.error.message == "Session terminated":
|
||||
return (
|
||||
"Failed to connect to MCP server: the MCP session was terminated. "
|
||||
"Check that the URL points to an MCP endpoint and matches the selected transport, "
|
||||
"then retry to start a new session."
|
||||
)
|
||||
return (
|
||||
f"Failed to connect to MCP server: the MCP request failed (JSON-RPC code {exc.error.code}). "
|
||||
"Check that the endpoint supports MCP initialization and tool listing, and check the upstream server logs."
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
if MCP_AVAILABLE:
|
||||
from mcp.shared.exceptions import McpError
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.experimental_mcp_client.client import MCPClient, as_mcp_read_timeout
|
||||
from litellm.llms.litellm_proxy.skills.skill_search import (
|
||||
DEFAULT_SKILL_SEARCH_TOP_K,
|
||||
)
|
||||
|
|
@ -1342,11 +1407,18 @@ if MCP_AVAILABLE:
|
|||
except (KeyboardInterrupt, SystemExit, asyncio.CancelledError):
|
||||
raise
|
||||
except BaseException as e:
|
||||
verbose_logger.error("Error in MCP operation: %s", e, exc_info=True)
|
||||
effective_timeout: Final = (
|
||||
min(request.timeout if request.timeout is not None else MCP_CLIENT_TIMEOUT, timeout_seconds)
|
||||
if any(
|
||||
isinstance(cause, McpError) and as_mcp_read_timeout(cause) is not None
|
||||
for cause in iter_exception_tree(e)
|
||||
)
|
||||
else timeout_seconds
|
||||
)
|
||||
return {
|
||||
"status": "error",
|
||||
"error": True,
|
||||
"message": _connection_error_message(e, request.url, timeout_seconds),
|
||||
"message": _connection_error_message(e, request.url, effective_timeout),
|
||||
}
|
||||
|
||||
async def _preview_openapi_tools(spec_path: str) -> dict:
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import json
|
|||
import os
|
||||
from collections.abc import Callable, Mapping
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple
|
||||
|
||||
import httpx
|
||||
|
|
@ -1294,6 +1295,13 @@ class UpdateKeyRequest(KeyRequestBase):
|
|||
rotation_interval: str | None = None
|
||||
organization_id: str | None = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def drop_blank_team_id(cls, values: object) -> object:
|
||||
if isinstance(values, Mapping) and values.get("team_id") == "":
|
||||
return MappingProxyType({k: v for k, v in values.items() if k != "team_id"})
|
||||
return values
|
||||
|
||||
@field_validator("organization_id", mode="before")
|
||||
@classmethod
|
||||
def treat_cleared_organization_id_as_unset(cls, v: object) -> object:
|
||||
|
|
@ -5148,9 +5156,26 @@ class CostEstimateRequest(LiteLLMPydanticObjectBase):
|
|||
model: str = Field(description="Model name (from /model_group/info)")
|
||||
input_tokens: int = Field(description="Expected input tokens per request", ge=0)
|
||||
output_tokens: int = Field(description="Expected output tokens per request", ge=0)
|
||||
cache_read_input_tokens: int = Field(
|
||||
default=0, description="Input tokens read from the prompt cache; counted within input_tokens", ge=0
|
||||
)
|
||||
cache_creation_input_tokens: int = Field(
|
||||
default=0, description="Input tokens written to the prompt cache; counted within input_tokens", ge=0
|
||||
)
|
||||
reasoning_tokens: int = Field(
|
||||
default=0, description="Reasoning tokens the model emits; counted within output_tokens", ge=0
|
||||
)
|
||||
num_requests_per_day: int | None = Field(default=None, description="Number of requests per day", ge=0)
|
||||
num_requests_per_month: int | None = Field(default=None, description="Number of requests per month", ge=0)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_token_subsets(self) -> "CostEstimateRequest":
|
||||
if self.cache_read_input_tokens + self.cache_creation_input_tokens > self.input_tokens:
|
||||
raise ValueError("cache_read_input_tokens plus cache_creation_input_tokens cannot exceed input_tokens")
|
||||
if self.reasoning_tokens > self.output_tokens:
|
||||
raise ValueError("reasoning_tokens cannot exceed output_tokens")
|
||||
return self
|
||||
|
||||
|
||||
class CostEstimateResponse(LiteLLMPydanticObjectBase):
|
||||
"""Response body for /cost/estimate endpoint."""
|
||||
|
|
@ -5158,6 +5183,9 @@ class CostEstimateResponse(LiteLLMPydanticObjectBase):
|
|||
model: str
|
||||
input_tokens: int
|
||||
output_tokens: int
|
||||
cache_read_input_tokens: int = 0
|
||||
cache_creation_input_tokens: int = 0
|
||||
reasoning_tokens: int = 0
|
||||
num_requests_per_day: int | None = None
|
||||
num_requests_per_month: int | None = None
|
||||
# Per-request costs
|
||||
|
|
@ -5165,17 +5193,33 @@ class CostEstimateResponse(LiteLLMPydanticObjectBase):
|
|||
input_cost_per_request: float = Field(description="Input token cost per request (before margin)")
|
||||
output_cost_per_request: float = Field(description="Output token cost per request (before margin)")
|
||||
margin_cost_per_request: float = Field(default=0.0, description="Margin/fee added per request")
|
||||
cache_read_cost_per_request: float = Field(default=0.0, description="Cache-read share of input_cost_per_request")
|
||||
cache_creation_cost_per_request: float = Field(
|
||||
default=0.0, description="Cache-write share of input_cost_per_request"
|
||||
)
|
||||
reasoning_cost_per_request: float = Field(default=0.0, description="Reasoning share of output_cost_per_request")
|
||||
# Daily costs (if num_requests_per_day provided)
|
||||
daily_cost: float | None = Field(default=None, description="Total daily cost (includes margin)")
|
||||
daily_input_cost: float | None = Field(default=None, description="Daily input token cost")
|
||||
daily_output_cost: float | None = Field(default=None, description="Daily output token cost")
|
||||
daily_margin_cost: float | None = Field(default=None, description="Daily margin/fee")
|
||||
daily_cache_read_cost: float | None = Field(default=None, description="Cache-read share of daily_input_cost")
|
||||
daily_cache_creation_cost: float | None = Field(default=None, description="Cache-write share of daily_input_cost")
|
||||
daily_reasoning_cost: float | None = Field(default=None, description="Reasoning share of daily_output_cost")
|
||||
# Monthly costs (if num_requests_per_month provided)
|
||||
monthly_cost: float | None = Field(default=None, description="Total monthly cost (includes margin)")
|
||||
monthly_input_cost: float | None = Field(default=None, description="Monthly input token cost")
|
||||
monthly_output_cost: float | None = Field(default=None, description="Monthly output token cost")
|
||||
monthly_margin_cost: float | None = Field(default=None, description="Monthly margin/fee")
|
||||
# Pricing info
|
||||
input_cost_per_token: float | None = None
|
||||
output_cost_per_token: float | None = None
|
||||
monthly_cache_read_cost: float | None = Field(default=None, description="Cache-read share of monthly_input_cost")
|
||||
monthly_cache_creation_cost: float | None = Field(
|
||||
default=None, description="Cache-write share of monthly_input_cost"
|
||||
)
|
||||
monthly_reasoning_cost: float | None = Field(default=None, description="Reasoning share of monthly_output_cost")
|
||||
# Pricing info: the rates this request's usage bills at, after token tiers and regional multipliers
|
||||
input_cost_per_token: float | None = Field(default=None, description="Rate billed per input token")
|
||||
output_cost_per_token: float | None = Field(default=None, description="Rate billed per output token")
|
||||
cache_read_input_token_cost: float | None = Field(default=None, description="Rate billed per cache-read token")
|
||||
cache_creation_input_token_cost: float | None = Field(default=None, description="Rate billed per cache-write token")
|
||||
output_cost_per_reasoning_token: float | None = Field(default=None, description="Rate billed per reasoning token")
|
||||
provider: str | None = None
|
||||
|
|
|
|||
|
|
@ -475,6 +475,7 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None
|
|||
|
||||
|
||||
_NO_MODEL_INFO: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
_TEAM_GRANT_RELATIONS: Final[Mapping[str, object]] = MappingProxyType({"litellm_model_table": True})
|
||||
|
||||
|
||||
def _has_ptu_flat_cost(model: str, llm_router: "Router") -> bool:
|
||||
|
|
@ -2858,7 +2859,9 @@ class TeamNotFoundError(HTTPException):
|
|||
async def _get_team_db_check(
|
||||
team_id: str, prisma_client: PrismaClient, team_id_upsert: bool | None = None
|
||||
) -> "_PrismaTeamRow | None":
|
||||
response = await _team_table(TeamRepository(prisma_client)).find_unique(where={"team_id": team_id})
|
||||
response = await _team_table(TeamRepository(prisma_client)).find_unique(
|
||||
where={"team_id": team_id}, include=_TEAM_GRANT_RELATIONS
|
||||
)
|
||||
|
||||
if response is None and team_id_upsert:
|
||||
from litellm.proxy.management_endpoints.team_endpoints import new_team
|
||||
|
|
@ -3158,7 +3161,9 @@ async def get_team_object_by_alias(
|
|||
|
||||
# Query database by team_alias
|
||||
try:
|
||||
teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(where={"team_alias": team_alias})
|
||||
teams: Final = await _team_table(TeamRepository(prisma_client)).find_many(
|
||||
where={"team_alias": team_alias}, include=_TEAM_GRANT_RELATIONS
|
||||
)
|
||||
|
||||
if not teams:
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -53,6 +53,7 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.auth_checks import can_team_access_model
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.auth.team_grants import team_model_aliases
|
||||
from litellm.proxy.common_utils.user_api_key_cache import (
|
||||
UserApiKeyCache,
|
||||
get_management_object_ttl,
|
||||
|
|
@ -1595,7 +1596,7 @@ class JWTAuthManager:
|
|||
model=requested_model,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=None,
|
||||
team_model_aliases=team_model_aliases(team_object),
|
||||
)
|
||||
):
|
||||
is_allowed = allowed_routes_check(
|
||||
|
|
@ -2132,7 +2133,7 @@ class JWTAuthManager:
|
|||
model=requested_model,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=None,
|
||||
team_model_aliases=team_model_aliases(team_object),
|
||||
)
|
||||
except ProxyException:
|
||||
continue
|
||||
|
|
|
|||
122
litellm/proxy/auth/team_grants.py
Normal file
122
litellm/proxy/auth/team_grants.py
Normal file
|
|
@ -0,0 +1,122 @@
|
|||
"""Project a team row (plus the caller's membership in it) onto the ``team_*`` fields of ``UserAPIKeyAuth``.
|
||||
|
||||
The virtual-key path gets these fields for free from the combined-view SQL join. Every other auth path
|
||||
starts from a ``LiteLLM_TeamTable`` object instead and has to copy them over by hand, which is how JWT
|
||||
callers kept losing grants (aliases, permissions, limits) one field at a time. Build the badge through
|
||||
``team_grants`` and the two paths cannot drift.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Annotated, Final
|
||||
|
||||
from pydantic import BaseModel, BeforeValidator, ConfigDict, TypeAdapter, ValidationError
|
||||
from pydantic.main import IncEx
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
Member,
|
||||
)
|
||||
|
||||
_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
|
||||
_JSON_COLUMNS: Final[Mapping[str, IncEx | bool]] = MappingProxyType(
|
||||
{"metadata": True, "litellm_model_table": MappingProxyType({"model_aliases": True})}
|
||||
)
|
||||
|
||||
|
||||
def _decode_model_aliases(value: object) -> object:
|
||||
"""``LiteLLM_ModelTable.model_aliases`` is typed ``str | dict``; writers hand Prisma ``json.dumps(...)``, so take both."""
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
try:
|
||||
return _MODEL_ALIASES_ADAPTER.validate_json(value)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
class TeamModelAliasTable(BaseModel):
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
model_aliases: Annotated[Mapping[str, str] | None, BeforeValidator(_decode_model_aliases)] = None
|
||||
|
||||
|
||||
class _TeamJsonColumns(BaseModel):
|
||||
"""The two loosely typed columns on ``LiteLLM_TeamTable``, re-read with the shape the badge needs."""
|
||||
|
||||
metadata: Mapping[str, object] | None = None
|
||||
litellm_model_table: TeamModelAliasTable | None = None
|
||||
|
||||
|
||||
class TeamGrants(TypedDict, total=False):
|
||||
"""Keyword arguments for ``UserAPIKeyAuth``. Empty when the caller has no team, so the model's own defaults apply."""
|
||||
|
||||
team_alias: ReadOnly[str | None]
|
||||
team_tpm_limit: ReadOnly[int | None]
|
||||
team_rpm_limit: ReadOnly[int | None]
|
||||
team_max_budget: ReadOnly[float | None]
|
||||
team_soft_budget: ReadOnly[float | None]
|
||||
team_spend: ReadOnly[float | None]
|
||||
team_models: ReadOnly[Sequence[str]]
|
||||
team_blocked: ReadOnly[bool]
|
||||
team_metadata: ReadOnly[Mapping[str, object] | None]
|
||||
team_model_aliases: ReadOnly[Mapping[str, str] | None]
|
||||
team_object_permission_id: ReadOnly[str | None]
|
||||
team_object_permission: ReadOnly[LiteLLM_ObjectPermissionTable | None]
|
||||
team_member: ReadOnly[Member | None]
|
||||
team_member_spend: ReadOnly[float | None]
|
||||
team_member_tpm_limit: ReadOnly[int | None]
|
||||
team_member_rpm_limit: ReadOnly[int | None]
|
||||
|
||||
|
||||
def _json_columns(team_object: LiteLLM_TeamTable) -> _TeamJsonColumns:
|
||||
try:
|
||||
return _TeamJsonColumns.model_validate(team_object.model_dump(include=_JSON_COLUMNS))
|
||||
except ValidationError:
|
||||
return _TeamJsonColumns()
|
||||
|
||||
|
||||
def team_model_aliases(team_object: LiteLLM_TeamTable | None) -> Mapping[str, str] | None:
|
||||
if team_object is None:
|
||||
return None
|
||||
alias_table: Final = _json_columns(team_object).litellm_model_table
|
||||
return alias_table.model_aliases if alias_table is not None else None
|
||||
|
||||
|
||||
def team_grants(
|
||||
team_object: LiteLLM_TeamTable | None,
|
||||
team_membership: LiteLLM_TeamMembership | None,
|
||||
user_id: str | None,
|
||||
) -> TeamGrants:
|
||||
if team_object is None:
|
||||
return TeamGrants()
|
||||
json_columns: Final = _json_columns(team_object)
|
||||
return TeamGrants(
|
||||
team_alias=team_object.team_alias,
|
||||
team_tpm_limit=team_object.tpm_limit,
|
||||
team_rpm_limit=team_object.rpm_limit,
|
||||
team_max_budget=team_object.max_budget,
|
||||
team_soft_budget=team_object.soft_budget,
|
||||
team_spend=team_object.spend,
|
||||
team_models=tuple(team_object.models),
|
||||
team_blocked=team_object.blocked,
|
||||
team_metadata=json_columns.metadata,
|
||||
team_model_aliases=(
|
||||
json_columns.litellm_model_table.model_aliases if json_columns.litellm_model_table is not None else None
|
||||
),
|
||||
team_object_permission_id=team_object.object_permission_id,
|
||||
team_object_permission=team_object.object_permission,
|
||||
team_member=next(
|
||||
(m for m in team_object.members_with_roles if user_id is not None and m.user_id == user_id),
|
||||
None,
|
||||
),
|
||||
team_member_spend=team_membership.spend if team_membership is not None else None,
|
||||
team_member_tpm_limit=(
|
||||
team_membership.safe_get_team_member_tpm_limit() if team_membership is not None else None
|
||||
),
|
||||
team_member_rpm_limit=(
|
||||
team_membership.safe_get_team_member_rpm_limit() if team_membership is not None else None
|
||||
),
|
||||
)
|
||||
|
|
@ -82,6 +82,7 @@ from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
|
|||
from litellm.proxy.auth.resolvers import CredentialRef, Principal
|
||||
from litellm.proxy.auth.resolvers.store import IdentityStore
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.auth.team_grants import team_grants
|
||||
from litellm.proxy.auth.trusted_proxy_utils import get_trusted_proxy_cidrs
|
||||
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
|
|
@ -1476,24 +1477,16 @@ async def _user_api_key_auth_builder(
|
|||
user_id=user_id,
|
||||
user_email=user_email,
|
||||
team_id=team_id,
|
||||
team_alias=(team_object.team_alias if team_object is not None else None),
|
||||
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
|
||||
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
|
||||
team_models=(team_object.models if team_object is not None else []),
|
||||
team_metadata=(team_object.metadata if team_object is not None else None),
|
||||
org_id=org_id,
|
||||
end_user_id=end_user_id,
|
||||
parent_otel_span=parent_otel_span,
|
||||
jwt_claims=jwt_claims,
|
||||
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
|
||||
)
|
||||
|
||||
valid_token = UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
team_id=team_id,
|
||||
team_alias=(team_object.team_alias if team_object is not None else None),
|
||||
team_tpm_limit=(team_object.tpm_limit if team_object is not None else None),
|
||||
team_rpm_limit=(team_object.rpm_limit if team_object is not None else None),
|
||||
team_models=(team_object.models if team_object is not None else []),
|
||||
user_role=(
|
||||
LitellmUserRoles(user_object.user_role)
|
||||
if user_object is not None and user_object.user_role is not None
|
||||
|
|
@ -1507,17 +1500,8 @@ async def _user_api_key_auth_builder(
|
|||
user_tpm_limit=(user_object.tpm_limit if user_object is not None else None),
|
||||
user_rpm_limit=(user_object.rpm_limit if user_object is not None else None),
|
||||
user_model_max_budget=(user_object.model_max_budget if user_object is not None else None),
|
||||
team_member_rpm_limit=(
|
||||
team_membership.safe_get_team_member_rpm_limit() if team_membership is not None else None
|
||||
),
|
||||
team_member_tpm_limit=(
|
||||
team_membership.safe_get_team_member_tpm_limit() if team_membership is not None else None
|
||||
),
|
||||
team_metadata=(team_object.metadata if team_object is not None else None),
|
||||
jwt_claims=jwt_claims,
|
||||
)
|
||||
valid_token.team_object_permission = (
|
||||
team_object.object_permission if team_object is not None else None
|
||||
**team_grants(team_object=team_object, team_membership=team_membership, user_id=user_id),
|
||||
)
|
||||
|
||||
# AUTO_REGISTER deferred from _resolve_jwt_to_virtual_key.
|
||||
|
|
|
|||
|
|
@ -34,12 +34,20 @@ writer's connection params (pool size, timeouts, pgbouncer mode) for the
|
|||
ones the reader URL does not pin itself.
|
||||
"""
|
||||
|
||||
import _ssl
|
||||
import hashlib
|
||||
import os
|
||||
import socket
|
||||
import ssl
|
||||
import struct
|
||||
import sys
|
||||
import tempfile
|
||||
import urllib.parse
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Annotated, Final, cast
|
||||
from typing import Annotated, Final, Protocol, TypeAlias, cast
|
||||
|
||||
from pydantic import AliasChoices, BeforeValidator, Field
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
|
@ -126,21 +134,100 @@ def add_missing_query_params(url: str, params: Mapping[str, str | int | float])
|
|||
|
||||
|
||||
LIBPQ_VERIFY_SSLMODES: Final[frozenset[str]] = frozenset({"verify-ca", "verify-full"})
|
||||
PEM_CERT_HEADER: Final = b"-----BEGIN CERTIFICATE-----"
|
||||
PG_SSL_REQUEST: Final = struct.pack("!ii", 8, 80877103)
|
||||
TLS_PROBE_TIMEOUT_SECONDS: Final = 10.0
|
||||
|
||||
RootCertResolver: TypeAlias = Callable[[str, str, int], str] # mutable-ok: Callable parameter syntax
|
||||
|
||||
|
||||
def translate_libpq_ssl_params(url: str) -> str:
|
||||
class _VerifiedChainSource(Protocol):
|
||||
def get_verified_chain(self) -> Sequence[_ssl.Certificate] | None: ...
|
||||
|
||||
|
||||
def _verified_chain_der(tls: ssl.SSLSocket) -> tuple[bytes, ...]:
|
||||
if sys.version_info >= (3, 13):
|
||||
return tuple(tls.get_verified_chain())
|
||||
legacy: Final = cast( # cast-ok: the stub omits _sslobj, the C object has get_verified_chain since 3.10
|
||||
"_VerifiedChainSource | None",
|
||||
tls._sslobj, # pyright: ignore[reportAttributeAccessIssue, reportUnknownMemberType] # public API only from 3.13
|
||||
)
|
||||
chain: Final = () if legacy is None else legacy.get_verified_chain() or ()
|
||||
return tuple(cert.public_bytes(_ssl.ENCODING_DER) for cert in chain)
|
||||
|
||||
|
||||
def _server_trust_anchor(cafile: str, host: str, port: int) -> bytes | None:
|
||||
try:
|
||||
context: Final = ssl.create_default_context(cafile=cafile)
|
||||
with socket.create_connection((host, port), timeout=TLS_PROBE_TIMEOUT_SECONDS) as raw:
|
||||
raw.sendall(PG_SSL_REQUEST)
|
||||
if raw.recv(1) != b"S":
|
||||
return None
|
||||
with context.wrap_socket(raw, server_hostname=host) as tls:
|
||||
chain: Final = _verified_chain_der(tls)
|
||||
except (OSError, ValueError):
|
||||
return None
|
||||
return chain[-1] if chain else None
|
||||
|
||||
|
||||
def pin_bundle_root(cert_path: str, host: str, port: int) -> str:
|
||||
"""Reduce a multi-root CA bundle to the one root that verifies ``host``.
|
||||
|
||||
Prisma's ``sslcert`` loads a single PEM certificate (native-tls
|
||||
``Certificate::from_pem``), so pointing it at a bundle such as the AWS RDS
|
||||
global bundle trusts only the first of its 108 regional roots and the
|
||||
handshake fails with "unable to get local issuer certificate" for every
|
||||
other region. A single-certificate file is returned as is. For a bundle,
|
||||
one verifying handshake (chain and hostname, whole bundle as trust store)
|
||||
identifies the trust anchor the server actually chains to, which is
|
||||
written to a single-certificate file for Prisma. If the probe fails the
|
||||
bundle path is returned unchanged, so Prisma fails closed exactly as
|
||||
before rather than trusting anything the bundle would not.
|
||||
"""
|
||||
try:
|
||||
if Path(cert_path).read_bytes().count(PEM_CERT_HEADER) < 2:
|
||||
return cert_path
|
||||
except OSError:
|
||||
return cert_path
|
||||
root: Final = _server_trust_anchor(cert_path, host, port)
|
||||
if root is None:
|
||||
return cert_path
|
||||
pinned: Final = Path(tempfile.gettempdir()) / f"litellm-sslcert-{hashlib.sha256(root).hexdigest()[:16]}.pem"
|
||||
return str(pinned) if _replace_file(pinned, ssl.DER_cert_to_PEM_cert(root)) else cert_path
|
||||
|
||||
|
||||
def _replace_file(target: Path, content: str) -> bool:
|
||||
"""Write ``content`` to a private temp file and rename it over ``target``, so
|
||||
readers never see a partial file and a symlink planted at ``target`` is
|
||||
replaced rather than followed."""
|
||||
try:
|
||||
fd, staged = tempfile.mkstemp(dir=target.parent, prefix=f"{target.name}.")
|
||||
except OSError:
|
||||
return False
|
||||
try:
|
||||
with os.fdopen(fd, "w") as handle:
|
||||
handle.write(content)
|
||||
os.replace(staged, target)
|
||||
except OSError:
|
||||
Path(staged).unlink(missing_ok=True)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def translate_libpq_ssl_params(url: str, resolve_root_cert: RootCertResolver = pin_bundle_root) -> str:
|
||||
"""Rewrite libpq's certificate-verification params into Prisma's dialect.
|
||||
|
||||
Prisma's engine only knows ``sslmode=disable|prefer|require``, ``sslcert``
|
||||
(the CA bundle) and ``sslaccept=strict``. It silently discards
|
||||
(a single CA certificate) and ``sslaccept=strict``. It silently discards
|
||||
``sslrootcert`` and downgrades ``sslmode=verify-ca`` / ``verify-full`` to
|
||||
``prefer``, so a URL copied from libpq / RDS docs connects over TLS with no
|
||||
certificate check at all. ``verify-ca`` and ``verify-full`` both become
|
||||
``require`` (Prisma has no CA-only mode), ``sslrootcert`` becomes
|
||||
``sslcert``, and either one turns on ``sslaccept=strict`` (chain and
|
||||
hostname), matching libpq where a root cert makes ``require`` verify.
|
||||
Prisma params the operator pinned themselves win; anything else is left
|
||||
untouched.
|
||||
``sslcert`` (run through ``resolve_root_cert``, which pins a multi-root
|
||||
bundle down to the server's root), and either one turns on
|
||||
``sslaccept=strict`` (chain and hostname), matching libpq where a root
|
||||
cert makes ``require`` verify. Prisma params the operator pinned
|
||||
themselves win; anything else is left untouched.
|
||||
"""
|
||||
parsed: Final = urllib.parse.urlsplit(url)
|
||||
pairs: Final = tuple(urllib.parse.parse_qsl(parsed.query, keep_blank_values=True))
|
||||
|
|
@ -154,7 +241,9 @@ def translate_libpq_ssl_params(url: str) -> str:
|
|||
if key != "sslrootcert"
|
||||
)
|
||||
root_cert: Final = tuple(
|
||||
("sslcert", value) for key, value in pairs if key == "sslrootcert" and "sslcert" not in keys
|
||||
("sslcert", resolve_root_cert(value, parsed.hostname or "", parsed.port or int(DEFAULT_POSTGRES_PORT)))
|
||||
for key, value in pairs
|
||||
if key == "sslrootcert" and "sslcert" not in keys
|
||||
)
|
||||
strict: Final = () if "sslaccept" in keys else (("sslaccept", "strict"),)
|
||||
query: Final = urllib.parse.urlencode(translated + root_cert + strict)
|
||||
|
|
|
|||
|
|
@ -10,13 +10,28 @@ strings rather than passing the raw path through. Nothing a caller sends can
|
|||
add a key, so the fold and the table it commits to are bounded by (days x
|
||||
routes) however much traffic arrives, and the response path carries no
|
||||
unbounded queue that would block once full.
|
||||
|
||||
A flush commits its whole snapshot as one multi-row ``INSERT ... ON CONFLICT DO
|
||||
UPDATE`` rather than one upsert per key, so a worker costs the primary one
|
||||
statement per interval however many routes it served. With
|
||||
``use_redis_transaction_buffer`` on, workers instead push their snapshot to a
|
||||
Redis list and one lock-holding pod folds every entry and writes the table, so
|
||||
the deployment as a whole costs the primary one statement per interval.
|
||||
"""
|
||||
|
||||
from dataclasses import asdict
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Iterable
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import RedisCache
|
||||
from litellm.constants import MAX_REDIS_BUFFER_DEQUEUE_COUNT, REDIS_GATEWAY_REQUESTS_BUFFER_KEY
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
|
||||
from litellm.proxy.middleware.billable_request_metrics_middleware import BillableCategory
|
||||
from litellm.types.proxy.gateway_requests import (
|
||||
GatewayRequestCounts,
|
||||
|
|
@ -28,6 +43,15 @@ if TYPE_CHECKING:
|
|||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
_EMPTY: Final = GatewayRequestCounts(successful_requests=0, failed_requests=0)
|
||||
_TABLE: Final = '"LiteLLM_DailyGatewayRequests"'
|
||||
_COLUMNS_PER_ROW: Final = 5
|
||||
_UTC_NOW: Final = "(NOW() AT TIME ZONE 'UTC')"
|
||||
GATEWAY_REQUESTS_JOB_NAME: Final = "update_gateway_requests_job"
|
||||
|
||||
_BufferedRows: TypeAlias = tuple[tuple[str, str, str, int, int], ...]
|
||||
_BUFFERED_ROWS: Final = TypeAdapter(_BufferedRows)
|
||||
_BUFFERED_ENTRIES: Final = TypeAdapter(tuple[str | bytes, ...])
|
||||
_NO_COUNTS: Final[GatewayRequestSnapshot] = MappingProxyType({})
|
||||
|
||||
|
||||
def _utc_date() -> str:
|
||||
|
|
@ -59,20 +83,54 @@ class GatewayRequestAccumulator:
|
|||
route) however long the database is unreachable.
|
||||
|
||||
This buys at-least-once, not exactly-once, and the cost is worth stating.
|
||||
The batch commits inside its context manager's ``__aexit__``, so a failure
|
||||
raised after the transaction committed (a connection dropped while reading
|
||||
the acknowledgement) restores counts that are already persisted, and the
|
||||
next flush increments them a second time. Exactly-once would need a dedup
|
||||
key the upserts could ignore on replay. For a traffic-volume metric a rare
|
||||
The statement commits on the server before its acknowledgement is read, so
|
||||
a failure raised after the commit (a connection dropped while reading the
|
||||
acknowledgement) restores counts that are already persisted, and the next
|
||||
flush increments them a second time. Exactly-once would need a dedup key
|
||||
the upsert could ignore on replay. For a traffic-volume metric a rare
|
||||
overcount on a dropped acknowledgement beats losing a whole interval to
|
||||
every database blip, so the trade is deliberate.
|
||||
"""
|
||||
for key, counts in snapshot.items():
|
||||
existing = self._counts.get(key, _EMPTY)
|
||||
self._counts[key] = GatewayRequestCounts(
|
||||
successful_requests=existing.successful_requests + counts.successful_requests,
|
||||
failed_requests=existing.failed_requests + counts.failed_requests,
|
||||
)
|
||||
self._counts = dict(fold_counts(chain(self._counts.items(), snapshot.items()))) # mutable-ok: fold replaced
|
||||
|
||||
|
||||
def fold_counts(items: Iterable[tuple[GatewayRequestKey, GatewayRequestCounts]]) -> GatewayRequestSnapshot:
|
||||
"""Sum counts key-wise; the result stays bounded by (date x category x route)."""
|
||||
folded: Final[dict[GatewayRequestKey, GatewayRequestCounts]] = {} # mutable-ok: local fold returned once
|
||||
for key, counts in items:
|
||||
existing = folded.get(key, _EMPTY)
|
||||
folded[key] = GatewayRequestCounts(
|
||||
successful_requests=existing.successful_requests + counts.successful_requests,
|
||||
failed_requests=existing.failed_requests + counts.failed_requests,
|
||||
)
|
||||
return folded
|
||||
|
||||
|
||||
def build_gateway_requests_upsert(snapshot: GatewayRequestSnapshot) -> tuple[str, tuple[str | int, ...]]:
|
||||
"""
|
||||
One ``INSERT ... ON CONFLICT DO UPDATE`` that increments every (date, category,
|
||||
route) in the snapshot. Rows are ordered by the conflict key so concurrent
|
||||
writers lock rows in the same order and cannot deadlock.
|
||||
"""
|
||||
ordered: Final = sorted(snapshot.items(), key=lambda item: (item[0].date, item[0].category, item[0].route))
|
||||
rows: Final = ", ".join(
|
||||
f"(${base + 1}::text, ${base + 2}::text, ${base + 3}::text, ${base + 4}::bigint, ${base + 5}::bigint, {_UTC_NOW})"
|
||||
for base in range(0, len(ordered) * _COLUMNS_PER_ROW, _COLUMNS_PER_ROW)
|
||||
)
|
||||
sql: Final = (
|
||||
f'INSERT INTO {_TABLE} ("date", "category", "route", "successful_requests", "failed_requests", "updated_at")\n'
|
||||
f"VALUES {rows}\n"
|
||||
'ON CONFLICT ("date", "category", "route") DO UPDATE SET\n'
|
||||
f' "successful_requests" = {_TABLE}."successful_requests" + EXCLUDED."successful_requests",\n'
|
||||
f' "failed_requests" = {_TABLE}."failed_requests" + EXCLUDED."failed_requests",\n'
|
||||
f' "updated_at" = {_UTC_NOW}'
|
||||
)
|
||||
params: Final[tuple[str | int, ...]] = tuple(
|
||||
value
|
||||
for key, counts in ordered
|
||||
for value in (key.date, key.category, key.route, counts.successful_requests, counts.failed_requests)
|
||||
)
|
||||
return sql, params
|
||||
|
||||
|
||||
async def commit_gateway_requests_to_db(
|
||||
|
|
@ -80,50 +138,130 @@ async def commit_gateway_requests_to_db(
|
|||
prisma_client: "PrismaClient",
|
||||
snapshot: GatewayRequestSnapshot,
|
||||
) -> None:
|
||||
"""Upsert one incrementing row per (date, category, route)."""
|
||||
"""Increment every (date, category, route) in the snapshot with a single statement."""
|
||||
if not snapshot:
|
||||
return
|
||||
|
||||
ordered: Final = sorted(snapshot.items(), key=lambda item: (item[0].date, item[0].category, item[0].route))
|
||||
sql, params = build_gateway_requests_upsert(snapshot)
|
||||
await prisma_client.db.execute_raw(sql, *params) # pyright: ignore[reportAny] # untyped prisma client
|
||||
|
||||
# pyright: ignore[reportAny] on both lines -- prisma's generated client is untyped,
|
||||
# so .db and every table action off it resolve to Any at this boundary. The dict
|
||||
# literals below are the shape prisma's generated inputs require.
|
||||
async with prisma_client.db.batch_() as batcher: # pyright: ignore[reportAny] # untyped prisma client
|
||||
for key, counts in ordered:
|
||||
columns = asdict(key)
|
||||
batcher.litellm_dailygatewayrequests.upsert( # pyright: ignore[reportAny] # untyped prisma client
|
||||
where={"date_category_route": columns}, # mutable-ok: prisma input is dict-shaped
|
||||
data={ # mutable-ok: prisma input is dict-shaped
|
||||
"create": { # mutable-ok: prisma input is dict-shaped
|
||||
**columns,
|
||||
"successful_requests": counts.successful_requests,
|
||||
"failed_requests": counts.failed_requests,
|
||||
},
|
||||
"update": { # mutable-ok: prisma input is dict-shaped
|
||||
"successful_requests": {"increment": counts.successful_requests}, # mutable-ok: as above
|
||||
"failed_requests": {"increment": counts.failed_requests}, # mutable-ok: as above
|
||||
},
|
||||
},
|
||||
verbose_proxy_logger.debug(
|
||||
"Gateway request tracking - committed %d aggregated rows in one statement", len(snapshot)
|
||||
)
|
||||
|
||||
|
||||
class GatewayRequestRedisBuffer:
|
||||
"""
|
||||
Folds every worker's snapshot through one Redis list so a single pod per
|
||||
interval writes the table, mirroring the spend writer's transaction buffer.
|
||||
|
||||
Each entry is one worker's snapshot as JSON rows; the lock holder pops them,
|
||||
sums them, and commits one statement. A commit failure pushes the summed
|
||||
rows back so the next holder retries, keeping the at-least-once guarantee.
|
||||
If that push fails too, the rows go back to the holder's own accumulator so
|
||||
they ride along with its next flush instead of vanishing with the pop.
|
||||
"""
|
||||
|
||||
def __init__(self, *, redis_cache: RedisCache, pod_lock_manager: PodLockManager) -> None:
|
||||
self._redis_cache: Final = redis_cache
|
||||
self._pod_lock_manager: Final = pod_lock_manager
|
||||
|
||||
async def push(self, snapshot: GatewayRequestSnapshot) -> None:
|
||||
if not snapshot:
|
||||
return
|
||||
rows: Final[_BufferedRows] = tuple(
|
||||
(key.date, key.category, key.route, counts.successful_requests, counts.failed_requests)
|
||||
for key, counts in snapshot.items()
|
||||
)
|
||||
await self._redis_cache.async_rpush(key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, values=(json.dumps(rows),))
|
||||
|
||||
async def _pop_batch(self) -> tuple[str | bytes, ...]:
|
||||
popped: Final[object] = await self._redis_cache.async_lpop( # pyright: ignore[reportAny] # redis returns Any
|
||||
key=REDIS_GATEWAY_REQUESTS_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT
|
||||
)
|
||||
if not popped:
|
||||
return ()
|
||||
return _BUFFERED_ENTRIES.validate_python(popped if isinstance(popped, list) else (popped,))
|
||||
|
||||
async def _pop_all(self) -> AsyncIterator[str | bytes]:
|
||||
while True:
|
||||
batch = await self._pop_batch()
|
||||
for entry in batch:
|
||||
yield entry
|
||||
if len(batch) < MAX_REDIS_BUFFER_DEQUEUE_COUNT:
|
||||
return
|
||||
|
||||
async def pop(self) -> GatewayRequestSnapshot:
|
||||
entries: Final = tuple([entry async for entry in self._pop_all()])
|
||||
return fold_counts(
|
||||
(
|
||||
GatewayRequestKey(date=date, category=category, route=route),
|
||||
GatewayRequestCounts(successful_requests=succeeded, failed_requests=failed),
|
||||
)
|
||||
for entry in entries
|
||||
for date, category, route, succeeded, failed in _BUFFERED_ROWS.validate_json(entry)
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug("Gateway request tracking - committed %d aggregated rows", len(ordered))
|
||||
async def commit_if_leader(self, prisma_client: "PrismaClient") -> GatewayRequestSnapshot:
|
||||
"""
|
||||
Drain the list and write it as one statement, but only on the pod holding the job lock.
|
||||
|
||||
The lock is a lease, never released: the holder re-enters it on every flush and
|
||||
keeps committing alone until the TTL lapses, so the primary sees one statement
|
||||
per flush interval deployment-wide instead of one per worker.
|
||||
|
||||
Returns the popped rows that could be neither committed nor re-queued, for the
|
||||
caller to keep in memory. Empty on success.
|
||||
"""
|
||||
if not await self._pod_lock_manager.acquire_lock(cronjob_id=GATEWAY_REQUESTS_JOB_NAME):
|
||||
return _NO_COUNTS
|
||||
buffered: Final = await self.pop()
|
||||
try:
|
||||
await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=buffered)
|
||||
except Exception: # noqa: BLE001 -- a failed commit must not stop the scheduler
|
||||
verbose_proxy_logger.warning(
|
||||
"Gateway request tracking - failed to commit %d buffered rows, re-queuing to Redis for the next flush",
|
||||
len(buffered),
|
||||
exc_info=True,
|
||||
)
|
||||
return await self._requeue(buffered)
|
||||
return _NO_COUNTS
|
||||
|
||||
async def _requeue(self, snapshot: GatewayRequestSnapshot) -> GatewayRequestSnapshot:
|
||||
try:
|
||||
await self.push(snapshot)
|
||||
except Exception: # noqa: BLE001 -- the rows go back to the caller's accumulator instead
|
||||
verbose_proxy_logger.warning(
|
||||
"Gateway request tracking - Redis re-queue failed, keeping %d rows in memory for the next flush",
|
||||
len(snapshot),
|
||||
exc_info=True,
|
||||
)
|
||||
return snapshot
|
||||
return _NO_COUNTS
|
||||
|
||||
|
||||
async def flush_gateway_requests(
|
||||
prisma_client: "PrismaClient",
|
||||
accumulator: GatewayRequestAccumulator,
|
||||
redis_buffer: GatewayRequestRedisBuffer | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Scheduler entrypoint. Never raises: a metering failure must not kill the job.
|
||||
|
||||
With ``redis_buffer`` the snapshot goes to Redis and only the lease holder
|
||||
writes to Postgres. Shutdown passes no buffer so a departing worker writes its
|
||||
own counts directly instead of parking them behind a lease it may not hold.
|
||||
|
||||
``CancelledError`` is deliberately not caught, so a flush cancelled during
|
||||
shutdown drops its snapshot rather than restoring counts onto an accumulator
|
||||
the process is about to discard.
|
||||
"""
|
||||
snapshot: Final = accumulator.drain()
|
||||
try:
|
||||
await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=snapshot)
|
||||
if redis_buffer is None:
|
||||
await commit_gateway_requests_to_db(prisma_client=prisma_client, snapshot=snapshot)
|
||||
else:
|
||||
await redis_buffer.push(snapshot)
|
||||
except Exception: # noqa: BLE001 -- a failed flush must not stop the scheduler
|
||||
accumulator.restore(snapshot)
|
||||
verbose_proxy_logger.warning(
|
||||
|
|
@ -131,3 +269,13 @@ async def flush_gateway_requests(
|
|||
len(snapshot),
|
||||
exc_info=True,
|
||||
)
|
||||
return
|
||||
if redis_buffer is None:
|
||||
return
|
||||
try:
|
||||
accumulator.restore(await redis_buffer.commit_if_leader(prisma_client))
|
||||
except Exception: # noqa: BLE001 -- entries still in Redis are drained by the next flush
|
||||
verbose_proxy_logger.warning(
|
||||
"Gateway request tracking - leader drain failed, buffered rows stay in Redis for the next flush",
|
||||
exc_info=True,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from fastapi import APIRouter, Depends, HTTPException
|
|||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import current_billing_time, pinned_billing_time
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -27,7 +28,15 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.utils import CostPerToken, LlmProvidersSet, ModelInfo
|
||||
from litellm.types.utils import (
|
||||
CostBreakdown,
|
||||
CostPerToken,
|
||||
LlmProvidersSet,
|
||||
ModelInfo,
|
||||
ModelResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
|
@ -46,13 +55,15 @@ def _configured_price(key: str, sources: tuple[Mapping[str, object], ...]) -> fl
|
|||
|
||||
|
||||
def _extract_custom_pricing(
|
||||
litellm_params: Mapping[str, object], model_info: Mapping[str, object]
|
||||
litellm_params: Mapping[str, object], model_info: Mapping[str, object], builtin: ModelInfo | None
|
||||
) -> CostPerToken | None:
|
||||
"""
|
||||
Pull per-token pricing configured on a deployment so on-prem / self-hosted
|
||||
models (absent from the public cost map) still estimate a real cost.
|
||||
Pricing may live on ``litellm_params`` or ``model_info``; ``litellm_params``
|
||||
wins, matching the router's cost-map registration precedence.
|
||||
wins, matching the router's cost-map registration precedence. Cache rates the
|
||||
deployment leaves unset come from the backend model's built-in entry, then its
|
||||
own input rate, again matching what the router registers for live billing.
|
||||
"""
|
||||
sources: Final = (litellm_params, model_info)
|
||||
input_price: Final = _configured_price("input_cost_per_token", sources)
|
||||
|
|
@ -61,15 +72,21 @@ def _extract_custom_pricing(
|
|||
if input_price is None and output_price is None:
|
||||
return None
|
||||
|
||||
input_rate: Final = input_price or 0.0
|
||||
cache_sources: Final = sources if builtin is None else (*sources, builtin)
|
||||
cache_read_price: Final = _configured_price("cache_read_input_token_cost", cache_sources)
|
||||
cache_creation_price: Final = _configured_price("cache_creation_input_token_cost", cache_sources)
|
||||
return CostPerToken(
|
||||
input_cost_per_token=input_price or 0.0,
|
||||
input_cost_per_token=input_rate,
|
||||
output_cost_per_token=output_price or 0.0,
|
||||
cache_read_input_token_cost=input_rate if cache_read_price is None else cache_read_price,
|
||||
cache_creation_input_token_cost=input_rate if cache_creation_price is None else cache_creation_price,
|
||||
)
|
||||
|
||||
|
||||
def _lookup_model_info(model: str) -> ModelInfo | None:
|
||||
def _lookup_model_info(model: str, custom_llm_provider: str | None = None) -> ModelInfo | None:
|
||||
try:
|
||||
return litellm.get_model_info(model=model)
|
||||
return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
|
@ -98,17 +115,14 @@ def _resolve_model_for_cost_lookup(model: str) -> ResolvedCostModel:
|
|||
model_info: Final = first_deployment.get("model_info", {})
|
||||
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
|
||||
provider: Final = str(custom_llm_provider) if custom_llm_provider is not None else None
|
||||
custom_cost_per_token: Final = _extract_custom_pricing(litellm_params, model_info)
|
||||
|
||||
# Check base_model first (needed for Azure custom deployment names)
|
||||
# base_model wins (needed for Azure custom deployment names)
|
||||
base_model: Final = model_info.get("base_model") or litellm_params.get("base_model")
|
||||
if base_model:
|
||||
verbose_proxy_logger.debug("Resolved model '%s' to base_model '%s' from router", model, base_model)
|
||||
return ResolvedCostModel(str(base_model), provider, custom_cost_per_token)
|
||||
|
||||
resolved_model: Final = litellm_params.get("model")
|
||||
resolved_model: Final = base_model or litellm_params.get("model")
|
||||
if resolved_model:
|
||||
verbose_proxy_logger.debug("Resolved model '%s' to '%s' from router", model, resolved_model)
|
||||
custom_cost_per_token: Final = _extract_custom_pricing(
|
||||
litellm_params, model_info, _lookup_model_info(str(resolved_model), provider)
|
||||
)
|
||||
return ResolvedCostModel(str(resolved_model), provider, custom_cost_per_token)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Could not resolve model '%s' from router: %s", model, e)
|
||||
|
|
@ -117,19 +131,59 @@ def _resolve_model_for_cost_lookup(model: str) -> ResolvedCostModel:
|
|||
return ResolvedCostModel(model, None, None)
|
||||
|
||||
|
||||
def _calculate_period_costs(num_requests, cost_per_request, input_cost, output_cost, margin_cost):
|
||||
"""
|
||||
Calculate costs for a given number of requests.
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CostLines:
|
||||
"""Cost of one request split the way the spend logs split it: the cache lines are
|
||||
shares of input_cost and the reasoning line is a share of output_cost."""
|
||||
|
||||
Returns tuple of (total_cost, input_cost, output_cost, margin_cost) or all None if num_requests is None/0.
|
||||
"""
|
||||
if not num_requests:
|
||||
return None, None, None, None
|
||||
return (
|
||||
cost_per_request * num_requests,
|
||||
input_cost * num_requests,
|
||||
output_cost * num_requests,
|
||||
margin_cost * num_requests,
|
||||
total_cost: float
|
||||
input_cost: float
|
||||
output_cost: float
|
||||
margin_cost: float
|
||||
cache_read_cost: float
|
||||
cache_creation_cost: float
|
||||
reasoning_cost: float
|
||||
|
||||
def times(self, num_requests: int | None) -> "CostLines | None":
|
||||
if not num_requests:
|
||||
return None
|
||||
return CostLines(
|
||||
total_cost=self.total_cost * num_requests,
|
||||
input_cost=self.input_cost * num_requests,
|
||||
output_cost=self.output_cost * num_requests,
|
||||
margin_cost=self.margin_cost * num_requests,
|
||||
cache_read_cost=self.cache_read_cost * num_requests,
|
||||
cache_creation_cost=self.cache_creation_cost * num_requests,
|
||||
reasoning_cost=self.reasoning_cost * num_requests,
|
||||
)
|
||||
|
||||
|
||||
def _cost_lines(cost_per_request: float, cost_breakdown: CostBreakdown | None) -> CostLines:
|
||||
breakdown: Final = cost_breakdown if cost_breakdown is not None else CostBreakdown()
|
||||
return CostLines(
|
||||
total_cost=cost_per_request,
|
||||
input_cost=breakdown.get("input_cost", 0.0),
|
||||
output_cost=breakdown.get("output_cost", 0.0),
|
||||
margin_cost=breakdown.get("margin_total_amount", 0.0),
|
||||
cache_read_cost=breakdown.get("cache_read_cost", 0.0),
|
||||
cache_creation_cost=breakdown.get("cache_creation_cost", 0.0),
|
||||
reasoning_cost=breakdown.get("reasoning_cost", 0.0),
|
||||
)
|
||||
|
||||
|
||||
def _usage_for_estimate(request: CostEstimateRequest) -> Usage:
|
||||
cache_tokens: Final = request.cache_read_input_tokens + request.cache_creation_input_tokens
|
||||
return Usage(
|
||||
prompt_tokens=request.input_tokens,
|
||||
completion_tokens=request.output_tokens,
|
||||
total_tokens=request.input_tokens + request.output_tokens,
|
||||
reasoning_tokens=request.reasoning_tokens,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
cached_tokens=request.cache_read_input_tokens,
|
||||
cache_creation_tokens=request.cache_creation_input_tokens,
|
||||
)
|
||||
if cache_tokens
|
||||
else None,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -530,11 +584,14 @@ async def estimate_cost(
|
|||
- model: Model name (e.g., "gpt-4", "claude-3-opus")
|
||||
- input_tokens: Expected input tokens per request
|
||||
- output_tokens: Expected output tokens per request
|
||||
- cache_read_input_tokens: Cache-read tokens per request, counted within input_tokens (optional)
|
||||
- cache_creation_input_tokens: Cache-write tokens per request, counted within input_tokens (optional)
|
||||
- reasoning_tokens: Reasoning tokens per request, counted within output_tokens (optional)
|
||||
- num_requests_per_day: Number of requests per day (optional)
|
||||
- num_requests_per_month: Number of requests per month (optional)
|
||||
|
||||
Returns cost breakdown including:
|
||||
- Per-request costs (input, output, margin)
|
||||
- Per-request costs (input, output, margin, plus the cache-read, cache-write and reasoning shares)
|
||||
- Daily costs (if num_requests_per_day provided)
|
||||
- Monthly costs (if num_requests_per_month provided)
|
||||
|
||||
|
|
@ -543,14 +600,15 @@ async def estimate_cost(
|
|||
{
|
||||
"model": "gpt-4",
|
||||
"input_tokens": 1000,
|
||||
"cache_read_input_tokens": 800,
|
||||
"output_tokens": 500,
|
||||
"reasoning_tokens": 200,
|
||||
"num_requests_per_day": 100,
|
||||
"num_requests_per_month": 3000
|
||||
}
|
||||
```
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
# Resolve model name (handles router aliases like 'e-model-router' -> 'azure_ai/gpt-4')
|
||||
resolved: Final = _resolve_model_for_cost_lookup(request.model)
|
||||
|
|
@ -559,15 +617,8 @@ async def estimate_cost(
|
|||
|
||||
verbose_proxy_logger.debug("Cost estimate: request.model='%s' resolved to '%s'", request.model, resolved_model)
|
||||
|
||||
# Create a mock response with usage for completion_cost
|
||||
mock_response: Final = ModelResponse(
|
||||
model=resolved_model,
|
||||
usage=Usage(
|
||||
prompt_tokens=request.input_tokens,
|
||||
completion_tokens=request.output_tokens,
|
||||
total_tokens=request.input_tokens + request.output_tokens,
|
||||
),
|
||||
)
|
||||
usage: Final = _usage_for_estimate(request)
|
||||
mock_response: Final = ModelResponse(model=resolved_model, usage=usage)
|
||||
|
||||
# Create a logging object to capture cost breakdown
|
||||
litellm_logging_obj: Final = LiteLLMLoggingObj(
|
||||
|
|
@ -580,92 +631,73 @@ async def estimate_cost(
|
|||
function_id="cost-estimate",
|
||||
)
|
||||
|
||||
# Use completion_cost which handles all the logic including margins/discounts
|
||||
try:
|
||||
cost_per_request: Final = completion_cost(
|
||||
completion_response=mock_response,
|
||||
model=resolved_model,
|
||||
custom_llm_provider=resolved_provider,
|
||||
custom_cost_per_token=resolved.custom_cost_per_token,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={
|
||||
"error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {e}"
|
||||
},
|
||||
)
|
||||
# Pinning one moment keeps an off-peak window that opens mid-quote from pricing the totals on
|
||||
# one side of it and the reported rates on the other.
|
||||
with pinned_billing_time(current_billing_time()):
|
||||
# Use completion_cost which handles all the logic including margins/discounts
|
||||
try:
|
||||
cost_per_request: Final = completion_cost(
|
||||
completion_response=mock_response,
|
||||
model=resolved_model,
|
||||
custom_llm_provider=resolved_provider,
|
||||
custom_cost_per_token=resolved.custom_cost_per_token,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={
|
||||
"error": f"Could not calculate cost for model '{request.model}' (resolved to '{resolved_model}'): {e}"
|
||||
},
|
||||
)
|
||||
|
||||
# Get cost breakdown from the logging object
|
||||
cost_breakdown: Final = litellm_logging_obj.cost_breakdown
|
||||
# The rates come back from the pricing call itself rather than a second lookup, so they are the
|
||||
# ones the cost lines above billed at even when completion_cost infers a provider this endpoint
|
||||
# never resolved (an unrouted "xai/grok-4" prices on xai's inclusive tier thresholds; a lookup
|
||||
# here without that provider would report the sub-200k rate for a line billed above it).
|
||||
rates: Final = litellm_logging_obj.billed_token_rates
|
||||
per_request: Final = _cost_lines(cost_per_request, litellm_logging_obj.cost_breakdown)
|
||||
daily: Final = per_request.times(request.num_requests_per_day)
|
||||
monthly: Final = per_request.times(request.num_requests_per_month)
|
||||
|
||||
input_cost: Final = cost_breakdown.get("input_cost", 0.0) if cost_breakdown else 0.0
|
||||
output_cost: Final = cost_breakdown.get("output_cost", 0.0) if cost_breakdown else 0.0
|
||||
margin_cost: Final = cost_breakdown.get("margin_total_amount", 0.0) if cost_breakdown else 0.0
|
||||
|
||||
model_info: Final = _lookup_model_info(resolved_model)
|
||||
mapped_input_price: Final = model_info.get("input_cost_per_token") if model_info is not None else None
|
||||
mapped_output_price: Final = model_info.get("output_cost_per_token") if model_info is not None else None
|
||||
model_info: Final = _lookup_model_info(resolved_model, resolved_provider)
|
||||
mapped_provider: Final = model_info.get("litellm_provider") if model_info is not None else None
|
||||
|
||||
input_cost_per_token: Final = (
|
||||
resolved.custom_cost_per_token["input_cost_per_token"]
|
||||
if resolved.custom_cost_per_token is not None
|
||||
else mapped_input_price
|
||||
)
|
||||
output_cost_per_token: Final = (
|
||||
resolved.custom_cost_per_token["output_cost_per_token"]
|
||||
if resolved.custom_cost_per_token is not None
|
||||
else mapped_output_price
|
||||
)
|
||||
custom_llm_provider: Final = mapped_provider if mapped_provider is not None else resolved_provider
|
||||
|
||||
# Calculate daily and monthly costs
|
||||
(
|
||||
daily_cost,
|
||||
daily_input_cost,
|
||||
daily_output_cost,
|
||||
daily_margin_cost,
|
||||
) = _calculate_period_costs(
|
||||
num_requests=request.num_requests_per_day,
|
||||
cost_per_request=cost_per_request,
|
||||
input_cost=input_cost,
|
||||
output_cost=output_cost,
|
||||
margin_cost=margin_cost,
|
||||
)
|
||||
(
|
||||
monthly_cost,
|
||||
monthly_input_cost,
|
||||
monthly_output_cost,
|
||||
monthly_margin_cost,
|
||||
) = _calculate_period_costs(
|
||||
num_requests=request.num_requests_per_month,
|
||||
cost_per_request=cost_per_request,
|
||||
input_cost=input_cost,
|
||||
output_cost=output_cost,
|
||||
margin_cost=margin_cost,
|
||||
)
|
||||
|
||||
return CostEstimateResponse(
|
||||
model=request.model,
|
||||
input_tokens=request.input_tokens,
|
||||
output_tokens=request.output_tokens,
|
||||
cache_read_input_tokens=request.cache_read_input_tokens,
|
||||
cache_creation_input_tokens=request.cache_creation_input_tokens,
|
||||
reasoning_tokens=request.reasoning_tokens,
|
||||
num_requests_per_day=request.num_requests_per_day,
|
||||
num_requests_per_month=request.num_requests_per_month,
|
||||
cost_per_request=cost_per_request,
|
||||
input_cost_per_request=input_cost,
|
||||
output_cost_per_request=output_cost,
|
||||
margin_cost_per_request=margin_cost,
|
||||
daily_cost=daily_cost,
|
||||
daily_input_cost=daily_input_cost,
|
||||
daily_output_cost=daily_output_cost,
|
||||
daily_margin_cost=daily_margin_cost,
|
||||
monthly_cost=monthly_cost,
|
||||
monthly_input_cost=monthly_input_cost,
|
||||
monthly_output_cost=monthly_output_cost,
|
||||
monthly_margin_cost=monthly_margin_cost,
|
||||
input_cost_per_token=input_cost_per_token,
|
||||
output_cost_per_token=output_cost_per_token,
|
||||
cost_per_request=per_request.total_cost,
|
||||
input_cost_per_request=per_request.input_cost,
|
||||
output_cost_per_request=per_request.output_cost,
|
||||
margin_cost_per_request=per_request.margin_cost,
|
||||
cache_read_cost_per_request=per_request.cache_read_cost,
|
||||
cache_creation_cost_per_request=per_request.cache_creation_cost,
|
||||
reasoning_cost_per_request=per_request.reasoning_cost,
|
||||
daily_cost=daily.total_cost if daily is not None else None,
|
||||
daily_input_cost=daily.input_cost if daily is not None else None,
|
||||
daily_output_cost=daily.output_cost if daily is not None else None,
|
||||
daily_margin_cost=daily.margin_cost if daily is not None else None,
|
||||
daily_cache_read_cost=daily.cache_read_cost if daily is not None else None,
|
||||
daily_cache_creation_cost=daily.cache_creation_cost if daily is not None else None,
|
||||
daily_reasoning_cost=daily.reasoning_cost if daily is not None else None,
|
||||
monthly_cost=monthly.total_cost if monthly is not None else None,
|
||||
monthly_input_cost=monthly.input_cost if monthly is not None else None,
|
||||
monthly_output_cost=monthly.output_cost if monthly is not None else None,
|
||||
monthly_margin_cost=monthly.margin_cost if monthly is not None else None,
|
||||
monthly_cache_read_cost=monthly.cache_read_cost if monthly is not None else None,
|
||||
monthly_cache_creation_cost=monthly.cache_creation_cost if monthly is not None else None,
|
||||
monthly_reasoning_cost=monthly.reasoning_cost if monthly is not None else None,
|
||||
input_cost_per_token=rates.input_cost_per_token if rates is not None else None,
|
||||
output_cost_per_token=rates.output_cost_per_token if rates is not None else None,
|
||||
cache_read_input_token_cost=rates.cache_read_input_token_cost if rates is not None else None,
|
||||
cache_creation_input_token_cost=rates.cache_creation_input_token_cost if rates is not None else None,
|
||||
output_cost_per_reasoning_token=rates.output_cost_per_reasoning_token if rates is not None else None,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -5601,7 +5601,7 @@ async def team_model_add(
|
|||
updated_team: Final = await _team_db(prisma_client).update(
|
||||
where={"team_id": data.team_id},
|
||||
data={"updated_at": datetime.now(timezone.utc)},
|
||||
include={"object_permission": True},
|
||||
include={"litellm_model_table": True, "object_permission": True},
|
||||
)
|
||||
if updated_team is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -5688,7 +5688,7 @@ async def team_model_delete(
|
|||
updated_team: Final = await _team_db(prisma_client).update(
|
||||
where={"team_id": data.team_id},
|
||||
data={"models": updated_models},
|
||||
include={"object_permission": True},
|
||||
include={"litellm_model_table": True, "object_permission": True},
|
||||
)
|
||||
if updated_team is None:
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -22,7 +22,6 @@ from html import escape
|
|||
from types import MappingProxyType
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Annotated,
|
||||
Any,
|
||||
Final,
|
||||
Literal,
|
||||
|
|
@ -42,7 +41,7 @@ if TYPE_CHECKING:
|
|||
import jwt
|
||||
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response, status
|
||||
from fastapi.responses import RedirectResponse
|
||||
from pydantic import BaseModel, BeforeValidator, ConfigDict, TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -95,6 +94,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
)
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.auth.team_grants import TeamModelAliasTable
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.admin_ui_utils import (
|
||||
admin_ui_disabled,
|
||||
|
|
@ -209,31 +209,14 @@ def _team_detail_db(repo: TeamRepository) -> "TableActions[_TeamDetailRow]":
|
|||
return repo.table
|
||||
|
||||
|
||||
_MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str])
|
||||
_SSO_TOKEN_CLAIMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def _decode_model_aliases(value: object) -> object:
|
||||
"""``/team/new`` stores team model aliases as a JSON-encoded string in the Json column."""
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
try:
|
||||
return _MODEL_ALIASES_ADAPTER.validate_json(value)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
class _TeamModelAliasTable(BaseModel):
|
||||
model_config = ConfigDict(protected_namespaces=())
|
||||
|
||||
model_aliases: Annotated[Mapping[str, str] | None, BeforeValidator(_decode_model_aliases)] = None
|
||||
|
||||
|
||||
class _TeamRowGrants(BaseModel):
|
||||
team_id: str
|
||||
team_alias: str | None = None
|
||||
models: tuple[str, ...] = ()
|
||||
litellm_model_table: _TeamModelAliasTable | None = None
|
||||
litellm_model_table: TeamModelAliasTable | None = None
|
||||
|
||||
|
||||
class CliSsoTeamDetail(BaseModel):
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import time
|
|||
import traceback
|
||||
import warnings
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Collection, Mapping, MutableMapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import MappingProxyType, UnionType
|
||||
from typing import (
|
||||
|
|
@ -131,6 +132,7 @@ from litellm.router_utils.auto_router_tuning_baseline import (
|
|||
snapshot_tuning_baselines,
|
||||
tuning_limit_violation,
|
||||
)
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.utils import (
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
|
|
@ -138,11 +140,7 @@ from litellm.types.utils import (
|
|||
TextCompletionResponse,
|
||||
TokenCountResponse,
|
||||
)
|
||||
from litellm.utils import (
|
||||
_invalidate_model_cost_lowercase_map,
|
||||
load_credentials_from_list,
|
||||
reapply_runtime_model_cost_registrations,
|
||||
)
|
||||
from litellm.utils import load_credentials_from_list
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from aiohttp import ClientSession
|
||||
|
|
@ -426,6 +424,7 @@ from litellm.proxy.db.exception_handler import (
|
|||
)
|
||||
from litellm.proxy.db.gateway_request_tracking import (
|
||||
GatewayRequestAccumulator,
|
||||
GatewayRequestRedisBuffer,
|
||||
flush_gateway_requests,
|
||||
)
|
||||
from litellm.proxy.db.proxy_worker_heartbeat import (
|
||||
|
|
@ -2357,6 +2356,17 @@ open_telemetry_logger: OpenTelemetry | None = None
|
|||
gateway_request_accumulator: Final = GatewayRequestAccumulator()
|
||||
### INITIALIZE GLOBAL LOGGING OBJECT ###
|
||||
proxy_logging_obj: ProxyLogging = ProxyLogging(user_api_key_cache=user_api_key_cache, premium_user=premium_user)
|
||||
|
||||
|
||||
def _gateway_request_redis_buffer() -> GatewayRequestRedisBuffer | None:
|
||||
"""Shares the spend writer's transaction-buffer Redis and pod lock when use_redis_transaction_buffer is on."""
|
||||
writer: Final = proxy_logging_obj.db_spend_update_writer
|
||||
redis_cache: Final = writer.redis_update_buffer.redis_cache
|
||||
if redis_cache is None or not writer.redis_update_buffer._should_commit_spend_updates_to_redis():
|
||||
return None
|
||||
return GatewayRequestRedisBuffer(redis_cache=redis_cache, pod_lock_manager=writer.pod_lock_manager)
|
||||
|
||||
|
||||
### REDIS QUEUE ###
|
||||
async_result: Final = None
|
||||
celery_app_conn: Final = None
|
||||
|
|
@ -2707,6 +2717,12 @@ async def _read_spend_counter_estimate(counter_key: str, fallback_spend: float)
|
|||
return fallback_spend, False
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PendingSpendIncrement:
|
||||
counter_key: str
|
||||
increment: float
|
||||
|
||||
|
||||
async def increment_spend_counters(
|
||||
token: str | None,
|
||||
team_id: str | None,
|
||||
|
|
@ -2741,7 +2757,7 @@ async def increment_spend_counters(
|
|||
|
||||
cost: Final[float] = response_cost
|
||||
|
||||
async def _key_scope(key_token: str) -> None:
|
||||
async def _key_scope(key_token: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
# key_token arrives pre-hashed from metadata["user_api_key"] (auth flow
|
||||
# hashes raw "sk-..." keys before they reach the callback). The
|
||||
# startswith("sk-") check is a safety net matching update_cache —
|
||||
|
|
@ -2752,30 +2768,29 @@ async def increment_spend_counters(
|
|||
hash_token(token=key_token) if isinstance(key_token, str) and key_token.startswith("sk-") else key_token
|
||||
)
|
||||
key_counter_key: Final = f"spend:key:{hashed_token}"
|
||||
if key_counter_key not in reserved_counter_keys:
|
||||
await _init_and_increment_spend_counter(
|
||||
counter_key=key_counter_key,
|
||||
source_cache_key=hashed_token,
|
||||
increment=cost,
|
||||
key_pending: Final[tuple[_PendingSpendIncrement, ...]] = (
|
||||
()
|
||||
if key_counter_key in reserved_counter_keys
|
||||
else (
|
||||
await _prepare_spend_counter_increment(
|
||||
counter_key=key_counter_key,
|
||||
source_cache_key=hashed_token,
|
||||
increment=cost,
|
||||
),
|
||||
)
|
||||
|
||||
key_obj: Final[object] = await user_api_key_cache.async_get_cache(key=hashed_token)
|
||||
if key_obj is None:
|
||||
return
|
||||
key_budget_limits = getattr(key_obj, "budget_limits", None) or (
|
||||
key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
|
||||
)
|
||||
if isinstance(key_budget_limits, str):
|
||||
key_budget_limits = json.loads(key_budget_limits)
|
||||
if not isinstance(key_budget_limits, list):
|
||||
return
|
||||
for window in key_budget_limits:
|
||||
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
|
||||
key_window_reset_at = window.get("reset_at") if isinstance(window, dict) else window.reset_at
|
||||
key_window_counter = f"spend:key:{hashed_token}:window:{duration}"
|
||||
|
||||
async def _key_window_increment(window: object) -> _PendingSpendIncrement | None:
|
||||
duration = (
|
||||
window["budget_duration"] if isinstance(window, dict) else getattr(window, "budget_duration", None)
|
||||
)
|
||||
key_window_reset_at = (
|
||||
window.get("reset_at") if isinstance(window, dict) else getattr(window, "reset_at", None)
|
||||
)
|
||||
key_window_counter: Final = f"spend:key:{hashed_token}:window:{duration}"
|
||||
key_window_start = get_budget_window_start(window)
|
||||
if key_window_counter not in reserved_counter_keys:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending_window: Final = (
|
||||
await _prepare_window_spend_counter_increment(
|
||||
counter_key=key_window_counter,
|
||||
entity_type="Key",
|
||||
entity_id=hashed_token,
|
||||
|
|
@ -2783,6 +2798,9 @@ async def increment_spend_counters(
|
|||
window_start=key_window_start,
|
||||
increment=cost,
|
||||
)
|
||||
if key_window_counter not in reserved_counter_keys
|
||||
else None
|
||||
)
|
||||
await _enqueue_window_spend_row_update(
|
||||
entity_type=Litellm_EntityType.KEY,
|
||||
entity_id=hashed_token,
|
||||
|
|
@ -2792,33 +2810,48 @@ async def increment_spend_counters(
|
|||
increment=cost,
|
||||
request_started_at=request_started_at,
|
||||
)
|
||||
return pending_window
|
||||
|
||||
async def _team_scope(scope_team_id: str) -> None:
|
||||
team_counter_key: Final = f"spend:team:{scope_team_id}"
|
||||
if team_counter_key not in reserved_counter_keys:
|
||||
await _init_and_increment_spend_counter(
|
||||
counter_key=team_counter_key,
|
||||
source_cache_key=f"team_id:{scope_team_id}",
|
||||
increment=cost,
|
||||
)
|
||||
|
||||
team_obj: Final[object] = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}")
|
||||
if team_obj is None:
|
||||
return
|
||||
team_budget_limits = getattr(team_obj, "budget_limits", None) or (
|
||||
team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
|
||||
key_obj: Final[object] = await user_api_key_cache.async_get_cache(key=hashed_token)
|
||||
if key_obj is None:
|
||||
return key_pending
|
||||
key_budget_limits = getattr(key_obj, "budget_limits", None) or (
|
||||
key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
|
||||
)
|
||||
if isinstance(team_budget_limits, str):
|
||||
team_budget_limits = json.loads(team_budget_limits)
|
||||
if not isinstance(team_budget_limits, list):
|
||||
return
|
||||
for window in team_budget_limits:
|
||||
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
|
||||
team_window_reset_at = window.get("reset_at") if isinstance(window, dict) else window.reset_at
|
||||
team_window_counter = f"spend:team:{scope_team_id}:window:{duration}"
|
||||
if isinstance(key_budget_limits, str):
|
||||
key_budget_limits = json.loads(key_budget_limits)
|
||||
if not isinstance(key_budget_limits, list):
|
||||
return key_pending
|
||||
window_pending: Final = await asyncio.gather(
|
||||
*(_key_window_increment(window) for window in key_budget_limits), return_exceptions=True
|
||||
)
|
||||
return key_pending + tuple(item for item in window_pending if item is not None)
|
||||
|
||||
async def _team_scope(scope_team_id: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
team_counter_key: Final = f"spend:team:{scope_team_id}"
|
||||
team_pending: Final[tuple[_PendingSpendIncrement, ...]] = (
|
||||
()
|
||||
if team_counter_key in reserved_counter_keys
|
||||
else (
|
||||
await _prepare_spend_counter_increment(
|
||||
counter_key=team_counter_key,
|
||||
source_cache_key=f"team_id:{scope_team_id}",
|
||||
increment=cost,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
async def _team_window_increment(window: object) -> _PendingSpendIncrement | None:
|
||||
duration = (
|
||||
window["budget_duration"] if isinstance(window, dict) else getattr(window, "budget_duration", None)
|
||||
)
|
||||
team_window_reset_at = (
|
||||
window.get("reset_at") if isinstance(window, dict) else getattr(window, "reset_at", None)
|
||||
)
|
||||
team_window_counter: Final = f"spend:team:{scope_team_id}:window:{duration}"
|
||||
team_window_start = get_budget_window_start(window)
|
||||
if team_window_counter not in reserved_counter_keys:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending_window: Final = (
|
||||
await _prepare_window_spend_counter_increment(
|
||||
counter_key=team_window_counter,
|
||||
entity_type="Team",
|
||||
entity_id=scope_team_id,
|
||||
|
|
@ -2826,6 +2859,9 @@ async def increment_spend_counters(
|
|||
window_start=team_window_start,
|
||||
increment=cost,
|
||||
)
|
||||
if team_window_counter not in reserved_counter_keys
|
||||
else None
|
||||
)
|
||||
await _enqueue_window_spend_row_update(
|
||||
entity_type=Litellm_EntityType.TEAM,
|
||||
entity_id=scope_team_id,
|
||||
|
|
@ -2835,25 +2871,47 @@ async def increment_spend_counters(
|
|||
increment=cost,
|
||||
request_started_at=request_started_at,
|
||||
)
|
||||
return pending_window
|
||||
|
||||
async def _team_member_scope(scope_user_id: str, scope_team_id: str) -> None:
|
||||
team_obj: Final[object] = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}")
|
||||
if team_obj is None:
|
||||
return team_pending
|
||||
team_budget_limits = getattr(team_obj, "budget_limits", None) or (
|
||||
team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
|
||||
)
|
||||
if isinstance(team_budget_limits, str):
|
||||
team_budget_limits = json.loads(team_budget_limits)
|
||||
if not isinstance(team_budget_limits, list):
|
||||
return team_pending
|
||||
window_pending: Final = await asyncio.gather(
|
||||
*(_team_window_increment(window) for window in team_budget_limits), return_exceptions=True
|
||||
)
|
||||
return team_pending + tuple(item for item in window_pending if item is not None)
|
||||
|
||||
async def _team_member_scope(
|
||||
scope_user_id: str, scope_team_id: str
|
||||
) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
team_member_counter_key: Final = f"spend:team_member:{scope_user_id}:{scope_team_id}"
|
||||
if team_member_counter_key in reserved_counter_keys:
|
||||
return
|
||||
await _init_and_increment_spend_counter(
|
||||
counter_key=team_member_counter_key,
|
||||
source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}",
|
||||
increment=cost,
|
||||
return ()
|
||||
return (
|
||||
await _prepare_spend_counter_increment(
|
||||
counter_key=team_member_counter_key,
|
||||
source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}",
|
||||
increment=cost,
|
||||
),
|
||||
)
|
||||
|
||||
async def _user_scope(scope_user_id: str) -> None:
|
||||
async def _user_scope(scope_user_id: str) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
user_counter_key: Final = f"spend:user:{scope_user_id}"
|
||||
if user_counter_key in reserved_counter_keys:
|
||||
return
|
||||
await _init_and_increment_spend_counter(
|
||||
counter_key=user_counter_key,
|
||||
source_cache_key=scope_user_id,
|
||||
increment=cost,
|
||||
return ()
|
||||
return (
|
||||
await _prepare_spend_counter_increment(
|
||||
counter_key=user_counter_key,
|
||||
source_cache_key=scope_user_id,
|
||||
increment=cost,
|
||||
),
|
||||
)
|
||||
|
||||
scope_coros: Final = tuple(
|
||||
|
|
@ -2863,7 +2921,7 @@ async def increment_spend_counters(
|
|||
_team_scope(team_id) if team_id is not None else None,
|
||||
_team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None,
|
||||
_user_scope(user_id) if user_id is not None else None,
|
||||
_increment_end_user_and_tag_spend_counters(
|
||||
_prepare_end_user_and_tag_spend_increments(
|
||||
end_user_id=end_user_id,
|
||||
tags=tags,
|
||||
response_cost=cost,
|
||||
|
|
@ -2871,14 +2929,14 @@ async def increment_spend_counters(
|
|||
)
|
||||
if end_user_id is not None or tags is not None
|
||||
else None,
|
||||
_increment_model_access_group_spend_counters(
|
||||
_prepare_model_access_group_spend_increments(
|
||||
model_access_groups=model_access_groups,
|
||||
response_cost=cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
if model_access_groups
|
||||
else None,
|
||||
_increment_org_spend_counter(
|
||||
_prepare_org_spend_increment(
|
||||
org_id=org_id,
|
||||
response_cost=cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
|
|
@ -2893,7 +2951,20 @@ async def increment_spend_counters(
|
|||
# as orphaned tasks that race the caller's reservation-counter invalidation;
|
||||
# all scopes settle, then the first error propagates as before.
|
||||
scope_results: Final = await asyncio.gather(*scope_coros, return_exceptions=True)
|
||||
scope_errors: Final = [r for r in scope_results if isinstance(r, BaseException)]
|
||||
scope_errors: Final = tuple(
|
||||
item
|
||||
for scope in scope_results
|
||||
for item in (scope if isinstance(scope, tuple) else (scope,))
|
||||
if isinstance(item, BaseException)
|
||||
)
|
||||
pending: Final = tuple(
|
||||
item
|
||||
for scope in scope_results
|
||||
if not isinstance(scope, BaseException)
|
||||
for item in scope
|
||||
if not isinstance(item, BaseException)
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=pending)
|
||||
if scope_errors:
|
||||
raise scope_errors[0]
|
||||
|
||||
|
|
@ -2936,41 +3007,49 @@ async def _reconcile_budget_reservation_for_counter_update(
|
|||
return reserved_counter_keys
|
||||
|
||||
|
||||
async def _increment_end_user_and_tag_spend_counters(
|
||||
async def _prepare_end_user_and_tag_spend_increments(
|
||||
end_user_id: str | None,
|
||||
tags: list[str] | None,
|
||||
response_cost: float,
|
||||
reserved_counter_keys: set[str],
|
||||
) -> None:
|
||||
if end_user_id is not None:
|
||||
await _init_and_increment_unreserved_spend_counter(
|
||||
counter_key=f"spend:end_user:{end_user_id}",
|
||||
source_cache_key=end_user_cache_key(end_user_id),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
|
||||
if tags is None:
|
||||
return
|
||||
|
||||
seen_tags: Final[set[str]] = set()
|
||||
for tag_name in tags:
|
||||
if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags:
|
||||
continue
|
||||
seen_tags.add(tag_name)
|
||||
await _init_and_increment_unreserved_spend_counter(
|
||||
counter_key=f"spend:tag:{tag_name}",
|
||||
source_cache_key=tag_cache_key(tag_name),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
unique_tags: Final = (
|
||||
tuple(dict.fromkeys(tag for tag in tags if tag and isinstance(tag, str))) if tags is not None else ()
|
||||
)
|
||||
results: Final = await asyncio.gather(
|
||||
*(
|
||||
coro
|
||||
for coro in (
|
||||
_prepare_unreserved_spend_counter_increment(
|
||||
counter_key=f"spend:end_user:{end_user_id}",
|
||||
source_cache_key=end_user_cache_key(end_user_id),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
if end_user_id is not None
|
||||
else None,
|
||||
*(
|
||||
_prepare_unreserved_spend_counter_increment(
|
||||
counter_key=f"spend:tag:{tag_name}",
|
||||
source_cache_key=tag_cache_key(tag_name),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
for tag_name in unique_tags
|
||||
),
|
||||
)
|
||||
if coro is not None
|
||||
),
|
||||
return_exceptions=True,
|
||||
)
|
||||
return tuple(item for item in results if item is not None)
|
||||
|
||||
|
||||
async def _increment_model_access_group_spend_counters(
|
||||
async def _prepare_model_access_group_spend_increments(
|
||||
model_access_groups: Sequence[object],
|
||||
response_cost: float,
|
||||
reserved_counter_keys: set[str],
|
||||
) -> None:
|
||||
) -> tuple[_PendingSpendIncrement | BaseException, ...]:
|
||||
"""Charge the model access groups that authorized this request.
|
||||
|
||||
Without this the counter auth reads is written only by the reservation path, so
|
||||
|
|
@ -2984,55 +3063,63 @@ async def _increment_model_access_group_spend_counters(
|
|||
unique_groups: Final = tuple(
|
||||
dict.fromkeys(group for group in model_access_groups if group and isinstance(group, str))
|
||||
)
|
||||
for group in unique_groups:
|
||||
await _init_and_increment_unreserved_spend_counter(
|
||||
counter_key=model_access_group_spend_counter_key(group),
|
||||
source_cache_key=model_access_group_cache_key(group),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
results: Final = await asyncio.gather(
|
||||
*(
|
||||
_prepare_unreserved_spend_counter_increment(
|
||||
counter_key=model_access_group_spend_counter_key(group),
|
||||
source_cache_key=model_access_group_cache_key(group),
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
for group in unique_groups
|
||||
),
|
||||
return_exceptions=True,
|
||||
)
|
||||
return tuple(item for item in results if item is not None)
|
||||
|
||||
|
||||
async def _increment_org_spend_counter(
|
||||
async def _prepare_org_spend_increment(
|
||||
org_id: str | None,
|
||||
response_cost: float,
|
||||
reserved_counter_keys: set[str],
|
||||
) -> None:
|
||||
) -> tuple[_PendingSpendIncrement, ...]:
|
||||
if org_id is None:
|
||||
return
|
||||
return ()
|
||||
|
||||
await _init_and_increment_unreserved_spend_counter(
|
||||
pending: Final = await _prepare_unreserved_spend_counter_increment(
|
||||
counter_key=f"spend:org:{org_id}",
|
||||
source_cache_key=[f"org_id:{org_id}:with_budget", f"org_id:{org_id}"],
|
||||
increment=response_cost,
|
||||
reserved_counter_keys=reserved_counter_keys,
|
||||
)
|
||||
return (pending,) if pending is not None else ()
|
||||
|
||||
|
||||
async def _init_and_increment_unreserved_spend_counter(
|
||||
async def _prepare_unreserved_spend_counter_increment(
|
||||
counter_key: str,
|
||||
source_cache_key: str | list[str],
|
||||
increment: float,
|
||||
reserved_counter_keys: set[str],
|
||||
) -> None:
|
||||
) -> _PendingSpendIncrement | None:
|
||||
if counter_key in reserved_counter_keys:
|
||||
return
|
||||
return None
|
||||
|
||||
await _init_and_increment_spend_counter(
|
||||
return await _prepare_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
source_cache_key=source_cache_key,
|
||||
increment=increment,
|
||||
)
|
||||
|
||||
|
||||
async def _init_and_increment_spend_counter(
|
||||
async def _prepare_spend_counter_increment(
|
||||
counter_key: str,
|
||||
source_cache_key: str | list[str],
|
||||
increment: float,
|
||||
):
|
||||
) -> _PendingSpendIncrement:
|
||||
"""
|
||||
Initialize counter from the authoritative DB spend value if not yet
|
||||
set, then atomically increment in both in-memory and Redis.
|
||||
set, then return the pending increment for the caller to apply in one
|
||||
pipelined Redis call.
|
||||
|
||||
On first access per pod:
|
||||
1. Check spend_counter_cache (in-memory -> Redis via DualCache)
|
||||
|
|
@ -3044,13 +3131,13 @@ async def _init_and_increment_spend_counter(
|
|||
the counter as absent and seed it. Using increment means the worst case
|
||||
is over-counting (conservative, blocks slightly early) rather than
|
||||
under-counting (would allow overspend).
|
||||
4. Increment atomically (both in-memory + Redis)
|
||||
4. Increment is returned for the caller to apply via pipeline
|
||||
"""
|
||||
await _ensure_spend_counter_initialized(
|
||||
counter_key=counter_key,
|
||||
source_cache_key=source_cache_key,
|
||||
)
|
||||
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
|
||||
return _PendingSpendIncrement(counter_key=counter_key, increment=increment)
|
||||
|
||||
|
||||
async def _enqueue_window_spend_row_update(
|
||||
|
|
@ -3102,20 +3189,20 @@ async def _enqueue_window_spend_row_update(
|
|||
)
|
||||
|
||||
|
||||
async def _init_and_increment_window_spend_counter(
|
||||
async def _prepare_window_spend_counter_increment(
|
||||
counter_key: str,
|
||||
entity_type: str,
|
||||
entity_id: str,
|
||||
window_duration: str | None,
|
||||
window_start: datetime | None,
|
||||
increment: float,
|
||||
):
|
||||
) -> _PendingSpendIncrement | None:
|
||||
if window_start is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Skipping spend counter increment for invalid budget window %s",
|
||||
counter_key,
|
||||
)
|
||||
return
|
||||
return None
|
||||
|
||||
initialized: Final = await _ensure_window_spend_counter_initialized(
|
||||
counter_key=counter_key,
|
||||
|
|
@ -3125,8 +3212,8 @@ async def _init_and_increment_window_spend_counter(
|
|||
window_start=window_start,
|
||||
)
|
||||
if initialized is False:
|
||||
return
|
||||
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
|
||||
return None
|
||||
return _PendingSpendIncrement(counter_key=counter_key, increment=increment)
|
||||
|
||||
|
||||
async def _ensure_spend_counter_initialized(
|
||||
|
|
@ -3259,6 +3346,32 @@ async def _invalidate_spend_counter(counter_key: str):
|
|||
)
|
||||
|
||||
|
||||
async def _apply_spend_counter_increments(pending: Sequence[_PendingSpendIncrement]) -> None:
|
||||
if not pending:
|
||||
return
|
||||
redis_cache: Final = spend_counter_cache.redis_cache
|
||||
if redis_cache is None:
|
||||
for item in pending:
|
||||
await spend_counter_cache.async_increment_cache(
|
||||
key=item.counter_key,
|
||||
value=item.increment,
|
||||
refresh_ttl=True,
|
||||
)
|
||||
return
|
||||
ttl: Final = redis_cache.get_ttl()
|
||||
increment_list: Final = [ # mutable-ok: async_increment_pipeline signature requires list[RedisPipelineIncrementOperation]
|
||||
RedisPipelineIncrementOperation(key=item.counter_key, increment_value=item.increment, ttl=ttl)
|
||||
for item in pending
|
||||
]
|
||||
try:
|
||||
results: Final = await redis_cache.async_increment_pipeline(increment_list=increment_list)
|
||||
except Exception:
|
||||
await asyncio.gather(*(_invalidate_spend_counter(counter_key=item.counter_key) for item in pending))
|
||||
raise
|
||||
for item, current_value in zip(pending, results or ()):
|
||||
spend_counter_cache.in_memory_cache.set_cache(key=item.counter_key, value=current_value)
|
||||
|
||||
|
||||
async def update_cache(
|
||||
token: str | None,
|
||||
user_id: str | None,
|
||||
|
|
@ -4436,20 +4549,9 @@ def resolve_classifier_plugin(
|
|||
|
||||
|
||||
def _swap_in_model_cost_map(new_model_cost_map: dict) -> int:
|
||||
"""Adopt a freshly fetched cost map into this process's litellm state, return the model count"""
|
||||
litellm.model_cost = new_model_cost_map
|
||||
# Invalidate case-insensitive lookup map since model_cost was replaced
|
||||
_invalidate_model_cost_lowercase_map()
|
||||
# Repopulate provider model sets (e.g. litellm.anthropic_models) so that
|
||||
# wildcard patterns like "anthropic/*" include any newly added models.
|
||||
litellm.add_known_models(model_cost_map=new_model_cost_map)
|
||||
# Counted before the re-apply below, which writes into this same dict, so the
|
||||
# number reported describes the fetched price data alone.
|
||||
fetched_model_count: Final = len(new_model_cost_map) if new_model_cost_map else 0
|
||||
# The swap discards everything registered at runtime (deployment model_info,
|
||||
# register_model overrides), so put it back on top of the fresh catalog.
|
||||
reapply_runtime_model_cost_registrations()
|
||||
return fetched_model_count
|
||||
from litellm.litellm_core_utils.get_model_cost_map import adopt_model_cost_map
|
||||
|
||||
return adopt_model_cost_map(new_model_cost_map)
|
||||
|
||||
|
||||
def should_load_db_object(object_type: str | SupportedDBObjectType) -> bool:
|
||||
|
|
@ -9543,7 +9645,7 @@ class ProxyStartupEvent:
|
|||
flush_gateway_requests,
|
||||
"interval",
|
||||
seconds=batch_writing_interval,
|
||||
args=(prisma_client, gateway_request_accumulator),
|
||||
args=(prisma_client, gateway_request_accumulator, _gateway_request_redis_buffer()),
|
||||
id="update_gateway_requests_job",
|
||||
replace_existing=True,
|
||||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
|
|
|
|||
|
|
@ -543,6 +543,49 @@ class ResponsesAPIRequestUtils:
|
|||
|
||||
return request_input
|
||||
|
||||
@staticmethod
|
||||
def strip_encrypted_reasoning_from_input(request_input: object) -> None:
|
||||
"""Drop reasoning items the routed deployment cannot decrypt, keeping their readable summary.
|
||||
|
||||
Mutates ``request_input`` in place: the router's fallback snapshot shares this
|
||||
list object, so a rebound list would replay the stripped items on the fallback hop.
|
||||
"""
|
||||
if not isinstance(request_input, list):
|
||||
return
|
||||
items: Final = cast(list[object], request_input) # cast-ok: untyped client json
|
||||
stripped: Final = tuple(ResponsesAPIRequestUtils._without_encrypted_reasoning(item) for item in items)
|
||||
items[:] = (item for item in stripped if item is not None) # rebind-ok: list shared with fallback snapshot
|
||||
|
||||
@staticmethod
|
||||
def _without_encrypted_reasoning(item: object) -> object | None:
|
||||
if not isinstance(item, dict):
|
||||
return item
|
||||
reasoning: Final = cast(Mapping[str, object], item) # cast-ok: untyped client json
|
||||
if reasoning.get("type") != "reasoning" or not reasoning.get("encrypted_content"):
|
||||
return reasoning
|
||||
readable: Final = any(
|
||||
ResponsesAPIRequestUtils._has_readable_text(reasoning.get(key)) for key in ("summary", "content")
|
||||
)
|
||||
if not readable:
|
||||
return None
|
||||
kept: Final[dict[str, object]] = { # mutable-ok: request item rebuilt without the undecryptable keys
|
||||
key: value for key, value in reasoning.items() if key not in ("encrypted_content", "id")
|
||||
}
|
||||
return kept
|
||||
|
||||
@staticmethod
|
||||
def _has_readable_text(value: object) -> bool:
|
||||
"""A reasoning item's ``summary``/``content`` carries readable text: a non-empty string, or a
|
||||
list holding at least one block with a non-empty ``text`` field (summary_text / output_text)."""
|
||||
if isinstance(value, str):
|
||||
return bool(value.strip())
|
||||
if isinstance(value, list):
|
||||
return any(
|
||||
isinstance(block, dict) and bool(cast(Mapping[str, object], block).get("text")) # cast-ok: untyped json
|
||||
for block in value
|
||||
)
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _build_responses_api_response_id(
|
||||
custom_llm_provider: str | None,
|
||||
|
|
|
|||
|
|
@ -1317,6 +1317,43 @@ class Router:
|
|||
if isinstance(litellm.input_callback, list):
|
||||
litellm.input_callback = [c for c in litellm.input_callback if id(c) not in selector_ids]
|
||||
|
||||
def _apply_updated_routing_strategy_args(self) -> None:
|
||||
"""
|
||||
Re-link the default group's selector to the current `routing_strategy_args`.
|
||||
|
||||
Selectors freeze their `RoutingArgs` at construction, so a runtime args
|
||||
update would otherwise keep serving the boot-time values until restart.
|
||||
Latency/usage state survives the rebuild: it lives in the shared router
|
||||
cache, not on the selector.
|
||||
"""
|
||||
strategy: Final = self._normalize_strategy(self.routing_strategy)
|
||||
if strategy == "lar1":
|
||||
from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy
|
||||
|
||||
apply_lar1_routing_strategy(self, self.routing_strategy_args)
|
||||
return
|
||||
|
||||
attr: Final = self._DEFAULT_SELECTOR_ATTR_BY_STRATEGY.get(strategy or "")
|
||||
current: Final = getattr(self, attr, None) if attr is not None else None
|
||||
if attr is None or current is None:
|
||||
return
|
||||
|
||||
try:
|
||||
rebuilt: Final = self._build_strategy_selector(
|
||||
strategy=strategy or "",
|
||||
routing_strategy_args=self.routing_strategy_args,
|
||||
)
|
||||
except (TypeError, ValidationError):
|
||||
verbose_router_logger.exception(
|
||||
"Invalid routing_strategy_args %s for '%s'; keeping the previous ones",
|
||||
self.routing_strategy_args,
|
||||
strategy,
|
||||
)
|
||||
return
|
||||
|
||||
self._unregister_router_selectors((current,))
|
||||
setattr(self, attr, rebuilt)
|
||||
|
||||
def routing_strategy_init(self, routing_strategy: RoutingStrategy | str, routing_strategy_args: dict):
|
||||
verbose_router_logger.info("Routing strategy: %s", routing_strategy)
|
||||
self._validate_routing_strategy(routing_strategy)
|
||||
|
|
@ -11130,6 +11167,40 @@ class Router:
|
|||
|
||||
return ids
|
||||
|
||||
def get_candidate_model_ids_for_route(self, model: str, team_id: str | None = None) -> frozenset[str]:
|
||||
"""
|
||||
Deployment ids that could serve ``model`` for ``team_id``, unioned across the paths
|
||||
the router resolves a route through: ``model_group_alias``, a routing group, the
|
||||
``model_name`` and team indexes, and wildcard pattern routes. Read-only and
|
||||
side-effect-free, unlike ``_common_checks_available_deployment`` which also applies
|
||||
fallbacks and can raise. Lets a pre-call check tell a genuine cross-group route from
|
||||
same-group unavailability without re-deriving that precedence at the call site, and
|
||||
without leaking deployment ids into request kwargs bound for the provider.
|
||||
"""
|
||||
resolved: Final = self._get_model_from_alias(model=model) or model
|
||||
routing_group_members: Final = self._get_routing_group_deployments(model=resolved, team_id=team_id)
|
||||
if routing_group_members is not None:
|
||||
return self._deployment_ids(routing_group_members)
|
||||
if resolved in self.model_names:
|
||||
return self._deployment_ids(self._get_all_deployments(model_name=resolved, team_id=team_id))
|
||||
team_router: Final = self.team_pattern_routers.get(team_id) if team_id is not None else None
|
||||
return self._deployment_ids(
|
||||
(
|
||||
*self._get_all_deployments(model_name=resolved, team_id=team_id),
|
||||
*(self.pattern_router.route(resolved) or ()),
|
||||
*((team_router.route(resolved) or ()) if team_router is not None else ()),
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _deployment_ids(deployments: Sequence[Mapping[str, object]]) -> frozenset[str]:
|
||||
return frozenset(
|
||||
str(model_info["id"])
|
||||
for deployment in deployments
|
||||
for model_info in (deployment.get("model_info"),)
|
||||
if isinstance(model_info, Mapping) and model_info.get("id") is not None
|
||||
)
|
||||
|
||||
def has_model_id(self, candidate_id: str) -> bool:
|
||||
"""
|
||||
O(1) membership check for a deployment ID without allocating large lists.
|
||||
|
|
@ -11847,7 +11918,7 @@ class Router:
|
|||
|
||||
_existing_router_settings: Final = self.get_settings()
|
||||
rebuild_routing_groups = False
|
||||
relink_lar1_from_args = False
|
||||
routing_args_updated = False
|
||||
for var in kwargs:
|
||||
if var in RUNTIME_UPDATABLE_ROUTER_SETTINGS:
|
||||
if var in _int_settings:
|
||||
|
|
@ -11886,15 +11957,13 @@ class Router:
|
|||
)
|
||||
rebuild_routing_groups = True
|
||||
elif var == "routing_strategy_args":
|
||||
relink_lar1_from_args = True
|
||||
routing_args_updated = True
|
||||
setattr(self, var, value)
|
||||
else:
|
||||
verbose_router_logger.debug("Setting %s is not allowed", var)
|
||||
|
||||
if relink_lar1_from_args and self._normalize_strategy(self.routing_strategy) == "lar1":
|
||||
from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy
|
||||
|
||||
apply_lar1_routing_strategy(self, self.routing_strategy_args)
|
||||
if routing_args_updated:
|
||||
self._apply_updated_routing_strategy_args()
|
||||
|
||||
if rebuild_routing_groups:
|
||||
self._init_routing_groups(self._routing_groups_input)
|
||||
|
|
|
|||
|
|
@ -3,8 +3,11 @@
|
|||
import random
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime, timedelta
|
||||
from math import ceil
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
import litellm
|
||||
from litellm import ModelResponse, token_counter, verbose_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -24,6 +27,7 @@ class RoutingArgs(LiteLLMPydanticObjectBase):
|
|||
ttl: float = 1 * 60 * 60 # 1 hour
|
||||
lowest_latency_buffer: float = 0
|
||||
max_latency_list_size: int = 10
|
||||
ttft_percentile: float | None = Field(default=None, gt=0, le=1)
|
||||
|
||||
|
||||
def _average_latency(samples: Sequence[float]) -> float:
|
||||
|
|
@ -32,6 +36,12 @@ def _average_latency(samples: Sequence[float]) -> float:
|
|||
return sum(samples) / len(samples)
|
||||
|
||||
|
||||
def _percentile_latency(samples: Sequence[float], percentile: float) -> float:
|
||||
values: Final = sorted(samples)
|
||||
index: Final = ceil(len(values) * percentile) - 1
|
||||
return values[index]
|
||||
|
||||
|
||||
def _ttft_seconds(elapsed: timedelta | float) -> float:
|
||||
if isinstance(elapsed, timedelta):
|
||||
return elapsed.total_seconds()
|
||||
|
|
@ -427,14 +437,17 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
item_rpm = item_map.get(precise_minute, {}).get("rpm", 0)
|
||||
item_tpm = item_map.get(precise_minute, {}).get("tpm", 0)
|
||||
|
||||
# get average latency or average ttft (depending on streaming/non-streaming)
|
||||
use_ttft = (
|
||||
request_kwargs is not None
|
||||
and request_kwargs.get("stream", None) is not None
|
||||
and request_kwargs["stream"] is True
|
||||
and len(item_ttft_latency) > 0
|
||||
)
|
||||
average_latency = _average_latency(item_ttft_latency if use_ttft else item_latency)
|
||||
selected_latency = (
|
||||
_percentile_latency(item_ttft_latency, self.routing_args.ttft_percentile)
|
||||
if use_ttft and self.routing_args.ttft_percentile is not None
|
||||
else _average_latency(item_ttft_latency if use_ttft else item_latency)
|
||||
)
|
||||
|
||||
# -------------- #
|
||||
# Debugging Logic
|
||||
|
|
@ -443,7 +456,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
# this helps a user to debug why the router picked a specfic deployment #
|
||||
_deployment_api_base = _deployment.get("litellm_params", {}).get("api_base", "")
|
||||
if _deployment_api_base is not None:
|
||||
_latency_per_deployment[_deployment_api_base] = average_latency
|
||||
_latency_per_deployment[_deployment_api_base] = selected_latency
|
||||
# -------------- #
|
||||
# End of Debugging Logic
|
||||
# -------------- #
|
||||
|
|
@ -453,7 +466,7 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
): # if user passed in tpm / rpm in the model_list
|
||||
continue
|
||||
else:
|
||||
potential_deployments.append((_deployment, average_latency))
|
||||
potential_deployments.append((_deployment, selected_latency))
|
||||
|
||||
if len(potential_deployments) == 0:
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -37,13 +37,13 @@ Safe to enable globally:
|
|||
"""
|
||||
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Final, Optional, Protocol, cast
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.exceptions import (
|
||||
BadRequestError,
|
||||
RateLimitError,
|
||||
ServiceUnavailableError,
|
||||
)
|
||||
|
|
@ -158,6 +158,23 @@ class EncryptedContentAffinityCheck(CustomLogger):
|
|||
return deployment
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _request_team_id(request_kwargs: Mapping[str, object]) -> str | None:
|
||||
containers: Final = (request_kwargs.get("metadata"), request_kwargs.get("litellm_metadata"))
|
||||
team_ids: Final = (c.get("user_api_key_team_id") for c in containers if isinstance(c, Mapping))
|
||||
return next((tid for tid in team_ids if isinstance(tid, str)), None)
|
||||
|
||||
def _routed_group_candidate_model_ids(self, request_kwargs: Mapping[str, object], model: str) -> frozenset[str]:
|
||||
"""
|
||||
Deployment ids that could serve this turn's routed ``model``, as the router
|
||||
resolves a route (model_group_alias / routing group / model_name / team /
|
||||
pattern). Delegates to the router so the full precedence is not re-derived here
|
||||
and no deployment ids are written into request kwargs bound for the provider.
|
||||
"""
|
||||
if self.router is None:
|
||||
return frozenset()
|
||||
return self.router.get_candidate_model_ids_for_route(model=model, team_id=self._request_team_id(request_kwargs))
|
||||
|
||||
@staticmethod
|
||||
def _encryption_boundary_key(
|
||||
litellm_params: object,
|
||||
|
|
@ -225,10 +242,14 @@ class EncryptedContentAffinityCheck(CustomLogger):
|
|||
"""
|
||||
If the request ``input`` contains litellm-encoded item IDs, decode the
|
||||
embedded ``model_id`` and pin the request to that deployment. Raises
|
||||
``RateLimitError`` / ``ServiceUnavailableError`` / ``BadRequestError``
|
||||
when the originating deployment is unavailable and no encryption-boundary
|
||||
peer exists, rather than dispatching a doomed request to a non-peer
|
||||
deployment. The 429/503 split mirrors the originating cooldown's status:
|
||||
``RateLimitError`` / ``ServiceUnavailableError`` when the originating
|
||||
deployment is a member of the routed model group but currently unavailable
|
||||
and no encryption-boundary peer exists, rather than dispatching a doomed
|
||||
request to a non-peer deployment. When the origin is not a member of the
|
||||
routed group (an auto-router tier change, a model switch with no peer, a
|
||||
removed deployment, or an unknown/forged marker), the encrypted reasoning is
|
||||
stripped and the request dispatches with its readable history instead. The
|
||||
429/503 split mirrors the originating cooldown's status:
|
||||
a 429-induced cooldown surfaces as 429 (with ``Retry-After`` set to the
|
||||
remaining cooldown window) so OpenAI-compatible clients back off and
|
||||
retry after the deployment is eligible again.
|
||||
|
|
@ -285,12 +306,34 @@ class EncryptedContentAffinityCheck(CustomLogger):
|
|||
request_kwargs["_encrypted_content_affinity_pinned"] = True
|
||||
return boundary_matches
|
||||
|
||||
# Dispatching to a non-peer would guarantee an upstream
|
||||
# `invalid_encrypted_content` 400, so fail fast with a clearer error.
|
||||
# The origin cannot serve this turn's routed group and no peer shares the boundary, so its
|
||||
# encrypted reasoning can never decrypt here. Strip it, keep the readable history, and dispatch
|
||||
# to the routed group instead of failing. Membership is tested by deployment id against the set
|
||||
# the router actually resolved for this route, not by model-group name, so an alias, a
|
||||
# provider-qualified spelling, a team-public name, or a pattern route of the same group is not
|
||||
# mistaken for a tier change. An unknown origin (a removed deployment, or a forged marker) is
|
||||
# treated the same as a cross-group one, which also denies an authenticated caller a
|
||||
# deployment-id existence oracle: a real cross-group id and a nonexistent id both strip and
|
||||
# dispatch rather than returning distinguishable responses. Only a genuine same-group member
|
||||
# that is currently unavailable falls through to the fail-fast, preserving the cooldown contract.
|
||||
routed_group_model_ids: Final = (
|
||||
self._routed_group_candidate_model_ids(request_kwargs, model) if originating is not None else frozenset()
|
||||
)
|
||||
if str(model_id) not in routed_group_model_ids:
|
||||
verbose_router_logger.debug(
|
||||
"EncryptedContentAffinityCheck: model_id=%s is not a candidate for the routed group %s; "
|
||||
"forwarding without its encrypted reasoning",
|
||||
model_id,
|
||||
model,
|
||||
)
|
||||
ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
|
||||
return typed_healthy_deployments
|
||||
|
||||
# The origin is a member of the routed group but currently unavailable (cooled down); fail fast
|
||||
# rather than dispatching to a non-peer, which would guarantee an upstream 400.
|
||||
raise await self._unavailable_origin_error(
|
||||
model=model,
|
||||
model_id=model_id,
|
||||
originating=originating,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
|
|
@ -298,25 +341,11 @@ class EncryptedContentAffinityCheck(CustomLogger):
|
|||
self,
|
||||
model: str,
|
||||
model_id: str,
|
||||
originating: Deployment | None,
|
||||
parent_otel_span: Span | None,
|
||||
) -> Exception:
|
||||
# Public error messages intentionally omit the originating ``model_id`` so
|
||||
# an authenticated caller forging encrypted-content markers cannot use the
|
||||
# error surface to enumerate which deployment IDs exist on this router.
|
||||
if originating is None:
|
||||
return BadRequestError(
|
||||
message=(
|
||||
"The deployment that produced this encrypted_content is no "
|
||||
"longer configured on this router, and no deployment on the "
|
||||
"same encryption boundary is available. Re-issue the request "
|
||||
"without the stale encrypted_content items, or restore the "
|
||||
"originating deployment."
|
||||
),
|
||||
model=model,
|
||||
llm_provider="",
|
||||
)
|
||||
|
||||
cooldown: Final = await self._get_origin_cooldown(model_id=model_id, parent_otel_span=parent_otel_span)
|
||||
|
||||
if cooldown is not None and str(cooldown.get("status_code")) == "429":
|
||||
|
|
|
|||
|
|
@ -3079,7 +3079,7 @@ def register_model(
|
|||
# Convert stringified numbers to appropriate numeric types
|
||||
loaded_model_cost = model_cost
|
||||
elif isinstance(model_cost, str):
|
||||
loaded_model_cost = litellm.get_model_cost_map(url=model_cost)
|
||||
loaded_model_cost = litellm.get_model_cost_map(url=model_cost, max_attempts=1)
|
||||
|
||||
if persist_across_reloads:
|
||||
_registrations: Final[Mapping[str, Mapping[str, object]]] = loaded_model_cost
|
||||
|
|
|
|||
|
|
@ -87,6 +87,7 @@ ignored_function_names = [
|
|||
"_delete_claude_code_session_router_binding", # Tested through Redis cleanup failure in test_router.py
|
||||
"_resolve_claude_code_session_router", # Tested through Claude Code session routing in test_router.py
|
||||
"_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py
|
||||
"_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name)
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,9 +1,12 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from collections.abc import AsyncIterator
|
||||
from importlib import metadata
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import anyio
|
||||
|
|
@ -11,15 +14,19 @@ import httpx
|
|||
import pytest
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth
|
||||
from mcp import McpError
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
from pydantic import ValidationError
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.types import (
|
||||
LATEST_PROTOCOL_VERSION,
|
||||
CallToolResult,
|
||||
ErrorData,
|
||||
Implementation,
|
||||
InitializeResult,
|
||||
JSONRPCError,
|
||||
JSONRPCMessage,
|
||||
JSONRPCResponse,
|
||||
LoggingMessageNotificationParams,
|
||||
ServerCapabilities,
|
||||
)
|
||||
|
||||
|
|
@ -29,8 +36,9 @@ import litellm.experimental_mcp_client.client as mcp_client_module
|
|||
from litellm.experimental_mcp_client.client import (
|
||||
MCP_STREAMABLE_HTTP_REQUIREMENT,
|
||||
MCPClient,
|
||||
_as_read_timeout,
|
||||
_first_non_cancelled_cause,
|
||||
_TransportContext,
|
||||
as_mcp_read_timeout,
|
||||
missing_streamable_http_client_error,
|
||||
strip_auth_scheme,
|
||||
)
|
||||
|
|
@ -859,25 +867,25 @@ def _raise_mcp_error_while_handling_a_timeout(code: int, message: str) -> McpErr
|
|||
return raised
|
||||
|
||||
|
||||
def test_as_read_timeout_separates_the_sdk_timeout_from_a_relayed_upstream_error():
|
||||
def test_as_mcp_read_timeout_separates_the_sdk_timeout_from_a_relayed_upstream_error():
|
||||
"""Neither signal alone is enough. The code alone cannot separate the SDK's own timeout from an
|
||||
upstream JSON-RPC error that happens to use 408, and the context chain alone cannot separate it
|
||||
from any other relayed error that surfaces while a timeout is being handled, so both must hold.
|
||||
"""
|
||||
timeout_code = int(httpx.codes.REQUEST_TIMEOUT)
|
||||
|
||||
translated = _as_read_timeout(_raise_mcp_error_while_handling_a_timeout(timeout_code, "Timed out while waiting"))
|
||||
translated = as_mcp_read_timeout(_raise_mcp_error_while_handling_a_timeout(timeout_code, "Timed out while waiting"))
|
||||
assert isinstance(translated, TimeoutError)
|
||||
assert str(translated) == "Timed out while waiting"
|
||||
|
||||
relayed_408 = McpError(ErrorData(code=timeout_code, message="upstream said 408"))
|
||||
assert _as_read_timeout(relayed_408) is None, "an upstream 408 with no elapsed timeout is not our timeout"
|
||||
assert as_mcp_read_timeout(relayed_408) is None, "an upstream 408 with no elapsed timeout is not our timeout"
|
||||
|
||||
relayed_other = _raise_mcp_error_while_handling_a_timeout(-32603, "upstream internal error")
|
||||
assert _as_read_timeout(relayed_other) is None, "a non-timeout code is not our timeout, whatever the chain"
|
||||
assert as_mcp_read_timeout(relayed_other) is None, "a non-timeout code is not our timeout, whatever the chain"
|
||||
|
||||
assert _as_read_timeout(McpError(ErrorData(code=-32603, message="boom"))) is None
|
||||
assert _as_read_timeout(RuntimeError("not an McpError")) is None
|
||||
assert as_mcp_read_timeout(McpError(ErrorData(code=-32603, message="boom"))) is None
|
||||
assert as_mcp_read_timeout(RuntimeError("not an McpError")) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1224,14 +1232,14 @@ def test_without_a_configured_slot_the_existing_precedence_is_unchanged():
|
|||
|
||||
|
||||
_REDIRECT_CASES = [
|
||||
("https://upstream.example.com/mcp", "https://upstream.example.com/other"), # same origin
|
||||
("https://upstream.example.com/mcp", "https://upstream.example.com/other"), # same origin
|
||||
("https://upstream.example.com/mcp", "https://upstream.example.com:443/other"), # explicit default port
|
||||
("https://upstream.example.com/mcp", "https://attacker.example.com/collect"), # different host
|
||||
("https://upstream.example.com/mcp", "http://upstream.example.com/collect"), # scheme downgrade
|
||||
("https://upstream.example.com/mcp", "https://upstream.example.com:8443/other"), # different port
|
||||
("https://upstream.example.com/mcp", "https://sub.upstream.example.com/x"), # different host
|
||||
("http://upstream.example.com/mcp", "https://upstream.example.com/other"), # http -> https upgrade
|
||||
("http://upstream.example.com/mcp", "http://upstream.example.com/other"), # same origin, plain http
|
||||
("https://upstream.example.com/mcp", "https://attacker.example.com/collect"), # different host
|
||||
("https://upstream.example.com/mcp", "http://upstream.example.com/collect"), # scheme downgrade
|
||||
("https://upstream.example.com/mcp", "https://upstream.example.com:8443/other"), # different port
|
||||
("https://upstream.example.com/mcp", "https://sub.upstream.example.com/x"), # different host
|
||||
("http://upstream.example.com/mcp", "https://upstream.example.com/other"), # http -> https upgrade
|
||||
("http://upstream.example.com/mcp", "http://upstream.example.com/other"), # same origin, plain http
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -1283,3 +1291,398 @@ def test_a_differently_cased_injected_header_cannot_shadow_the_slot() -> None:
|
|||
headers = client._get_auth_headers()
|
||||
assert [v for k, v in headers.items() if k.lower() == "esb-oauth"] == ["Bearer minted-token"]
|
||||
assert headers["X-Trace"] == "keep"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("content_type", "body", "expected_type"),
|
||||
[
|
||||
("text/html", b"<html>secret-page</html>", ValueError),
|
||||
("application/json", b"secret-invalid-json", ValidationError),
|
||||
("application/json", b"", ValidationError),
|
||||
("application/json", b'{"secret":"invalid-rpc"}', ValidationError),
|
||||
("application/json", b'{"jsonrpc":"2.0","id":0,"result":{"secret":"invalid-schema"}}', ValidationError),
|
||||
],
|
||||
)
|
||||
async def test_invalid_http_response_surfaces_without_waiting_for_timeout(
|
||||
content_type: str, body: bytes, expected_type: type[Exception]
|
||||
) -> None:
|
||||
from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, headers={"Content-Type": content_type}, content=body)
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
|
||||
client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
|
||||
with pytest.raises(expected_type) as caught:
|
||||
await asyncio.wait_for(
|
||||
client._execute_session_operation(
|
||||
streamable_http_client(client.server_url, http_client=http_client),
|
||||
lambda session: session.list_tools(),
|
||||
),
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
message: Final = _connection_error_message(caught.value, client.server_url, 30)
|
||||
assert "unsupported content type" in message or "invalid MCP response" in message
|
||||
assert "secret" not in message
|
||||
assert "timed out" not in message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("status_code", [200, 401, 503])
|
||||
async def test_http_response_handler_preserves_success_and_http_errors(status_code: int) -> None:
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
if request.method == "DELETE":
|
||||
return httpx.Response(200)
|
||||
payload: Final = json.loads(request.content)
|
||||
if "id" not in payload:
|
||||
return httpx.Response(202)
|
||||
result: Final = (
|
||||
{
|
||||
"protocolVersion": LATEST_PROTOCOL_VERSION,
|
||||
"capabilities": {},
|
||||
"serverInfo": {"name": "test", "version": "1"},
|
||||
}
|
||||
if payload["method"] == "initialize"
|
||||
else {"tools": []}
|
||||
)
|
||||
return httpx.Response(status_code, json={"jsonrpc": "2.0", "id": payload["id"], "result": result})
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
|
||||
client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
|
||||
operation: Final = client._execute_session_operation(
|
||||
streamable_http_client(client.server_url, http_client=http_client), lambda session: session.list_tools()
|
||||
)
|
||||
if status_code == 200:
|
||||
result: Final = await asyncio.wait_for(operation, timeout=3)
|
||||
assert result.tools == []
|
||||
else:
|
||||
with pytest.raises(httpx.HTTPStatusError) as caught:
|
||||
await asyncio.wait_for(operation, timeout=3)
|
||||
assert caught.value.response.status_code == status_code
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_http_response_handler_preserves_notifications_and_tool_listing() -> None:
|
||||
notification: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"method": "notifications/message",
|
||||
"params": {"level": "info", "data": "Listing tools"},
|
||||
}
|
||||
logging_callback: Final = AsyncMock()
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
if request.method == "DELETE":
|
||||
return httpx.Response(200)
|
||||
payload: Final = json.loads(request.content)
|
||||
if "id" not in payload:
|
||||
return httpx.Response(202)
|
||||
if payload["method"] == "initialize":
|
||||
return httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"jsonrpc": "2.0",
|
||||
"id": payload["id"],
|
||||
"result": {
|
||||
"protocolVersion": LATEST_PROTOCOL_VERSION,
|
||||
"capabilities": {"logging": {}, "tools": {}},
|
||||
"serverInfo": {"name": "test", "version": "1"},
|
||||
},
|
||||
},
|
||||
)
|
||||
response: Final = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": payload["id"],
|
||||
"result": {"tools": [{"name": "search", "inputSchema": {"type": "object"}}]},
|
||||
}
|
||||
return httpx.Response(
|
||||
200,
|
||||
headers={"Content-Type": "text/event-stream"},
|
||||
content="".join(f"event: message\ndata: {json.dumps(message)}\n\n" for message in (notification, response)),
|
||||
)
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
|
||||
client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30, logging_callback=logging_callback)
|
||||
result: Final = await asyncio.wait_for(
|
||||
client._execute_session_operation(
|
||||
streamable_http_client(client.server_url, http_client=http_client), lambda session: session.list_tools()
|
||||
),
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
assert [tool.name for tool in result.tools] == ["search"]
|
||||
logging_callback.assert_awaited_once_with(LoggingMessageNotificationParams(level="info", data="Listing tools"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_tool_list_schema_is_identified_as_an_upstream_response() -> None:
|
||||
from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
if request.method == "DELETE":
|
||||
return httpx.Response(200)
|
||||
payload: Final = json.loads(request.content)
|
||||
if "id" not in payload:
|
||||
return httpx.Response(202)
|
||||
result: Final = (
|
||||
{
|
||||
"protocolVersion": LATEST_PROTOCOL_VERSION,
|
||||
"capabilities": {},
|
||||
"serverInfo": {"name": "test", "version": "1"},
|
||||
}
|
||||
if payload["method"] == "initialize"
|
||||
else {"tools": "secret-invalid-tools"}
|
||||
)
|
||||
return httpx.Response(200, json={"jsonrpc": "2.0", "id": payload["id"], "result": result})
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
|
||||
client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
|
||||
with pytest.raises(ValidationError) as caught:
|
||||
await asyncio.wait_for(
|
||||
client._execute_session_operation(
|
||||
streamable_http_client(client.server_url, http_client=http_client),
|
||||
lambda session: session.list_tools(),
|
||||
),
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
message: Final = _connection_error_message(caught.value, client.server_url, 30)
|
||||
assert "invalid MCP response" in message
|
||||
assert "secret" not in message
|
||||
|
||||
|
||||
class _DiagnosticSSEStream(httpx.AsyncByteStream):
|
||||
def __init__(self, messages: asyncio.Queue[bytes | Exception | None]) -> None:
|
||||
self.messages = messages
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
yield b"event: endpoint\ndata: /messages\n\n"
|
||||
while True:
|
||||
message: Final = await self.messages.get()
|
||||
if message is None:
|
||||
return
|
||||
if isinstance(message, Exception):
|
||||
raise message
|
||||
yield b"event: message\ndata: " + message + b"\n\n"
|
||||
|
||||
|
||||
_DIAGNOSTIC_STDIO_SERVER: Final = """
|
||||
import json, sys
|
||||
mode, failure_method = sys.argv[1:]
|
||||
for line in sys.stdin:
|
||||
request = json.loads(line)
|
||||
if "method" not in request or "id" not in request:
|
||||
continue
|
||||
if request["method"] == failure_method:
|
||||
if mode == "bad-json":
|
||||
print("secret-invalid-json", flush=True)
|
||||
continue
|
||||
if mode == "closed":
|
||||
sys.exit(0)
|
||||
if mode == "silent":
|
||||
print(json.dumps({"jsonrpc": "2.0", "method": "notifications/message", "params": {"level": "info", "data": "Waiting"}}), flush=True)
|
||||
continue
|
||||
if request["method"] == "initialize":
|
||||
result = {"protocolVersion": request["params"]["protocolVersion"], "capabilities": {"tools": {}, "logging": {}}, "serverInfo": {"name": "diagnostic", "version": "1"}}
|
||||
elif request["method"] == "tools/list":
|
||||
print(json.dumps({"jsonrpc": "2.0", "method": "notifications/message", "params": {"level": "info", "data": "Listing tools"}}), flush=True)
|
||||
print(json.dumps({"jsonrpc": "2.0", "id": "unmatched", "result": {}}), flush=True)
|
||||
print(json.dumps({"jsonrpc": "2.0", "id": "server-ping", "method": "ping"}), flush=True)
|
||||
result = {"tools": [{"name": "ping", "inputSchema": {"type": "object"}}]}
|
||||
else:
|
||||
result = {"content": [{"type": "text", "text": "pong"}], "isError": False}
|
||||
print(json.dumps({"jsonrpc": "2.0", "id": request["id"], "result": result}), flush=True)
|
||||
"""
|
||||
|
||||
|
||||
def _diagnostic_transport(transport: MCPTransport, mode: str, failure_method: str) -> _TransportContext:
|
||||
from mcp import StdioServerParameters
|
||||
from mcp.client.sse import sse_client
|
||||
from mcp.client.stdio import stdio_client
|
||||
|
||||
if transport == MCPTransport.stdio:
|
||||
return stdio_client(
|
||||
StdioServerParameters(
|
||||
command=sys.executable, args=["-u", "-c", _DIAGNOSTIC_STDIO_SERVER, mode, failure_method]
|
||||
)
|
||||
)
|
||||
messages: Final[asyncio.Queue[bytes | Exception | None]] = asyncio.Queue()
|
||||
|
||||
async def respond(request: httpx.Request) -> httpx.Response:
|
||||
if request.method == "GET":
|
||||
return httpx.Response(
|
||||
200, headers={"Content-Type": "text/event-stream"}, stream=_DiagnosticSSEStream(messages)
|
||||
)
|
||||
payload: Final = json.loads(request.content)
|
||||
if "method" not in payload or "id" not in payload:
|
||||
return httpx.Response(202)
|
||||
if payload["method"] == failure_method and mode != "ok":
|
||||
if mode == "bad-json":
|
||||
await messages.put(b"secret-invalid-json")
|
||||
elif mode == "io-error":
|
||||
await messages.put(httpx.ReadError("secret-read-error"))
|
||||
elif mode == "closed":
|
||||
await messages.put(None)
|
||||
elif mode == "silent":
|
||||
await messages.put(
|
||||
b'{"jsonrpc":"2.0","method":"notifications/message","params":{"level":"info","data":"Waiting"}}'
|
||||
)
|
||||
return httpx.Response(202)
|
||||
if payload["method"] == "tools/list":
|
||||
for message in (
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"method": "notifications/message",
|
||||
"params": {"level": "info", "data": "Listing tools"},
|
||||
},
|
||||
{"jsonrpc": "2.0", "id": "unmatched", "result": {}},
|
||||
{"jsonrpc": "2.0", "id": "server-ping", "method": "ping"},
|
||||
):
|
||||
await messages.put(json.dumps(message).encode())
|
||||
result: Final = (
|
||||
{
|
||||
"protocolVersion": LATEST_PROTOCOL_VERSION,
|
||||
"capabilities": {"tools": {}, "logging": {}},
|
||||
"serverInfo": {"name": "diagnostic", "version": "1"},
|
||||
}
|
||||
if payload["method"] == "initialize"
|
||||
else {"tools": [{"name": "ping", "inputSchema": {"type": "object"}}]}
|
||||
if payload["method"] == "tools/list"
|
||||
else {"content": [{"type": "text", "text": "pong"}], "isError": False}
|
||||
)
|
||||
await messages.put(json.dumps({"jsonrpc": "2.0", "id": payload["id"], "result": result}).encode())
|
||||
return httpx.Response(202)
|
||||
|
||||
def factory(
|
||||
headers: dict[str, str] | None = None,
|
||||
timeout: httpx.Timeout | None = None,
|
||||
auth: httpx.Auth | None = None,
|
||||
) -> httpx.AsyncClient:
|
||||
return httpx.AsyncClient(transport=httpx.MockTransport(respond), headers=headers, timeout=timeout, auth=auth)
|
||||
|
||||
return sse_client("https://example.com/sse", httpx_client_factory=factory)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio])
|
||||
@pytest.mark.parametrize("failure_method", ["initialize", "tools/list"])
|
||||
async def test_transport_parsing_failure_is_preserved(transport: MCPTransport, failure_method: str) -> None:
|
||||
client: Final = MCPClient(server_url="https://example.com/sse", transport_type=transport, timeout=0.2)
|
||||
with pytest.raises(ValidationError):
|
||||
await asyncio.wait_for(
|
||||
client._execute_session_operation(
|
||||
_diagnostic_transport(transport, "bad-json", failure_method), lambda session: session.list_tools()
|
||||
),
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sse_read_failure_is_preserved() -> None:
|
||||
client: Final = MCPClient(server_url="https://example.com/sse", transport_type=MCPTransport.sse, timeout=0.2)
|
||||
with pytest.raises(httpx.ReadError, match="secret-read-error"):
|
||||
await asyncio.wait_for(
|
||||
client._execute_session_operation(
|
||||
_diagnostic_transport(MCPTransport.sse, "io-error", "tools/list"), lambda session: session.list_tools()
|
||||
),
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio])
|
||||
@pytest.mark.parametrize("mode", ["ok", "closed", "silent"])
|
||||
async def test_transport_completion_and_normal_messages(transport: MCPTransport, mode: str) -> None:
|
||||
from mcp import ClientSession
|
||||
from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message
|
||||
|
||||
logging_callback: Final = AsyncMock()
|
||||
client: Final = MCPClient(
|
||||
server_url="https://example.com/sse", transport_type=transport, timeout=0.2, logging_callback=logging_callback
|
||||
)
|
||||
|
||||
async def operation(session: ClientSession) -> CallToolResult:
|
||||
tools: Final = await session.list_tools()
|
||||
assert [tool.name for tool in tools.tools] == ["ping"]
|
||||
return await session.call_tool("ping", {})
|
||||
|
||||
pending: Final = client._execute_session_operation(_diagnostic_transport(transport, mode, "tools/list"), operation)
|
||||
if mode == "ok":
|
||||
result: Final = await asyncio.wait_for(pending, timeout=3)
|
||||
assert result.isError is False
|
||||
assert result.content[0].text == "pong"
|
||||
logging_callback.assert_awaited_once_with(LoggingMessageNotificationParams(level="info", data="Listing tools"))
|
||||
else:
|
||||
with pytest.raises(McpError) as caught:
|
||||
await asyncio.wait_for(pending, timeout=3)
|
||||
if mode == "closed":
|
||||
assert "connection was closed" in _connection_error_message(caught.value, client.server_url, 0.2)
|
||||
else:
|
||||
assert isinstance(as_mcp_read_timeout(caught.value), TimeoutError)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("transport", [MCPTransport.sse, MCPTransport.stdio])
|
||||
async def test_transport_cancellation_cleans_up_a_pending_request(transport: MCPTransport) -> None:
|
||||
ready: Final = asyncio.Event()
|
||||
|
||||
async def on_log(message: LoggingMessageNotificationParams) -> None:
|
||||
if message.data == "Waiting":
|
||||
ready.set()
|
||||
|
||||
client: Final = MCPClient(
|
||||
server_url="https://example.com/sse", transport_type=transport, timeout=30, logging_callback=on_log
|
||||
)
|
||||
task: Final = asyncio.create_task(
|
||||
client._execute_session_operation(
|
||||
_diagnostic_transport(transport, "silent", "tools/list"), lambda session: session.list_tools()
|
||||
)
|
||||
)
|
||||
try:
|
||||
await asyncio.wait_for(ready.wait(), timeout=3)
|
||||
finally:
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await asyncio.wait_for(task, timeout=3)
|
||||
|
||||
|
||||
class _InterruptedHTTPBody(httpx.AsyncByteStream):
|
||||
async def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
yield b'{"jsonrpc":'
|
||||
raise httpx.RemoteProtocolError("secret-incomplete-response")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_interrupted_http_response_preserves_the_transport_failure() -> None:
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, headers={"Content-Type": "application/json"}, stream=_InterruptedHTTPBody())
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
|
||||
client: Final = MCPClient(server_url="https://example.com/mcp", timeout=30)
|
||||
with pytest.raises(httpx.RemoteProtocolError, match="secret-incomplete-response"):
|
||||
await asyncio.wait_for(
|
||||
client._execute_session_operation(
|
||||
streamable_http_client(client.server_url, http_client=http_client),
|
||||
lambda session: session.list_tools(),
|
||||
),
|
||||
timeout=3,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_http_event_stream_uses_the_existing_request_deadline() -> None:
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, headers={"Content-Type": "text/event-stream"}, content=b"")
|
||||
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as http_client:
|
||||
client: Final = MCPClient(server_url="https://example.com/mcp", timeout=0.2)
|
||||
with pytest.raises(McpError) as caught:
|
||||
await asyncio.wait_for(
|
||||
client._execute_session_operation(
|
||||
streamable_http_client(client.server_url, http_client=http_client),
|
||||
lambda session: session.list_tools(),
|
||||
),
|
||||
timeout=3,
|
||||
)
|
||||
assert isinstance(as_mcp_read_timeout(caught.value), TimeoutError)
|
||||
|
|
|
|||
|
|
@ -1842,12 +1842,14 @@ class _ApplyStyleGuardrail(CustomGuardrail):
|
|||
self.block = block
|
||||
self.apply_called = False
|
||||
self.seen_texts = None
|
||||
self.seen_request_data = None
|
||||
|
||||
async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None):
|
||||
from fastapi import HTTPException
|
||||
|
||||
self.apply_called = True
|
||||
self.seen_texts = inputs.get("texts")
|
||||
self.seen_request_data = request_data
|
||||
if self.block:
|
||||
raise HTTPException(status_code=400, detail={"error": "Violated moderation policy"})
|
||||
return inputs
|
||||
|
|
@ -2646,6 +2648,91 @@ class TestCustomGuardrailPostCallSuccessDeploymentHook:
|
|||
which starved every later callback in litellm.callbacks (notably the lazily-appended
|
||||
VectorStorePreCallHook that attaches provider_specific_fields["search_results"])."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_retains_request_identity(self) -> None:
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
guardrail: Final = _ApplyStyleGuardrail(block=False)
|
||||
guardrail.event_hook = GuardrailEventHooks.post_call
|
||||
request_data: Final = {"guardrails": ["apply-style-guardrail"]}
|
||||
response: Final = ModelResponse(choices=[Choices(message=Message(content="review me"))])
|
||||
|
||||
await guardrail.async_post_call_success_deployment_hook(
|
||||
request_data=request_data, response=response, call_type=CallTypes.acompletion
|
||||
)
|
||||
|
||||
assert guardrail.seen_request_data is request_data
|
||||
assert guardrail.seen_texts == ["review me"]
|
||||
assert "guardrail_to_apply" not in request_data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("call_type", (None, CallTypes.acompletion))
|
||||
async def test_apply_guardrail_masks_response_and_records_metadata(self, call_type: CallTypes | None) -> None:
|
||||
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
|
||||
ContentFilterGuardrail,
|
||||
)
|
||||
from litellm.types.guardrails import BlockedWord, ContentFilterAction, GuardrailEventHooks
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
guardrail: Final = ContentFilterGuardrail(
|
||||
guardrail_name="response-filter",
|
||||
event_hook=GuardrailEventHooks.post_call,
|
||||
blocked_words=[BlockedWord(keyword="secret", action=ContentFilterAction.MASK)],
|
||||
)
|
||||
request_data: Final = {"guardrails": ["response-filter"]}
|
||||
response: Final = ModelResponse(choices=[Choices(message=Message(content="a secret"))])
|
||||
|
||||
result: Final = await guardrail.async_post_call_success_deployment_hook(
|
||||
request_data=request_data, response=response, call_type=call_type
|
||||
)
|
||||
|
||||
assert isinstance(result, ModelResponse)
|
||||
assert result.choices[0].message.content == f"a {guardrail.keyword_redaction_tag}"
|
||||
entries: Final = _guardrail_entries(request_data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_name"] == "response-filter"
|
||||
assert entries[0]["guardrail_mode"] == "post_call"
|
||||
assert "guardrail_to_apply" not in request_data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("error_type", (None, RuntimeError, asyncio.CancelledError))
|
||||
async def test_dispatch_cleans_up_request_on_every_exit(self, error_type: type[BaseException] | None) -> None:
|
||||
from contextlib import nullcontext
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.utils import LLMResponseTypes, ModelResponse
|
||||
|
||||
error: Final = error_type("dispatch interrupted") if error_type is not None else None
|
||||
|
||||
class Dispatch(CustomLogger):
|
||||
request_data: dict[str, object] | None = None
|
||||
|
||||
async def async_post_call_success_hook(
|
||||
self, data: dict[str, object], user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes
|
||||
) -> LLMResponseTypes:
|
||||
self.request_data = data
|
||||
if error is not None:
|
||||
raise error
|
||||
return response
|
||||
|
||||
dispatch: Final = Dispatch()
|
||||
|
||||
class Guardrail(_ApplyStyleGuardrail):
|
||||
def _deployment_hook_target(self) -> CustomLogger:
|
||||
return dispatch
|
||||
|
||||
guardrail: Final = Guardrail(block=False)
|
||||
guardrail.event_hook = GuardrailEventHooks.post_call
|
||||
request_data: Final = {"guardrails": ["apply-style-guardrail"]}
|
||||
with pytest.raises(error_type) if error_type is not None else nullcontext():
|
||||
await guardrail.async_post_call_success_deployment_hook(
|
||||
request_data=request_data, response=ModelResponse(), call_type=CallTypes.acompletion
|
||||
)
|
||||
|
||||
assert dispatch.request_data is request_data
|
||||
assert "guardrail_to_apply" not in request_data
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_none_when_request_has_no_guardrails(self):
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
|
@ -2740,4 +2827,5 @@ class TestCustomGuardrailPostCallSuccessDeploymentHook:
|
|||
|
||||
assert result is response
|
||||
assert response.choices[0].message.content == "filtered response"
|
||||
assert request_data == {"guardrails": ["test-guardrail"]}
|
||||
assert "guardrail_to_apply" not in request_data
|
||||
assert len(_guardrail_entries(request_data)) == 1
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
import json
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import pinned_billing_time
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
|
||||
StandardBuiltInToolCostTracking,
|
||||
)
|
||||
|
|
@ -27,10 +29,10 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
||||
BilledTokenRates,
|
||||
CostCalculatorUtils,
|
||||
PromptTokensDetailsResult,
|
||||
TokenRates,
|
||||
TokenTypeCostBreakdown,
|
||||
_calculate_input_cost,
|
||||
_get_token_base_cost,
|
||||
_is_off_peak,
|
||||
|
|
@ -38,6 +40,7 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import (
|
|||
apply_off_peak_pricing,
|
||||
calculate_cache_writing_cost,
|
||||
generic_cost_per_token,
|
||||
get_billed_token_rates,
|
||||
get_token_type_cost_breakdown,
|
||||
)
|
||||
from litellm.types.utils import CacheCreationTokenDetails, Usage
|
||||
|
|
@ -3906,6 +3909,200 @@ def test_token_type_cost_breakdown_reconciles_with_generic_total(_local_model_co
|
|||
assert text_input_cost + breakdown.cache_read_cost == pytest.approx(prompt_cost)
|
||||
|
||||
|
||||
def _custom_priced_usage() -> Usage:
|
||||
return Usage(
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=500,
|
||||
total_tokens=1500,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=800, cache_creation_tokens=100),
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=200),
|
||||
)
|
||||
|
||||
|
||||
def test_token_type_cost_breakdown_prices_custom_pricing_from_its_flat_rates():
|
||||
"""
|
||||
A custom-priced deployment, usually absent from the cost map, used to get zero cache and
|
||||
reasoning lines while its total already billed cache tokens at the custom cache rates.
|
||||
The lines must come from the same flat rates: a configured cache rate, else the input
|
||||
rate for cache tokens and the output rate for reasoning tokens.
|
||||
"""
|
||||
from litellm.types.utils import CostPerToken
|
||||
|
||||
breakdown = get_token_type_cost_breakdown(
|
||||
model="openai/onprem-model",
|
||||
custom_llm_provider="openai",
|
||||
usage=_custom_priced_usage(),
|
||||
custom_cost_per_token=CostPerToken(
|
||||
input_cost_per_token=1e-6, output_cost_per_token=2e-6, cache_read_input_token_cost=1e-7
|
||||
),
|
||||
)
|
||||
|
||||
assert breakdown.cache_read_cost == pytest.approx(800 * 1e-7)
|
||||
assert breakdown.cache_creation_cost == pytest.approx(100 * 1e-6)
|
||||
assert breakdown.reasoning_cost == pytest.approx(200 * 2e-6)
|
||||
|
||||
|
||||
def test_token_type_cost_breakdown_reconciles_with_custom_pricing_totals():
|
||||
from litellm.cost_calculator import cost_per_token
|
||||
from litellm.types.utils import CostPerToken
|
||||
|
||||
usage = _custom_priced_usage()
|
||||
custom_cost_per_token = CostPerToken(
|
||||
input_cost_per_token=1e-6,
|
||||
output_cost_per_token=2e-6,
|
||||
cache_read_input_token_cost=1e-7,
|
||||
cache_creation_input_token_cost=1.25e-6,
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
model="openai/onprem-model",
|
||||
custom_llm_provider="openai",
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=500,
|
||||
usage_object=usage,
|
||||
custom_cost_per_token=custom_cost_per_token,
|
||||
)
|
||||
breakdown = get_token_type_cost_breakdown(
|
||||
model="openai/onprem-model",
|
||||
custom_llm_provider="openai",
|
||||
usage=usage,
|
||||
custom_cost_per_token=custom_cost_per_token,
|
||||
)
|
||||
|
||||
assert 100 * 1e-6 + breakdown.cache_read_cost + breakdown.cache_creation_cost == pytest.approx(prompt_cost)
|
||||
assert 300 * 2e-6 + breakdown.reasoning_cost == pytest.approx(completion_cost)
|
||||
|
||||
|
||||
def test_billed_token_rates_follow_the_token_tier_the_breakdown_bills_at(monkeypatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"tiered-cache-model",
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"cache_read_input_token_cost": 3e-7,
|
||||
"cache_creation_input_token_cost": 3.75e-6,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-6,
|
||||
"output_cost_per_token_above_200k_tokens": 3e-5,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6e-7,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-6,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
usage = Usage(
|
||||
prompt_tokens=250_000,
|
||||
completion_tokens=1_000,
|
||||
total_tokens=251_000,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=200_000, cache_creation_tokens=10_000),
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=200),
|
||||
)
|
||||
|
||||
rates = get_billed_token_rates(model="tiered-cache-model", custom_llm_provider="openai", usage=usage)
|
||||
breakdown = get_token_type_cost_breakdown(model="tiered-cache-model", custom_llm_provider="openai", usage=usage)
|
||||
|
||||
assert rates == BilledTokenRates(
|
||||
input_cost_per_token=6e-6,
|
||||
output_cost_per_token=3e-5,
|
||||
cache_read_input_token_cost=6e-7,
|
||||
cache_creation_input_token_cost=7.5e-6,
|
||||
cache_creation_input_token_cost_above_1hr=0.0,
|
||||
output_cost_per_reasoning_token=3e-5,
|
||||
)
|
||||
assert breakdown.cache_read_cost == pytest.approx(200_000 * rates.cache_read_input_token_cost)
|
||||
assert breakdown.cache_creation_cost == pytest.approx(10_000 * rates.cache_creation_input_token_cost)
|
||||
assert breakdown.reasoning_cost == pytest.approx(200 * rates.output_cost_per_reasoning_token)
|
||||
|
||||
|
||||
def test_a_pinned_billing_time_prices_the_totals_and_the_reported_rates_at_one_moment(monkeypatch):
|
||||
"""Totals and reported rates resolve off-peak pricing on separate paths that each read the
|
||||
clock, so a window opening between the two reads used to leave them describing one request
|
||||
at two different prices. Pinned, both must answer for the pinned moment."""
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"off-peak-model",
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"off_peak_pricing": {
|
||||
"hours_utc": "02:00-03:00",
|
||||
"input_cost_per_token": 1e-6,
|
||||
"output_cost_per_token": 5e-6,
|
||||
},
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
|
||||
|
||||
with pinned_billing_time(datetime(2026, 1, 1, 2, 30, tzinfo=timezone.utc)):
|
||||
off_peak_prompt_cost, off_peak_completion_cost = generic_cost_per_token(
|
||||
model="off-peak-model", usage=usage, custom_llm_provider="openai"
|
||||
)
|
||||
off_peak_rates = get_billed_token_rates(model="off-peak-model", custom_llm_provider="openai", usage=usage)
|
||||
with pinned_billing_time(datetime(2026, 1, 1, 12, 30, tzinfo=timezone.utc)):
|
||||
peak_prompt_cost, peak_completion_cost = generic_cost_per_token(
|
||||
model="off-peak-model", usage=usage, custom_llm_provider="openai"
|
||||
)
|
||||
peak_rates = get_billed_token_rates(model="off-peak-model", custom_llm_provider="openai", usage=usage)
|
||||
|
||||
assert off_peak_rates.input_cost_per_token == pytest.approx(1e-6)
|
||||
assert peak_rates.input_cost_per_token == pytest.approx(3e-6)
|
||||
assert off_peak_prompt_cost == pytest.approx(1000 * off_peak_rates.input_cost_per_token)
|
||||
assert off_peak_completion_cost == pytest.approx(500 * off_peak_rates.output_cost_per_token)
|
||||
assert peak_prompt_cost == pytest.approx(1000 * peak_rates.input_cost_per_token)
|
||||
assert peak_completion_cost == pytest.approx(500 * peak_rates.output_cost_per_token)
|
||||
|
||||
|
||||
def test_the_token_type_breakdown_carries_the_rates_it_billed_at(monkeypatch):
|
||||
"""Callers that report both the lines and the rates read the rates off the breakdown rather than
|
||||
resolving them a second time, so the breakdown has to hand back exactly what it billed at."""
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"xai/tiered-model",
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"cache_read_input_token_cost": 3e-7,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-6,
|
||||
"output_cost_per_token_above_200k_tokens": 3e-5,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6e-7,
|
||||
"litellm_provider": "xai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
usage = Usage(
|
||||
prompt_tokens=200_000,
|
||||
completion_tokens=1_000,
|
||||
total_tokens=201_000,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=100_000),
|
||||
)
|
||||
|
||||
breakdown = get_token_type_cost_breakdown(model="xai/tiered-model", custom_llm_provider="xai", usage=usage)
|
||||
|
||||
assert breakdown.rates == get_billed_token_rates(
|
||||
model="xai/tiered-model", custom_llm_provider="xai", usage=usage
|
||||
)
|
||||
assert breakdown.rates.cache_read_input_token_cost == pytest.approx(6e-7)
|
||||
assert breakdown.cache_read_cost == pytest.approx(100_000 * breakdown.rates.cache_read_input_token_cost)
|
||||
|
||||
|
||||
def test_the_token_type_breakdown_reports_no_rates_for_an_unpriced_model():
|
||||
usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
|
||||
|
||||
breakdown = get_token_type_cost_breakdown(
|
||||
model="no-such-model-anywhere", custom_llm_provider="openai", usage=usage
|
||||
)
|
||||
|
||||
assert breakdown.rates is None
|
||||
|
||||
|
||||
def test_billed_token_rates_are_none_for_an_unpriced_model():
|
||||
usage = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
|
||||
|
||||
assert get_billed_token_rates(model="no-such-model-anywhere", custom_llm_provider="openai", usage=usage) is None
|
||||
|
||||
|
||||
def test_token_type_cost_breakdown_zero_without_special_tokens(_local_model_cost_map):
|
||||
|
||||
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
|
||||
|
|
@ -3913,9 +4110,7 @@ def test_token_type_cost_breakdown_zero_without_special_tokens(_local_model_cost
|
|||
model="gpt-4o", custom_llm_provider="openai", usage=usage
|
||||
)
|
||||
|
||||
assert breakdown == TokenTypeCostBreakdown(
|
||||
reasoning_cost=0.0, cache_read_cost=0.0, cache_creation_cost=0.0
|
||||
)
|
||||
assert (breakdown.reasoning_cost, breakdown.cache_read_cost, breakdown.cache_creation_cost) == (0.0, 0.0, 0.0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -3987,9 +4182,7 @@ def test_token_type_cost_breakdown_handles_unknown_model_gracefully():
|
|||
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=5),
|
||||
),
|
||||
)
|
||||
assert breakdown == TokenTypeCostBreakdown(
|
||||
reasoning_cost=0.0, cache_read_cost=0.0, cache_creation_cost=0.0
|
||||
)
|
||||
assert (breakdown.reasoning_cost, breakdown.cache_read_cost, breakdown.cache_creation_cost) == (0.0, 0.0, 0.0)
|
||||
|
||||
|
||||
def test_token_type_cost_breakdown_applies_regional_uplift(_local_model_cost_map):
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ count actual model entries, not reserved meta keys) and the extraction of the
|
|||
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -26,9 +27,7 @@ from litellm.litellm_core_utils.get_model_cost_map import (
|
|||
|
||||
|
||||
def _load_root_cost_map() -> dict:
|
||||
path = os.path.join(
|
||||
os.path.dirname(__file__), "../../../model_prices_and_context_window.json"
|
||||
)
|
||||
path = os.path.join(os.path.dirname(__file__), "../../../model_prices_and_context_window.json")
|
||||
with open(path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
|
@ -44,9 +43,7 @@ def test_git_blob_id_is_what_git_hash_object_prints():
|
|||
|
||||
|
||||
def _make_models(n: int) -> dict:
|
||||
return {
|
||||
f"model-{i}": {"litellm_provider": "openai", "mode": "chat"} for i in range(n)
|
||||
}
|
||||
return {f"model-{i}": {"litellm_provider": "openai", "mode": "chat"} for i in range(n)}
|
||||
|
||||
|
||||
def test_count_model_entries_excludes_reserved_keys():
|
||||
|
|
@ -129,9 +126,7 @@ def test_finalize_pops_key_and_installs_rules():
|
|||
def test_finalize_with_no_block_clears_rules():
|
||||
previous = list(get_fallback_generalization_rules())
|
||||
try:
|
||||
set_fallback_generalizations(
|
||||
[{"name": "stale", "pattern": r"^x", "model_info": {"a": 1}}]
|
||||
)
|
||||
set_fallback_generalizations([{"name": "stale", "pattern": r"^x", "model_info": {"a": 1}}])
|
||||
_finalize_model_cost_map(_make_models(2))
|
||||
assert match_capability_generalizations("x-1") is None
|
||||
finally:
|
||||
|
|
@ -317,9 +312,7 @@ def test_get_model_cost_map_stamps_loaded_at():
|
|||
|
||||
from litellm.litellm_core_utils import get_model_cost_map as module
|
||||
|
||||
client, _calls = _mock_client(
|
||||
[httpx.Response(200, content=_real_map_bytes())], client_cls=httpx.Client
|
||||
)
|
||||
client, _calls = _mock_client([httpx.Response(200, content=_real_map_bytes())], client_cls=httpx.Client)
|
||||
|
||||
before = datetime.now(timezone.utc)
|
||||
module.get_model_cost_map(url="https://example.invalid/cost_map.json", client=client)
|
||||
|
|
@ -328,6 +321,7 @@ def test_get_model_cost_map_stamps_loaded_at():
|
|||
assert loaded_at is not None
|
||||
assert before <= loaded_at <= datetime.now(timezone.utc)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# refetch_model_cost_map: retry/backoff behavior for runtime reloads
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -394,9 +388,7 @@ async def test_refetch_retries_429_honoring_retry_after():
|
|||
]
|
||||
)
|
||||
sleeper = _SleepRecorder()
|
||||
result = await refetch_model_cost_map(
|
||||
url=_URL, sleep=sleeper, rng=random.Random(0), client=client
|
||||
)
|
||||
result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
|
||||
assert isinstance(result, ModelCostMapReloaded)
|
||||
assert len(result.model_cost_map) > 100
|
||||
assert calls["count"] == 3
|
||||
|
|
@ -408,9 +400,7 @@ async def test_refetch_gives_up_after_max_attempts_with_exponential_backoff():
|
|||
"""All 429 without Retry-After: exponential backoff waits, then a failure value."""
|
||||
client, calls = _mock_client([httpx.Response(429)])
|
||||
sleeper = _SleepRecorder()
|
||||
result = await refetch_model_cost_map(
|
||||
url=_URL, sleep=sleeper, rng=random.Random(0), client=client
|
||||
)
|
||||
result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
|
||||
assert isinstance(result, ModelCostMapReloadUnavailable)
|
||||
assert "429" in result.reason
|
||||
assert "after 3 attempts" in result.reason
|
||||
|
|
@ -430,9 +420,7 @@ async def test_refetch_caps_retry_after_wait():
|
|||
]
|
||||
)
|
||||
sleeper = _SleepRecorder()
|
||||
result = await refetch_model_cost_map(
|
||||
url=_URL, sleep=sleeper, rng=random.Random(0), client=client
|
||||
)
|
||||
result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
|
||||
assert isinstance(result, ModelCostMapReloaded)
|
||||
assert sleeper.waits == [30.0]
|
||||
|
||||
|
|
@ -447,9 +435,7 @@ async def test_refetch_retries_transport_errors():
|
|||
]
|
||||
)
|
||||
sleeper = _SleepRecorder()
|
||||
result = await refetch_model_cost_map(
|
||||
url=_URL, sleep=sleeper, rng=random.Random(0), client=client
|
||||
)
|
||||
result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
|
||||
assert isinstance(result, ModelCostMapReloaded)
|
||||
assert calls["count"] == 2
|
||||
assert len(sleeper.waits) == 1
|
||||
|
|
@ -460,9 +446,7 @@ async def test_refetch_non_retryable_status_fails_immediately():
|
|||
"""A 404 is permanent: one attempt, no sleeps, failure value."""
|
||||
client, calls = _mock_client([httpx.Response(404)])
|
||||
sleeper = _SleepRecorder()
|
||||
result = await refetch_model_cost_map(
|
||||
url=_URL, sleep=sleeper, rng=random.Random(0), client=client
|
||||
)
|
||||
result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
|
||||
assert isinstance(result, ModelCostMapReloadUnavailable)
|
||||
assert "404" in result.reason
|
||||
assert calls["count"] == 1
|
||||
|
|
@ -473,9 +457,7 @@ async def test_refetch_non_retryable_status_fails_immediately():
|
|||
async def test_refetch_invalid_json_fails_immediately():
|
||||
client, calls = _mock_client([httpx.Response(200, content=b"not json")])
|
||||
sleeper = _SleepRecorder()
|
||||
result = await refetch_model_cost_map(
|
||||
url=_URL, sleep=sleeper, rng=random.Random(0), client=client
|
||||
)
|
||||
result = await refetch_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
|
||||
assert isinstance(result, ModelCostMapReloadUnavailable)
|
||||
assert "invalid JSON" in result.reason
|
||||
assert calls["count"] == 1
|
||||
|
|
@ -487,9 +469,7 @@ async def test_refetch_shrunk_map_fails_integrity_not_swapped_in():
|
|||
"""A drastically shrunk upstream file is rejected instead of being adopted."""
|
||||
tiny = json.dumps(_make_models(60)).encode()
|
||||
client, _calls = _mock_client([httpx.Response(200, content=tiny)])
|
||||
result = await refetch_model_cost_map(
|
||||
url=_URL, sleep=_SleepRecorder(), rng=random.Random(0), client=client
|
||||
)
|
||||
result = await refetch_model_cost_map(url=_URL, sleep=_SleepRecorder(), rng=random.Random(0), client=client)
|
||||
assert isinstance(result, ModelCostMapReloadUnavailable)
|
||||
assert "integrity validation" in result.reason
|
||||
|
||||
|
|
@ -586,69 +566,125 @@ from litellm.litellm_core_utils.get_model_cost_map import (
|
|||
class _SyncSleepRecorder:
|
||||
"""Injected in place of time.sleep so the boot path's waits are asserted without delay."""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, block=False):
|
||||
self.waits = []
|
||||
self.block = block
|
||||
self.started = threading.Event()
|
||||
self.release = threading.Event()
|
||||
|
||||
def __call__(self, seconds: float) -> None:
|
||||
if self.block:
|
||||
self.started.set()
|
||||
self.release.wait(timeout=10)
|
||||
self.waits.append(seconds)
|
||||
|
||||
|
||||
def test_boot_load_retries_transient_failures_instead_of_falling_back():
|
||||
"""A refused connection then a 503 at pod boot used to pin the process to the bundled
|
||||
backup for its lifetime; both are transient and must be retried before giving up."""
|
||||
def _retry_threads():
|
||||
return [thread for thread in threading.enumerate() if thread.name == "litellm-model-cost-map-retry"]
|
||||
|
||||
|
||||
def test_boot_load_success_does_not_start_background_retry():
|
||||
client, calls = _mock_client([httpx.Response(200, content=_real_map_bytes())], client_cls=httpx.Client)
|
||||
sleeper = _SyncSleepRecorder()
|
||||
|
||||
cost_map = get_model_cost_map(
|
||||
url=_URL,
|
||||
sleep=sleeper,
|
||||
rng=random.Random(0),
|
||||
client=client,
|
||||
)
|
||||
assert calls["count"] == 1
|
||||
assert sleeper.waits == []
|
||||
assert _retry_threads() == []
|
||||
assert cost_map.keys() >= _load_root_cost_map().keys() - {"sample_spec", FALLBACK_GENERALIZATIONS_KEY}
|
||||
assert get_model_cost_map_source_info()["source"] == "remote"
|
||||
|
||||
|
||||
def test_boot_load_transient_failure_returns_local_then_background_retry_adopts_remote(monkeypatch):
|
||||
import litellm
|
||||
from litellm import utils as litellm_utils
|
||||
from litellm.litellm_core_utils import get_model_cost_map as module
|
||||
|
||||
original_model_cost = litellm.model_cost
|
||||
monkeypatch.setattr(litellm, "model_cost", dict(original_model_cost))
|
||||
for name, provider_models in tuple(vars(litellm).items()):
|
||||
if name.endswith("_models") and isinstance(provider_models, set):
|
||||
monkeypatch.setattr(litellm, name, set(provider_models))
|
||||
monkeypatch.setattr(litellm, "models_by_provider", dict(litellm.models_by_provider))
|
||||
monkeypatch.setattr(
|
||||
litellm_utils,
|
||||
"_runtime_registered_model_cost",
|
||||
dict(litellm_utils._runtime_registered_model_cost),
|
||||
)
|
||||
source_info = module._cost_map_source_info
|
||||
for name in ("source", "url", "is_env_forced", "fallback_reason", "loaded_at", "source_revision", "etag"):
|
||||
monkeypatch.setattr(source_info, name, getattr(source_info, name))
|
||||
|
||||
remote_map = _load_root_cost_map()
|
||||
remote_map["claude-remote-only-test"] = {"litellm_provider": "anthropic", "mode": "chat"}
|
||||
client, calls = _mock_client(
|
||||
[
|
||||
httpx.ConnectError("connection refused"),
|
||||
httpx.Response(503),
|
||||
httpx.Response(200, content=_real_map_bytes()),
|
||||
httpx.Response(200, content=json.dumps(remote_map).encode()),
|
||||
],
|
||||
client_cls=httpx.Client,
|
||||
)
|
||||
sleeper = _SyncSleepRecorder()
|
||||
sleeper = _SyncSleepRecorder(block=True)
|
||||
litellm.register_model({"my-runtime-model": {"litellm_provider": "custom", "max_input_tokens": 4321}})
|
||||
|
||||
cost_map = get_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
|
||||
|
||||
assert calls["count"] == 3
|
||||
assert len(sleeper.waits) == 2
|
||||
assert 2.0 <= sleeper.waits[0] < 3.0
|
||||
assert 4.0 <= sleeper.waits[1] < 5.0
|
||||
source = get_model_cost_map_source_info()
|
||||
assert source["source"] == "remote"
|
||||
assert source["fallback_reason"] is None
|
||||
assert cost_map.keys() >= _load_root_cost_map().keys() - {"sample_spec", FALLBACK_GENERALIZATIONS_KEY}
|
||||
|
||||
|
||||
def test_boot_load_honors_retry_after_then_falls_back_after_max_attempts():
|
||||
"""An outage longer than the retry budget still ends on the bundled backup, and the
|
||||
recorded fallback reason says how many attempts were spent so operators can tell."""
|
||||
client, calls = _mock_client(
|
||||
[httpx.Response(429, headers={"Retry-After": "7"})], client_cls=httpx.Client
|
||||
cost_map = get_model_cost_map(
|
||||
url=_URL,
|
||||
max_attempts=3,
|
||||
sleep=sleeper,
|
||||
rng=random.Random(0),
|
||||
client=client,
|
||||
)
|
||||
sleeper = _SyncSleepRecorder()
|
||||
|
||||
cost_map = get_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
|
||||
|
||||
assert calls["count"] == 3
|
||||
assert sleeper.waits == [7.0, 7.0]
|
||||
source = get_model_cost_map_source_info()
|
||||
assert source["source"] == "local"
|
||||
assert "after 3 attempts" in source["fallback_reason"]
|
||||
assert len(cost_map) > 100
|
||||
assert calls["count"] == 1
|
||||
assert sleeper.waits == []
|
||||
assert sleeper.started.wait(timeout=10)
|
||||
threads = _retry_threads()
|
||||
try:
|
||||
assert len(threads) == 1
|
||||
assert "claude-remote-only-test" not in cost_map
|
||||
assert cost_map.keys() == _finalize_model_cost_map(GetModelCostMap.load_local_model_cost_map()).keys()
|
||||
source = get_model_cost_map_source_info()
|
||||
assert source["source"] == "local"
|
||||
assert source["fallback_reason"].startswith("Remote fetch failed:")
|
||||
sleeper.release.set()
|
||||
for thread in threads:
|
||||
thread.join(timeout=10)
|
||||
assert all(not thread.is_alive() for thread in threads)
|
||||
assert sleeper.waits and 2.0 <= sleeper.waits[0] < 3.0
|
||||
assert calls["count"] == 2
|
||||
assert "claude-remote-only-test" in litellm.model_cost
|
||||
assert "claude-remote-only-test" in litellm.anthropic_models
|
||||
assert "my-runtime-model" in litellm.model_cost
|
||||
source = get_model_cost_map_source_info()
|
||||
assert source["source"] == "remote"
|
||||
assert source["fallback_reason"] is None
|
||||
finally:
|
||||
sleeper.release.set()
|
||||
for thread in _retry_threads():
|
||||
thread.join(timeout=10)
|
||||
|
||||
|
||||
def test_boot_load_does_not_retry_permanent_failures():
|
||||
"""A 404 or a malformed URL cannot heal by waiting: one attempt, no sleeps, backup."""
|
||||
def test_boot_load_does_not_retry_non_retryable_failure():
|
||||
client, calls = _mock_client([httpx.Response(404)], client_cls=httpx.Client)
|
||||
sleeper = _SyncSleepRecorder()
|
||||
|
||||
get_model_cost_map(url=_URL, sleep=sleeper, rng=random.Random(0), client=client)
|
||||
get_model_cost_map(
|
||||
url=_URL,
|
||||
sleep=sleeper,
|
||||
rng=random.Random(0),
|
||||
client=client,
|
||||
)
|
||||
assert calls["count"] == 1
|
||||
assert sleeper.waits == []
|
||||
assert get_model_cost_map_source_info()["source"] == "local"
|
||||
|
||||
get_model_cost_map(url="not a url", sleep=sleeper, rng=random.Random(0))
|
||||
assert sleeper.waits == []
|
||||
assert get_model_cost_map_source_info()["source"] == "local"
|
||||
assert _retry_threads() == []
|
||||
source = get_model_cost_map_source_info()
|
||||
assert source["source"] == "local"
|
||||
assert source["fallback_reason"] is not None
|
||||
|
||||
|
||||
def test_boot_load_respects_local_env_override(monkeypatch):
|
||||
|
|
@ -701,7 +737,9 @@ def test_boot_load_that_fails_the_integrity_check_reports_the_backup_not_the_rej
|
|||
)
|
||||
get_model_cost_map(url=_URL, sleep=_SyncSleepRecorder(), rng=random.Random(0), client=remote)
|
||||
shrunk_body = b'{"gpt-5.4-mini": {"mode": "chat", "input_cost_per_token": 1e-06, "output_cost_per_token": 2e-06}}'
|
||||
shrunk, _ = _mock_client([httpx.Response(200, headers={"ETag": 'W/"shrunk"'}, content=shrunk_body)], client_cls=httpx.Client)
|
||||
shrunk, _ = _mock_client(
|
||||
[httpx.Response(200, headers={"ETag": 'W/"shrunk"'}, content=shrunk_body)], client_cls=httpx.Client
|
||||
)
|
||||
|
||||
get_model_cost_map(url=_URL, sleep=_SyncSleepRecorder(), rng=random.Random(0), client=shrunk)
|
||||
|
||||
|
|
|
|||
|
|
@ -38,6 +38,14 @@ def flush_shared_bedrock_iam_cache():
|
|||
yield
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clean_ssl_env(monkeypatch):
|
||||
"""get_ssl_verify reads these, so the sts client's verify= would otherwise depend on
|
||||
the ambient environment. The published images set SSL_CERT_FILE."""
|
||||
for env_var in ("SSL_CERT_FILE", "SSL_VERIFY"):
|
||||
monkeypatch.delenv(env_var, raising=False)
|
||||
|
||||
|
||||
def test_base_aws_llm_instances_share_process_wide_iam_cache():
|
||||
"""Regression LIT-2662: new instances must reuse iam_cache (Bedrock passthrough is per-request)."""
|
||||
first = BaseAWSLLM()
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import inspect
|
|||
import json
|
||||
import sys
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional
|
||||
from typing import Any, Dict, Final, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11
|
||||
|
|
@ -113,7 +113,7 @@ class TestExecuteWithMcpClient:
|
|||
assert "stack_trace" not in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_timeout_caps_hanging_operation_and_names_url(self, monkeypatch):
|
||||
async def test_timeout_caps_hanging_operation_and_names_origin(self, monkeypatch):
|
||||
async def fake_create_client(*args, **kwargs):
|
||||
return object()
|
||||
|
||||
|
|
@ -138,7 +138,7 @@ class TestExecuteWithMcpClient:
|
|||
)
|
||||
|
||||
assert result["error"] is True
|
||||
assert "https://mcp.example.com/mcp/" in result["message"]
|
||||
assert "https://mcp.example.com" in result["message"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_timeout_covers_client_creation(self, monkeypatch):
|
||||
|
|
@ -166,15 +166,15 @@ class TestExecuteWithMcpClient:
|
|||
)
|
||||
|
||||
assert result["error"] is True
|
||||
assert "https://mcp.example.com/mcp/" in result["message"]
|
||||
assert "https://mcp.example.com" in result["message"]
|
||||
|
||||
def test_timeout_defaults_to_tool_listing_timeout(self):
|
||||
default = inspect.signature(rest_endpoints._execute_with_mcp_client).parameters["timeout_seconds"].default
|
||||
assert default == MCP_TOOL_LISTING_TIMEOUT
|
||||
|
||||
def test_connection_error_message_timeout_names_url_and_budget(self):
|
||||
def test_connection_error_message_timeout_names_origin_and_budget(self):
|
||||
message = rest_endpoints._connection_error_message(TimeoutError(), "https://api.example.com/mcp/", 30.0)
|
||||
assert "https://api.example.com/mcp/" in message
|
||||
assert "https://api.example.com" in message
|
||||
assert "30s" in message
|
||||
|
||||
def test_connection_error_message_hides_arbitrary_http_exception_detail(self):
|
||||
|
|
@ -592,7 +592,7 @@ class TestExecuteWithMcpClient:
|
|||
|
||||
assert result["status"] == "error"
|
||||
assert result["error"] is True
|
||||
assert "Failed to connect to MCP server" in result["message"]
|
||||
assert "reference" in result["message"]
|
||||
# Error message must not leak raw exception details
|
||||
assert "cancel scope" not in result["message"]
|
||||
|
||||
|
|
@ -3427,10 +3427,191 @@ class TestConnectionErrorMessage:
|
|||
message = rest_endpoints._connection_error_message(exc, "https://example.com", 30.0)
|
||||
assert "503" in message
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error_type", [httpx.ReadError, httpx.WriteError, httpx.RemoteProtocolError, ConnectionResetError]
|
||||
)
|
||||
def test_interrupted_connection_message_is_safe(self, error_type: type[Exception]) -> None:
|
||||
message: Final = rest_endpoints._connection_error_message(
|
||||
error_type("secret-transport-detail"), "https://example.com/?token=secret-query", 30
|
||||
)
|
||||
assert "connection was interrupted" in message
|
||||
assert "secret" not in message
|
||||
|
||||
def test_closed_connection_explains_incomplete_request(self) -> None:
|
||||
from mcp import McpError
|
||||
from mcp.types import ErrorData
|
||||
|
||||
message: Final = rest_endpoints._connection_error_message(
|
||||
McpError(ErrorData(code=-32000, message="Connection closed", data="secret-data")), None, 30
|
||||
)
|
||||
assert "connection was closed before the request completed" in message
|
||||
assert "secret" not in message
|
||||
|
||||
def test_timeout_does_not_claim_the_server_sent_nothing(self) -> None:
|
||||
message: Final = rest_endpoints._connection_error_message(TimeoutError(), None, 30)
|
||||
assert "no valid MCP response received" in message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("sdk_timeout", [True, False])
|
||||
@pytest.mark.parametrize("read_timeout", [0, 1])
|
||||
async def test_timeout_message_uses_the_deadline_that_expired(self, sdk_timeout: bool, read_timeout: int) -> None:
|
||||
from mcp import McpError
|
||||
from mcp.types import ErrorData
|
||||
|
||||
async def operation(client: rest_endpoints.MCPClient) -> dict[str, object]:
|
||||
try:
|
||||
raise TimeoutError("secret-timeout")
|
||||
except TimeoutError as elapsed:
|
||||
if not sdk_timeout:
|
||||
raise
|
||||
try:
|
||||
raise McpError(ErrorData(code=408, message="secret-sdk-timeout")) from elapsed
|
||||
except McpError as sdk_error:
|
||||
raise TimeoutError() from sdk_error
|
||||
|
||||
payload: Final = NewMCPServerRequest(
|
||||
server_name="timeout", url="https://example.com", auth_type=MCPAuth.none, timeout=read_timeout
|
||||
)
|
||||
result: Final = await rest_endpoints._execute_with_mcp_client(payload, operation, timeout_seconds=30)
|
||||
assert (f"within {read_timeout}s" if sdk_timeout else "within 30s") in result["message"]
|
||||
assert "secret" not in result["message"]
|
||||
|
||||
def test_unknown_error_falls_back_to_generic(self):
|
||||
message = rest_endpoints._connection_error_message(RuntimeError("weird"), "https://example.com", 30.0)
|
||||
assert "weird" not in message
|
||||
assert "proxy logs" in message.lower()
|
||||
assert "reference" in message.lower()
|
||||
|
||||
def test_sdk_session_terminated_explains_endpoint_and_retry(self) -> None:
|
||||
from mcp.shared.exceptions import McpError
|
||||
from mcp.types import ErrorData
|
||||
|
||||
message: Final = rest_endpoints._connection_error_message(
|
||||
McpError(ErrorData(code=32600, message="Session terminated")), "https://example.com/mcp", 30.0
|
||||
)
|
||||
|
||||
assert "session was terminated" in message
|
||||
assert "MCP endpoint" in message
|
||||
assert "transport" in message
|
||||
assert "retry" in message
|
||||
assert "404" not in message
|
||||
|
||||
@pytest.mark.parametrize("code", [-32700, -32601, -32602, -32603, -32000, 32600, 408])
|
||||
def test_rpc_errors_include_code_without_echoing_upstream_data(self, code: int) -> None:
|
||||
from mcp.shared.exceptions import McpError
|
||||
from mcp.types import ErrorData
|
||||
|
||||
message: Final = rest_endpoints._connection_error_message(
|
||||
McpError(ErrorData(code=code, message="secret-message", data={"token": "secret-data"})),
|
||||
"https://example.com/secret-path?token=secret-query",
|
||||
30.0,
|
||||
)
|
||||
|
||||
assert f"JSON-RPC code {code}" in message
|
||||
assert "secret" not in message
|
||||
assert "timed out" not in message
|
||||
assert "session was terminated" not in message
|
||||
|
||||
@pytest.mark.parametrize("status_code", [401, 403, 404, 405, 429, 503])
|
||||
def test_wrapped_http_failures_preserve_status(self, status_code: int) -> None:
|
||||
response: Final = httpx.Response(status_code, text="secret-body")
|
||||
upstream: Final = httpx.HTTPStatusError(
|
||||
"secret-exception",
|
||||
request=httpx.Request("POST", "https://example.com/?token=secret-query"),
|
||||
response=response,
|
||||
)
|
||||
wrapped: Final = BaseExceptionGroup(
|
||||
"secret-group", [asyncio.CancelledError(), BaseExceptionGroup("nested", [upstream])]
|
||||
)
|
||||
|
||||
message: Final = rest_endpoints._connection_error_message(wrapped, "https://example.com", 30.0)
|
||||
|
||||
assert f"HTTP {status_code}" in message
|
||||
assert "secret" not in message
|
||||
|
||||
def test_explicit_cause_is_classified_before_incidental_context(self) -> None:
|
||||
wrapped: Final = RuntimeError("secret-wrapper")
|
||||
wrapped.__cause__ = httpx.ConnectError("secret-cause")
|
||||
wrapped.__context__ = TimeoutError("secret-context")
|
||||
|
||||
message: Final = rest_endpoints._connection_error_message(wrapped, "https://example.com", 30.0)
|
||||
|
||||
assert "unreachable" in message
|
||||
assert "secret" not in message
|
||||
|
||||
def test_timeout_url_redacts_credentials_path_query_and_fragment(self) -> None:
|
||||
message: Final = rest_endpoints._connection_error_message(
|
||||
TimeoutError("secret-error"),
|
||||
"https://secret-user:secret-pass@example.com:8443/secret-path?token=secret-query#secret-fragment",
|
||||
30.0,
|
||||
)
|
||||
|
||||
assert "https://example.com:8443" in message
|
||||
assert "30s" in message
|
||||
assert "secret" not in message
|
||||
|
||||
def test_unknown_failure_reference_matches_safe_diagnostics(self, caplog: pytest.LogCaptureFixture) -> None:
|
||||
import re
|
||||
|
||||
try:
|
||||
raise RuntimeError("secret-exception-body")
|
||||
except RuntimeError as exc:
|
||||
message: Final = rest_endpoints._connection_error_message(
|
||||
exc, "https://secret-user:secret-password@example.com/secret-path?token=secret-query", 30.0
|
||||
)
|
||||
|
||||
reference: Final = re.search(r"reference ([a-f0-9]{32})", message)
|
||||
assert reference is not None
|
||||
diagnostics: Final = tuple(
|
||||
record for record in caplog.records if "MCP connection test failed" in record.message
|
||||
)
|
||||
assert len(diagnostics) == 1
|
||||
assert reference.group(1) in diagnostics[0].message
|
||||
assert "RuntimeError" in diagnostics[0].message
|
||||
assert "test_unknown_failure_reference_matches_safe_diagnostics" in diagnostics[0].message
|
||||
assert diagnostics[0].exc_info is None
|
||||
assert "secret" not in message + diagnostics[0].message
|
||||
|
||||
@pytest.mark.parametrize("exc", [ValueError("secret-config"), HTTPException(500, "secret-detail")])
|
||||
def test_unrelated_errors_are_not_misreported_as_invalid_mcp(self, exc: Exception) -> None:
|
||||
message: Final = rest_endpoints._connection_error_message(exc, "https://example.com", 30.0)
|
||||
|
||||
assert "reference" in message
|
||||
assert "invalid MCP response" not in message
|
||||
assert "secret" not in message
|
||||
|
||||
def test_configuration_validation_error_uses_unknown_fallback(self) -> None:
|
||||
from pydantic import ValidationError
|
||||
|
||||
with pytest.raises(ValidationError) as caught:
|
||||
NewMCPServerRequest.model_validate({"server_name": "example", "transport": "secret-invalid-transport"})
|
||||
|
||||
message: Final = rest_endpoints._connection_error_message(caught.value, "https://example.com", 30.0)
|
||||
assert "reference" in message
|
||||
assert "invalid MCP response" not in message
|
||||
assert "secret" not in message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_test_preserves_cancellation(self) -> None:
|
||||
async def cancelled_operation(client: rest_endpoints.MCPClient) -> dict[str, object]:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
payload: Final = NewMCPServerRequest(server_name="cancelled", url="https://example.com", auth_type=MCPAuth.none)
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await rest_endpoints._execute_with_mcp_client(payload, cancelled_operation)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_failure_preserves_response_contract(self) -> None:
|
||||
async def failing_operation(client: rest_endpoints.MCPClient) -> dict[str, object]:
|
||||
raise RuntimeError("secret-operation")
|
||||
|
||||
payload: Final = NewMCPServerRequest(server_name="unknown", url="https://example.com", auth_type=MCPAuth.none)
|
||||
result: Final = await rest_endpoints._execute_with_mcp_client(payload, failing_operation)
|
||||
|
||||
assert result["error"] is True
|
||||
assert result["status"] == "error"
|
||||
assert "reference" in result["message"]
|
||||
assert "secret" not in result["message"]
|
||||
assert "stack_trace" not in result
|
||||
|
||||
|
||||
class TestGetServerAuthHeaderGroupDefault:
|
||||
|
|
|
|||
|
|
@ -2374,6 +2374,44 @@ def _mock_prisma_for_team_lookup(find_unique):
|
|||
return mock_prisma_client
|
||||
|
||||
|
||||
_TEAM_ALIAS_TABLE_ROW = {"id": 1, "model_aliases": '{"fast": "gpt-4o"}', "created_by": "admin", "updated_by": "admin"}
|
||||
|
||||
|
||||
def _prisma_team_row(include):
|
||||
"""Mimics Prisma: the `litellm_model_table` relation rides on the row only when the query `include`s it."""
|
||||
columns = {"team_id": "team-aliases", "team_alias": "aliases", "models": ["gpt-4o"]}
|
||||
row = (
|
||||
{**columns, "litellm_model_table": _TEAM_ALIAS_TABLE_ROW}
|
||||
if (include or {}).get("litellm_model_table")
|
||||
else columns
|
||||
)
|
||||
return SimpleNamespace(dict=lambda: row, model_dump=lambda: row)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_object_loads_model_aliases_relation():
|
||||
"""LIT-5858: the auth path read teams without `include`ing `litellm_model_table`, so every JWT
|
||||
team came back with `model_aliases=None` and alias requests 403'd."""
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
from litellm.proxy.auth.team_grants import team_model_aliases
|
||||
|
||||
async def find_unique(where, include=None):
|
||||
return _prisma_team_row(include)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
team = await get_team_object(
|
||||
team_id="team-aliases",
|
||||
prisma_client=_mock_prisma_for_team_lookup(AsyncMock(side_effect=find_unique)),
|
||||
user_api_key_cache=mock_cache,
|
||||
check_db_only=True,
|
||||
)
|
||||
|
||||
assert team_model_aliases(team) == {"fast": "gpt-4o"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_object_distinguishes_absent_team_from_unreadable_row():
|
||||
"""A deleted team and a database that would not answer both surface as a 404,
|
||||
|
|
@ -6195,6 +6233,32 @@ async def test_get_team_object_by_alias_db_fetch_returns_cached_obj():
|
|||
assert result.models == ["gpt-4"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_object_by_alias_loads_model_aliases_relation():
|
||||
"""LIT-5858: same regression as `test_get_team_object_loads_model_aliases_relation`, for the
|
||||
`team_alias_jwt_field` lookup."""
|
||||
from litellm.proxy.auth.auth_checks import get_team_object_by_alias
|
||||
from litellm.proxy.auth.team_grants import team_model_aliases
|
||||
|
||||
async def find_many(where, include=None):
|
||||
return [_prisma_team_row(include)]
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(side_effect=find_many)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
team = await get_team_object_by_alias(
|
||||
team_alias="aliases",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_cache,
|
||||
)
|
||||
|
||||
assert team_model_aliases(team) == {"fast": "gpt-4o"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_org_object_by_alias_db_fetch_returns_validated_org():
|
||||
from litellm.proxy._types import LiteLLM_OrganizationTable
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm.proxy._types import (
|
|||
DEFAULT_JWKS_STALE_TTL,
|
||||
JWTLiteLLMRoleMap,
|
||||
LiteLLM_JWTAuth,
|
||||
LiteLLM_ModelTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_UserTable,
|
||||
|
|
@ -1255,6 +1256,57 @@ async def test_find_team_with_model_access_model_group(monkeypatch):
|
|||
assert team_obj.team_id == "team-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"model_aliases",
|
||||
['{"fast": "gpt-4o"}', {"fast": "gpt-4o"}],
|
||||
ids=["json-string", "dict"],
|
||||
)
|
||||
async def test_find_team_with_model_access_resolves_team_model_alias(monkeypatch, model_aliases):
|
||||
"""LIT-5858: a JWT team that grants `gpt-4o` under the alias `fast` must resolve a request
|
||||
for `fast`. The JWT path used to pass `team_model_aliases=None`, so every alias request 403'd."""
|
||||
import sys
|
||||
import types
|
||||
|
||||
from litellm.caching import DualCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
from litellm.router import Router
|
||||
|
||||
router = Router(model_list=[{"model_name": "gpt-4o", "litellm_params": {"model": "gpt-4o"}}])
|
||||
proxy_server_module = types.ModuleType("proxy_server")
|
||||
proxy_server_module.llm_router = router
|
||||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module)
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team-aliases",
|
||||
models=["gpt-4o"],
|
||||
litellm_model_table=LiteLLM_ModelTable(model_aliases=model_aliases, created_by="admin", updated_by="admin"),
|
||||
)
|
||||
|
||||
async def mock_get_team_object(*args, **kwargs):
|
||||
return team
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object)
|
||||
|
||||
jwt_handler = JWTHandler()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth()
|
||||
user_api_key_cache = DualCache()
|
||||
|
||||
team_id, team_obj = await JWTAuthManager.find_team_with_model_access(
|
||||
team_ids={"team-aliases"},
|
||||
requested_model="fast",
|
||||
route="/chat/completions",
|
||||
jwt_handler=jwt_handler,
|
||||
prisma_client=None,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=ProxyLogging(user_api_key_cache=user_api_key_cache),
|
||||
)
|
||||
|
||||
assert team_id == "team-aliases"
|
||||
assert team_obj is team
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatch):
|
||||
"""Regression for #31189: a single-team JWT that grants the requested model
|
||||
|
|
|
|||
129
tests/test_litellm/proxy/auth/test_team_grants.py
Normal file
129
tests/test_litellm/proxy/auth/test_team_grants.py
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
import pytest
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_VerificationTokenView,
|
||||
Member,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.models.team import LiteLLM_ModelTable
|
||||
from litellm.proxy.auth.team_grants import team_grants, team_model_aliases
|
||||
|
||||
TEAM_ID = "team-grants"
|
||||
USER_ID = "user-in-team"
|
||||
ALIASES = {"fast": "gpt-4o-mini", "smart": "gpt-4o"}
|
||||
|
||||
|
||||
def _alias_table(model_aliases) -> LiteLLM_ModelTable:
|
||||
return LiteLLM_ModelTable(model_aliases=model_aliases, created_by="admin", updated_by="admin")
|
||||
|
||||
|
||||
def _full_team(model_aliases=ALIASES) -> LiteLLM_TeamTable:
|
||||
return LiteLLM_TeamTable(
|
||||
team_id=TEAM_ID,
|
||||
team_alias="grants-team",
|
||||
tpm_limit=1000,
|
||||
rpm_limit=10,
|
||||
max_budget=50.0,
|
||||
soft_budget=25.0,
|
||||
spend=12.5,
|
||||
models=["gpt-4o", "gpt-4o-mini"],
|
||||
blocked=True,
|
||||
metadata={"tier": "gold"},
|
||||
litellm_model_table=_alias_table(model_aliases),
|
||||
object_permission_id="op-1",
|
||||
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-1", mcp_servers=["mcp-a"]),
|
||||
members_with_roles=[
|
||||
Member(user_id="someone-else", role="user"),
|
||||
Member(user_id=USER_ID, role="admin"),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _membership() -> LiteLLM_TeamMembership:
|
||||
return LiteLLM_TeamMembership(
|
||||
user_id=USER_ID,
|
||||
team_id=TEAM_ID,
|
||||
spend=3.25,
|
||||
litellm_budget_table=LiteLLM_BudgetTable(tpm_limit=500, rpm_limit=5),
|
||||
)
|
||||
|
||||
|
||||
def test_team_grants_cover_every_team_field_the_key_path_gets():
|
||||
"""Class guard for LIT-5858 and its siblings: every ``team_*`` column the combined-view SQL hands the
|
||||
virtual-key path must come out of the projection too, with the team's actual value, so adding a column
|
||||
to ``LiteLLM_VerificationTokenView`` without teaching ``team_grants`` fails here instead of in prod."""
|
||||
team = _full_team()
|
||||
grants = team_grants(team_object=team, team_membership=_membership(), user_id=USER_ID)
|
||||
token = UserAPIKeyAuth(team_id=TEAM_ID, **grants)
|
||||
|
||||
view_team_fields = {name for name in LiteLLM_VerificationTokenView.model_fields if name.startswith("team_")}
|
||||
assert view_team_fields - {"team_id"} <= set(grants)
|
||||
assert all(grants[name] is not None for name in view_team_fields - {"team_id"})
|
||||
|
||||
assert token.team_alias == "grants-team"
|
||||
assert token.team_tpm_limit == 1000
|
||||
assert token.team_rpm_limit == 10
|
||||
assert token.team_max_budget == 50.0
|
||||
assert token.team_soft_budget == 25.0
|
||||
assert token.team_spend == 12.5
|
||||
assert token.team_models == ["gpt-4o", "gpt-4o-mini"]
|
||||
assert token.team_blocked is True
|
||||
assert token.team_metadata == {"tier": "gold"}
|
||||
assert token.team_model_aliases == ALIASES
|
||||
assert token.team_object_permission_id == "op-1"
|
||||
assert token.team_object_permission is not None
|
||||
assert token.team_object_permission.mcp_servers == ["mcp-a"]
|
||||
assert token.team_member == Member(user_id=USER_ID, role="admin")
|
||||
assert token.team_member_spend == 3.25
|
||||
assert token.team_member_tpm_limit == 500
|
||||
assert token.team_member_rpm_limit == 5
|
||||
|
||||
|
||||
def test_team_grants_without_team_leave_token_defaults():
|
||||
token = UserAPIKeyAuth(**team_grants(team_object=None, team_membership=None, user_id=USER_ID))
|
||||
assert token == UserAPIKeyAuth()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"stored_aliases",
|
||||
[ALIASES, '{"fast": "gpt-4o-mini", "smart": "gpt-4o"}'],
|
||||
ids=["json-object", "json-string-as-written-by-team-new"],
|
||||
)
|
||||
def test_team_model_aliases_decode_both_storage_shapes(stored_aliases):
|
||||
team = _full_team(model_aliases=stored_aliases)
|
||||
assert team_model_aliases(team) == ALIASES
|
||||
assert team_grants(team_object=team, team_membership=None, user_id=None)["team_model_aliases"] == ALIASES
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stored_aliases", [None, "not json", '["a", "b"]', {"fast": 3}], ids=str)
|
||||
def test_team_model_aliases_treat_unusable_column_as_no_aliases(stored_aliases):
|
||||
team = _full_team(model_aliases=stored_aliases)
|
||||
assert team_model_aliases(team) is None
|
||||
assert team_grants(team_object=team, team_membership=None, user_id=None)["team_model_aliases"] is None
|
||||
|
||||
|
||||
def test_team_model_aliases_none_without_relation_loaded():
|
||||
team = _full_team()
|
||||
team.litellm_model_table = None
|
||||
assert team_model_aliases(team) is None
|
||||
assert team_model_aliases(None) is None
|
||||
|
||||
|
||||
def test_team_member_is_the_callers_row_only():
|
||||
team = _full_team()
|
||||
assert team_grants(team_object=team, team_membership=None, user_id="someone-else")["team_member"] == Member(
|
||||
user_id="someone-else", role="user"
|
||||
)
|
||||
assert team_grants(team_object=team, team_membership=None, user_id="stranger")["team_member"] is None
|
||||
assert team_grants(team_object=team, team_membership=None, user_id=None)["team_member"] is None
|
||||
|
||||
|
||||
def test_membership_limits_absent_without_membership_row():
|
||||
grants = team_grants(team_object=_full_team(), team_membership=None, user_id=USER_ID)
|
||||
assert grants["team_member_spend"] is None
|
||||
assert grants["team_member_tpm_limit"] is None
|
||||
assert grants["team_member_rpm_limit"] is None
|
||||
|
|
@ -7133,3 +7133,119 @@ def test_user_api_key_auth_opens_a_datadog_span_for_accepted_and_rejected_keys(t
|
|||
assert report["outcomes"] == ["accepted", "rejected"]
|
||||
auth_span = "litellm.proxy.auth.user_api_key_auth.user_api_key_auth"
|
||||
assert [span for span in report["spans"] if span == auth_span] == [auth_span, auth_span]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("is_proxy_admin", [False, True], ids=["standard-return", "proxy-admin-return"])
|
||||
async def test_jwt_builder_returns_every_team_grant_the_key_path_gets(is_proxy_admin):
|
||||
"""LIT-5858: the team-based JWT path hand-built ``UserAPIKeyAuth`` from a short list of team fields, so the
|
||||
team's model aliases (and on the admin return, its object permission) never reached the token and alias
|
||||
requests 403'd. Both returns now go through ``team_grants``; pin the fields that used to be dropped."""
|
||||
import litellm.proxy.proxy_server as _proxy_server_mod
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
from litellm.models.team import LiteLLM_ModelTable
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
Member,
|
||||
)
|
||||
|
||||
class _AcceptEveryJwt(JWTHandler):
|
||||
def is_jwt(self, token: str) -> bool:
|
||||
return True
|
||||
|
||||
jwt_handler = _AcceptEveryJwt()
|
||||
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth()
|
||||
|
||||
team = LiteLLM_TeamTable(
|
||||
team_id="team-jwt-aliases",
|
||||
team_alias="jwt-aliases",
|
||||
models=["gpt-4o"],
|
||||
max_budget=40.0,
|
||||
spend=4.0,
|
||||
blocked=False,
|
||||
metadata={"tier": "gold"},
|
||||
litellm_model_table=LiteLLM_ModelTable(
|
||||
model_aliases='{"fast": "gpt-4o"}', created_by="admin", updated_by="admin"
|
||||
),
|
||||
object_permission_id="op-jwt",
|
||||
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-jwt", mcp_servers=["mcp-a"]),
|
||||
members_with_roles=[Member(user_id="jwt-user", role="admin")],
|
||||
)
|
||||
membership = LiteLLM_TeamMembership(user_id="jwt-user", team_id="team-jwt-aliases", spend=1.5)
|
||||
builder_result = {
|
||||
"is_proxy_admin": is_proxy_admin,
|
||||
"team_object": team,
|
||||
"user_object": None,
|
||||
"end_user_object": None,
|
||||
"org_object": None,
|
||||
"token": "jwt",
|
||||
"team_id": "team-jwt-aliases",
|
||||
"user_id": "jwt-user",
|
||||
"user_email": "jwt-user@example.com",
|
||||
"end_user_id": None,
|
||||
"org_id": None,
|
||||
"team_membership": membership,
|
||||
"jwt_claims": {"sub": "jwt-user"},
|
||||
}
|
||||
|
||||
mock_proxy_logging_obj = MagicMock()
|
||||
mock_proxy_logging_obj.internal_usage_cache = MagicMock()
|
||||
mock_proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
attrs = {
|
||||
"prisma_client": MagicMock(),
|
||||
"user_api_key_cache": DualCache(),
|
||||
"proxy_logging_obj": mock_proxy_logging_obj,
|
||||
"master_key": "sk-master-key",
|
||||
"general_settings": {"enable_jwt_auth": True},
|
||||
"llm_model_list": [],
|
||||
"llm_router": None,
|
||||
"open_telemetry_logger": None,
|
||||
"model_max_budget_limiter": MagicMock(),
|
||||
"user_custom_auth": None,
|
||||
"jwt_handler": jwt_handler,
|
||||
"premium_user": True,
|
||||
"litellm_proxy_admin_name": "admin",
|
||||
}
|
||||
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
|
||||
try:
|
||||
for k, v in attrs.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
request = Request(scope={"type": "http", "headers": [], "method": "POST"})
|
||||
request._url = URL(url="/chat/completions")
|
||||
with patch( # test-quality-ok: auth_builder is the claim-resolution seam; the regression is how its result is projected onto the token
|
||||
"litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder",
|
||||
new_callable=AsyncMock,
|
||||
return_value=builder_result,
|
||||
):
|
||||
token = await _user_api_key_auth_builder(
|
||||
request=request,
|
||||
api_key="Bearer header.payload.signature",
|
||||
azure_api_key_header="",
|
||||
anthropic_api_key_header=None,
|
||||
google_ai_studio_api_key_header=None,
|
||||
azure_apim_header=None,
|
||||
request_data={},
|
||||
)
|
||||
finally:
|
||||
for k, v in originals.items():
|
||||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
assert token.team_id == "team-jwt-aliases"
|
||||
assert token.user_role == (LitellmUserRoles.PROXY_ADMIN if is_proxy_admin else LitellmUserRoles.INTERNAL_USER)
|
||||
assert token.team_model_aliases == {"fast": "gpt-4o"}
|
||||
assert token.team_object_permission is not None
|
||||
assert token.team_object_permission.mcp_servers == ["mcp-a"]
|
||||
assert token.team_object_permission_id == "op-jwt"
|
||||
assert token.team_alias == "jwt-aliases"
|
||||
assert token.team_models == ["gpt-4o"]
|
||||
assert token.team_max_budget == 40.0
|
||||
assert token.team_spend == 4.0
|
||||
assert token.team_metadata == {"tier": "gold"}
|
||||
assert token.team_member == Member(user_id="jwt-user", role="admin")
|
||||
assert token.team_member_spend == 1.5
|
||||
assert token.jwt_claims == {"sub": "jwt-user"}
|
||||
|
|
|
|||
|
|
@ -11,15 +11,30 @@ clobber a pre-existing ``DATABASE_URL_READ_REPLICA``. A pre-existing
|
|||
``DATABASE_URL`` (password auth) is likewise left untouched.
|
||||
"""
|
||||
|
||||
import datetime
|
||||
import hashlib
|
||||
import os
|
||||
import socket
|
||||
import ssl
|
||||
import tempfile
|
||||
import threading
|
||||
import urllib.parse
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from cryptography import x509
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import ec
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.proxy.db.db_url_settings import (
|
||||
PG_SSL_REQUEST,
|
||||
DatabaseURLSettings,
|
||||
translate_libpq_ssl_params,
|
||||
unsupported_db_scheme,
|
||||
unsupported_db_scheme_message,
|
||||
)
|
||||
|
|
@ -381,9 +396,7 @@ def test_writer_password_is_percent_encoded(monkeypatch):
|
|||
def test_writer_url_not_clobbered_when_already_set(monkeypatch):
|
||||
"""An operator-pinned DATABASE_URL (e.g. helm's $(VAR) assembly) always
|
||||
wins over the discrete fields."""
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL", "postgresql://pinned:url@db.example.com:5432/litellm_db"
|
||||
)
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://pinned:url@db.example.com:5432/litellm_db")
|
||||
monkeypatch.setenv("DATABASE_HOST", "writer.example.com")
|
||||
monkeypatch.setenv("DATABASE_USER", "litellm")
|
||||
monkeypatch.setenv("DATABASE_NAME", "litellm_db")
|
||||
|
|
@ -515,9 +528,7 @@ def test_apply_to_env_rejects_pinned_sqlite_direct_url(monkeypatch):
|
|||
|
||||
def test_apply_to_env_rejects_pinned_non_postgres_reader(monkeypatch):
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@writer.example.com:5432/db")
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL_READ_REPLICA", "mysql://u:p@reader.example.com:3306/db"
|
||||
)
|
||||
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", "mysql://u:p@reader.example.com:3306/db")
|
||||
|
||||
with pytest.raises(RuntimeError, match=r"DATABASE_URL_READ_REPLICA.*mysql"):
|
||||
_apply()
|
||||
|
|
@ -542,15 +553,11 @@ def test_reader_inherits_writer_connection_params(monkeypatch):
|
|||
"DATABASE_URL",
|
||||
"postgresql://u:p@writer.example.com:5432/db?connection_limit=3&pool_timeout=20&pgbouncer=true",
|
||||
)
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL_READ_REPLICA", "postgresql://u:p@reader.example.com:5432/db"
|
||||
)
|
||||
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", "postgresql://u:p@reader.example.com:5432/db")
|
||||
|
||||
_apply()
|
||||
|
||||
query = urllib.parse.parse_qs(
|
||||
urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query
|
||||
)
|
||||
query = urllib.parse.parse_qs(urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query)
|
||||
assert query["connection_limit"] == ["3"]
|
||||
assert query["pool_timeout"] == ["20"]
|
||||
assert query["pgbouncer"] == ["true"]
|
||||
|
|
@ -568,9 +575,7 @@ def test_reader_keeps_its_own_pinned_connection_params(monkeypatch):
|
|||
|
||||
_apply()
|
||||
|
||||
query = urllib.parse.parse_qs(
|
||||
urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query
|
||||
)
|
||||
query = urllib.parse.parse_qs(urllib.parse.urlsplit(os.environ["DATABASE_URL_READ_REPLICA"]).query)
|
||||
assert query["connection_limit"] == ["50"]
|
||||
assert query["pool_timeout"] == ["20"]
|
||||
|
||||
|
|
@ -776,6 +781,170 @@ def test_libpq_verify_full_and_sslrootcert_become_prisma_strict_sslcert(monkeypa
|
|||
}
|
||||
|
||||
|
||||
def _issue_cert(
|
||||
subject: str, issuer: x509.Certificate | None, issuer_key: ec.EllipticCurvePrivateKey | None, ca: bool
|
||||
) -> tuple[x509.Certificate, ec.EllipticCurvePrivateKey]:
|
||||
key: Final = ec.generate_private_key(ec.SECP256R1())
|
||||
name: Final = x509.Name((x509.NameAttribute(x509.NameOID.COMMON_NAME, subject),))
|
||||
now: Final = datetime.datetime.now(datetime.timezone.utc)
|
||||
builder: Final = (
|
||||
x509.CertificateBuilder()
|
||||
.subject_name(name)
|
||||
.issuer_name(issuer.subject if issuer else name)
|
||||
.public_key(key.public_key())
|
||||
.serial_number(x509.random_serial_number())
|
||||
.not_valid_before(now - datetime.timedelta(minutes=5))
|
||||
.not_valid_after(now + datetime.timedelta(days=1))
|
||||
.add_extension(x509.BasicConstraints(ca=ca, path_length=None), critical=True)
|
||||
.add_extension(x509.SubjectAlternativeName((x509.DNSName("localhost"),)), critical=False)
|
||||
)
|
||||
return builder.sign(issuer_key or key, hashes.SHA256()), key
|
||||
|
||||
|
||||
def _pem(cert: x509.Certificate) -> bytes:
|
||||
return cert.public_bytes(serialization.Encoding.PEM)
|
||||
|
||||
|
||||
class _TlsPostgresStub:
|
||||
"""Answers one libpq ``SSLRequest`` with ``S`` and serves ``leaf + intermediate``."""
|
||||
|
||||
def __init__(self, chain_pem: Path, key_pem: Path) -> None:
|
||||
self.context: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||||
self.context.load_cert_chain(str(chain_pem), str(key_pem))
|
||||
self.listener: Final = socket.create_server(("127.0.0.1", 0))
|
||||
self.port: Final[int] = self.listener.getsockname()[1]
|
||||
self.thread: Final = threading.Thread(target=self._serve, daemon=True)
|
||||
self.thread.start()
|
||||
|
||||
def _serve(self) -> None:
|
||||
with self.listener:
|
||||
while True:
|
||||
try:
|
||||
conn: socket.socket = self.listener.accept()[0]
|
||||
except OSError:
|
||||
return
|
||||
with conn:
|
||||
try:
|
||||
if conn.recv(8) == PG_SSL_REQUEST:
|
||||
conn.sendall(b"S")
|
||||
with self.context.wrap_socket(conn, server_side=True) as tls:
|
||||
tls.recv(1)
|
||||
except OSError:
|
||||
continue
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RdsLikePki:
|
||||
bundle: Path
|
||||
wrong_bundle: Path
|
||||
root: Path
|
||||
port: int
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def rds_like_pki(tmp_path: Path) -> Iterator[_RdsLikePki]:
|
||||
"""An RDS-shaped trust setup: the server sends leaf + intermediate, the
|
||||
bundle holds only self-signed roots, and the right root is not first."""
|
||||
root, root_key = _issue_cert("Real Root CA", None, None, ca=True)
|
||||
decoys: Final = tuple(_issue_cert(f"Decoy Root CA {i}", None, None, ca=True)[0] for i in range(3))
|
||||
intermediate, intermediate_key = _issue_cert("Intermediate CA", root, root_key, ca=True)
|
||||
leaf, leaf_key = _issue_cert("localhost", intermediate, intermediate_key, ca=False)
|
||||
chain_pem: Final = tmp_path / "server-chain.pem"
|
||||
chain_pem.write_bytes(_pem(leaf) + _pem(intermediate))
|
||||
key_pem: Final = tmp_path / "server.key"
|
||||
key_pem.write_bytes(
|
||||
leaf_key.private_bytes(
|
||||
serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()
|
||||
)
|
||||
)
|
||||
bundle: Final = tmp_path / "global-bundle.pem"
|
||||
bundle.write_bytes(b"".join(_pem(decoy) for decoy in decoys) + _pem(root))
|
||||
wrong_bundle: Final = tmp_path / "wrong-bundle.pem"
|
||||
wrong_bundle.write_bytes(b"".join(_pem(decoy) for decoy in decoys))
|
||||
root_pem: Final = tmp_path / "root.pem"
|
||||
root_pem.write_bytes(_pem(root))
|
||||
stub: Final = _TlsPostgresStub(chain_pem, key_pem)
|
||||
yield _RdsLikePki(bundle=bundle, wrong_bundle=wrong_bundle, root=root_pem, port=stub.port)
|
||||
stub.listener.close()
|
||||
|
||||
|
||||
def _params(url: str) -> tuple[tuple[str, str], ...]:
|
||||
return tuple(urllib.parse.parse_qsl(urllib.parse.urlsplit(url).query, keep_blank_values=True))
|
||||
|
||||
|
||||
def test_multi_root_bundle_is_pinned_to_the_root_the_server_chains_to(
|
||||
monkeypatch: pytest.MonkeyPatch, rds_like_pki: _RdsLikePki
|
||||
):
|
||||
"""Prisma's ``sslcert`` loads only the first certificate of the file, so
|
||||
handing it the whole RDS bundle trusts one region's root and fails with
|
||||
"unable to get local issuer certificate" everywhere else. The URL Prisma
|
||||
receives must point at a single-certificate file holding the server's root."""
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL",
|
||||
f"postgresql://u:p@localhost:{rds_like_pki.port}/litellm_db?sslmode=verify-full&sslrootcert={rds_like_pki.bundle}",
|
||||
)
|
||||
|
||||
_apply()
|
||||
|
||||
(sslmode, sslcert, sslaccept, _) = _params(os.environ["DATABASE_URL"])
|
||||
assert (sslmode, sslaccept) == (("sslmode", "require"), ("sslaccept", "strict"))
|
||||
assert sslcert[0] == "sslcert" and sslcert[1] != str(rds_like_pki.bundle)
|
||||
assert Path(sslcert[1]).read_bytes() == rds_like_pki.root.read_bytes()
|
||||
|
||||
|
||||
def test_pinned_root_replaces_a_planted_symlink_instead_of_writing_through_it(
|
||||
monkeypatch: pytest.MonkeyPatch, rds_like_pki: _RdsLikePki, tmp_path: Path
|
||||
):
|
||||
"""The pinned file has a predictable name in a shared temp dir, so a symlink
|
||||
planted there must not redirect the write onto its target."""
|
||||
monkeypatch.setattr(tempfile, "tempdir", str(tmp_path))
|
||||
root_der: Final = x509.load_pem_x509_certificate(rds_like_pki.root.read_bytes()).public_bytes(
|
||||
serialization.Encoding.DER
|
||||
)
|
||||
pinned: Final = tmp_path / f"litellm-sslcert-{hashlib.sha256(root_der).hexdigest()[:16]}.pem"
|
||||
victim: Final = tmp_path / "victim.txt"
|
||||
victim.write_text("untouched")
|
||||
pinned.symlink_to(victim)
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL",
|
||||
f"postgresql://u:p@localhost:{rds_like_pki.port}/litellm_db?sslmode=verify-full&sslrootcert={rds_like_pki.bundle}",
|
||||
)
|
||||
|
||||
_apply()
|
||||
|
||||
assert ("sslcert", str(pinned)) in _params(os.environ["DATABASE_URL"])
|
||||
assert victim.read_text() == "untouched"
|
||||
assert not pinned.is_symlink() and pinned.read_bytes() == rds_like_pki.root.read_bytes()
|
||||
|
||||
|
||||
def test_bundle_without_the_servers_root_is_passed_through_unchanged(
|
||||
monkeypatch: pytest.MonkeyPatch, rds_like_pki: _RdsLikePki
|
||||
):
|
||||
"""Nothing in the bundle verifies the server, so no root is pinned and
|
||||
Prisma keeps rejecting the connection instead of trusting a root the
|
||||
operator never shipped."""
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL",
|
||||
f"postgresql://u:p@localhost:{rds_like_pki.port}/litellm_db"
|
||||
f"?sslmode=verify-full&sslrootcert={rds_like_pki.wrong_bundle}",
|
||||
)
|
||||
|
||||
_apply()
|
||||
|
||||
assert ("sslcert", str(rds_like_pki.wrong_bundle)) in _params(os.environ["DATABASE_URL"])
|
||||
|
||||
|
||||
def test_root_cert_resolver_receives_the_urls_host_and_default_port():
|
||||
def resolver(cert_path: str, host: str, port: int) -> str:
|
||||
return f"/pinned/{host}/{port}{cert_path}"
|
||||
|
||||
url: Final = translate_libpq_ssl_params(
|
||||
"postgresql://u:p@db.example.com/litellm_db?sslmode=verify-full&sslrootcert=/certs/bundle.pem", resolver
|
||||
)
|
||||
|
||||
assert ("sslcert", "/pinned/db.example.com/5432/certs/bundle.pem") in _params(url)
|
||||
|
||||
|
||||
def test_libpq_verify_ca_becomes_prisma_strict(monkeypatch):
|
||||
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@db.example.com:5432/litellm_db?sslmode=verify-ca")
|
||||
|
||||
|
|
|
|||
|
|
@ -8,8 +8,11 @@ from datetime import datetime, timezone
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.constants import MAX_REDIS_BUFFER_DEQUEUE_COUNT, REDIS_GATEWAY_REQUESTS_BUFFER_KEY
|
||||
from litellm.proxy.db.gateway_request_tracking import (
|
||||
GATEWAY_REQUESTS_JOB_NAME,
|
||||
GatewayRequestAccumulator,
|
||||
GatewayRequestRedisBuffer,
|
||||
commit_gateway_requests_to_db,
|
||||
flush_gateway_requests,
|
||||
)
|
||||
|
|
@ -83,40 +86,47 @@ def test_drain_snapshot_is_not_mutated_by_later_records():
|
|||
# ── commit ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class FakeTable:
|
||||
def __init__(self) -> None:
|
||||
self.upserts: list[dict] = []
|
||||
|
||||
def upsert(self, *, where: dict, data: dict) -> None:
|
||||
self.upserts.append({"where": where, "data": data})
|
||||
|
||||
|
||||
class FakeBatcher:
|
||||
def __init__(self, table: FakeTable) -> None:
|
||||
self.litellm_dailygatewayrequests = table
|
||||
|
||||
async def __aenter__(self) -> "FakeBatcher":
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *args: object) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
class FakeDB:
|
||||
def __init__(self, table: FakeTable) -> None:
|
||||
self._table = table
|
||||
def __init__(self) -> None:
|
||||
self.statements: list[tuple[str, tuple[object, ...]]] = []
|
||||
|
||||
def batch_(self) -> FakeBatcher:
|
||||
return FakeBatcher(self._table)
|
||||
async def execute_raw(self, query: str, *args: object) -> int:
|
||||
self.statements.append((query, args))
|
||||
return len(args) // 5
|
||||
|
||||
|
||||
class FakePrismaClient:
|
||||
def __init__(self) -> None:
|
||||
self.table = FakeTable()
|
||||
self.db = FakeDB(self.table)
|
||||
self.db = FakeDB()
|
||||
|
||||
|
||||
def test_commit_upserts_one_incrementing_row_per_key():
|
||||
def _rows_written(client: FakePrismaClient) -> list[tuple[object, ...]]:
|
||||
"""Every (date, category, route, successful, failed) tuple the database received, in statement order."""
|
||||
return [params[i : i + 5] for _, params in client.db.statements for i in range(0, len(params), 5)]
|
||||
|
||||
|
||||
def test_commit_increments_with_a_single_statement_for_the_whole_snapshot():
|
||||
"""One statement per flush is the whole point: the previous per-key upsert cost
|
||||
the primary (workers x routes) statements per interval."""
|
||||
client = FakePrismaClient()
|
||||
snapshot = {
|
||||
GatewayRequestKey(date="2026-08-01", category="llm", route=route): (
|
||||
GatewayRequestCounts(successful_requests=7, failed_requests=2)
|
||||
)
|
||||
for route in ("/chat/completions", "/embeddings", "/responses", "/v1/messages", "/mcp")
|
||||
}
|
||||
|
||||
asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot=snapshot))
|
||||
|
||||
assert len(client.db.statements) == 1
|
||||
sql, params = client.db.statements[0]
|
||||
assert sql.count("ON CONFLICT") == 1
|
||||
assert sql.count("(NOW() AT TIME ZONE 'UTC'))") == 5
|
||||
assert len(params) == 25
|
||||
|
||||
|
||||
def test_commit_sql_adds_to_the_existing_row_instead_of_replacing_it():
|
||||
"""A worker only knows its own share; the SQL must add EXCLUDED onto the stored count."""
|
||||
client = FakePrismaClient()
|
||||
snapshot = {
|
||||
GatewayRequestKey(date="2026-08-01", category="llm", route="/chat/completions"): (
|
||||
|
|
@ -126,20 +136,36 @@ def test_commit_upserts_one_incrementing_row_per_key():
|
|||
|
||||
asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot=snapshot))
|
||||
|
||||
assert len(client.table.upserts) == 1
|
||||
written = client.table.upserts[0]
|
||||
assert written["where"] == {
|
||||
"date_category_route": {
|
||||
"date": "2026-08-01",
|
||||
"category": "llm",
|
||||
"route": "/chat/completions",
|
||||
}
|
||||
sql, params = client.db.statements[0]
|
||||
assert 'INSERT INTO "LiteLLM_DailyGatewayRequests"' in sql
|
||||
assert 'ON CONFLICT ("date", "category", "route") DO UPDATE SET' in sql
|
||||
assert (
|
||||
'"successful_requests" = "LiteLLM_DailyGatewayRequests"."successful_requests" + EXCLUDED."successful_requests"'
|
||||
in sql
|
||||
)
|
||||
assert '"failed_requests" = "LiteLLM_DailyGatewayRequests"."failed_requests" + EXCLUDED."failed_requests"' in sql
|
||||
assert params == ("2026-08-01", "llm", "/chat/completions", 7, 2)
|
||||
|
||||
|
||||
def test_commit_placeholders_line_up_with_params():
|
||||
"""$n positions are generated per row; a drift here silently swaps a route for a count."""
|
||||
client = FakePrismaClient()
|
||||
snapshot = {
|
||||
GatewayRequestKey(date="2026-08-01", category="llm", route="/chat/completions"): (
|
||||
GatewayRequestCounts(successful_requests=1, failed_requests=0)
|
||||
),
|
||||
GatewayRequestKey(date="2026-08-01", category="mcp", route="/mcp"): (
|
||||
GatewayRequestCounts(successful_requests=0, failed_requests=3)
|
||||
),
|
||||
}
|
||||
assert written["data"]["update"] == {
|
||||
"successful_requests": {"increment": 7},
|
||||
"failed_requests": {"increment": 2},
|
||||
}
|
||||
assert written["data"]["create"]["successful_requests"] == 7
|
||||
|
||||
asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot=snapshot))
|
||||
|
||||
sql, params = client.db.statements[0]
|
||||
assert "($1::text, $2::text, $3::text, $4::bigint, $5::bigint," in sql
|
||||
assert "($6::text, $7::text, $8::text, $9::bigint, $10::bigint," in sql
|
||||
assert "$11" not in sql
|
||||
assert params == ("2026-08-01", "llm", "/chat/completions", 1, 0, "2026-08-01", "mcp", "/mcp", 0, 3)
|
||||
|
||||
|
||||
def test_commit_is_deterministically_ordered():
|
||||
|
|
@ -154,17 +180,14 @@ def test_commit_is_deterministically_ordered():
|
|||
|
||||
asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot=snapshot))
|
||||
|
||||
written_order = [
|
||||
(row["where"]["date_category_route"]["date"], row["where"]["date_category_route"]["category"])
|
||||
for row in client.table.upserts
|
||||
]
|
||||
written_order = [(row[0], row[1]) for row in _rows_written(client)]
|
||||
assert written_order == [("2026-08-01", "llm"), ("2026-08-01", "mcp"), ("2026-08-02", "llm")]
|
||||
|
||||
|
||||
def test_commit_skips_the_database_entirely_when_nothing_accumulated():
|
||||
client = FakePrismaClient()
|
||||
asyncio.run(commit_gateway_requests_to_db(prisma_client=client, snapshot={}))
|
||||
assert client.table.upserts == []
|
||||
assert client.db.statements == []
|
||||
|
||||
|
||||
# ── flush ─────────────────────────────────────────────────────────────────────
|
||||
|
|
@ -177,12 +200,12 @@ def test_flush_drains_and_commits():
|
|||
|
||||
asyncio.run(flush_gateway_requests(client, acc))
|
||||
|
||||
assert len(client.table.upserts) == 1
|
||||
assert len(client.db.statements) == 1
|
||||
assert acc.drain() == {}
|
||||
|
||||
|
||||
class ExplodingDB:
|
||||
def batch_(self):
|
||||
async def execute_raw(self, query: str, *args: object) -> int:
|
||||
raise RuntimeError("db gone")
|
||||
|
||||
|
||||
|
|
@ -208,10 +231,7 @@ def test_failed_flush_keeps_counts_for_the_next_attempt():
|
|||
client = FakePrismaClient()
|
||||
asyncio.run(flush_gateway_requests(client, acc))
|
||||
|
||||
assert client.table.upserts[0]["data"]["update"] == {
|
||||
"successful_requests": {"increment": 1},
|
||||
"failed_requests": {"increment": 1},
|
||||
}
|
||||
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 1, 1)]
|
||||
|
||||
|
||||
def test_restored_counts_merge_with_requests_recorded_meanwhile():
|
||||
|
|
@ -223,5 +243,272 @@ def test_restored_counts_merge_with_requests_recorded_meanwhile():
|
|||
client = FakePrismaClient()
|
||||
asyncio.run(flush_gateway_requests(client, acc))
|
||||
|
||||
assert len(client.table.upserts) == 1
|
||||
assert client.table.upserts[0]["data"]["update"]["successful_requests"] == {"increment": 2}
|
||||
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)]
|
||||
|
||||
|
||||
class ExplodingDBWithInFlightRequest:
|
||||
"""Fails the write after a request has been recorded while it was in flight."""
|
||||
|
||||
def __init__(self, accumulator: GatewayRequestAccumulator) -> None:
|
||||
self.accumulator = accumulator
|
||||
|
||||
async def execute_raw(self, query: str, *args: object) -> int:
|
||||
_record(self.accumulator, 500)
|
||||
raise RuntimeError("db gone")
|
||||
|
||||
|
||||
class ExplodingClientWithInFlightRequest:
|
||||
def __init__(self, accumulator: GatewayRequestAccumulator) -> None:
|
||||
self.db = ExplodingDBWithInFlightRequest(accumulator)
|
||||
|
||||
|
||||
def test_restore_keeps_requests_recorded_while_the_failed_write_was_in_flight():
|
||||
acc = GatewayRequestAccumulator()
|
||||
_record(acc, 200)
|
||||
asyncio.run(flush_gateway_requests(ExplodingClientWithInFlightRequest(acc), acc))
|
||||
|
||||
client = FakePrismaClient()
|
||||
asyncio.run(flush_gateway_requests(client, acc))
|
||||
|
||||
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 1, 1)]
|
||||
|
||||
|
||||
# ── redis buffer ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class FakeRedis:
|
||||
def __init__(self) -> None:
|
||||
self.lists: dict[str, list[str]] = {}
|
||||
|
||||
async def async_rpush(self, key: str, values: list[str]) -> int:
|
||||
self.lists.setdefault(key, []).extend(values)
|
||||
return len(self.lists[key])
|
||||
|
||||
async def async_lpop(self, key: str, count: int) -> list[str] | None:
|
||||
queue = self.lists.get(key, [])
|
||||
if not queue:
|
||||
return None
|
||||
popped, self.lists[key] = queue[:count], queue[count:]
|
||||
return popped
|
||||
|
||||
|
||||
class FakePodLock:
|
||||
def __init__(self, *, leader: bool) -> None:
|
||||
self.leader = leader
|
||||
self.held: list[str] = []
|
||||
self.released: list[str] = []
|
||||
|
||||
async def acquire_lock(self, cronjob_id: str) -> bool:
|
||||
self.held.append(cronjob_id)
|
||||
return self.leader
|
||||
|
||||
async def release_lock(self, cronjob_id: str) -> None:
|
||||
self.released.append(cronjob_id)
|
||||
|
||||
|
||||
class FakeLease:
|
||||
"""Redis-side view of the job lock: SET NX by pod id, re-entrant for the holder, freed only by release or TTL."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.holder: str | None = None
|
||||
|
||||
|
||||
class FakeLeasePodLock:
|
||||
def __init__(self, lease: FakeLease, pod_id: str) -> None:
|
||||
self.lease = lease
|
||||
self.pod_id = pod_id
|
||||
|
||||
async def acquire_lock(self, cronjob_id: str) -> bool:
|
||||
if self.lease.holder is None:
|
||||
self.lease.holder = self.pod_id
|
||||
return self.lease.holder == self.pod_id
|
||||
|
||||
async def release_lock(self, cronjob_id: str) -> None:
|
||||
if self.lease.holder == self.pod_id:
|
||||
self.lease.holder = None
|
||||
|
||||
|
||||
def _buffer(redis: FakeRedis, *, leader: bool) -> tuple[GatewayRequestRedisBuffer, FakePodLock]:
|
||||
lock = FakePodLock(leader=leader)
|
||||
return GatewayRequestRedisBuffer(redis_cache=redis, pod_lock_manager=lock), lock # pyright: ignore[reportArgumentType] # duck-typed fakes
|
||||
|
||||
|
||||
def test_non_leader_workers_push_to_redis_and_never_touch_the_database():
|
||||
redis = FakeRedis()
|
||||
client = FakePrismaClient()
|
||||
for _ in range(3):
|
||||
acc = GatewayRequestAccumulator()
|
||||
_record(acc, 200)
|
||||
buffer, _ = _buffer(redis, leader=False)
|
||||
asyncio.run(flush_gateway_requests(client, acc, buffer))
|
||||
|
||||
assert client.db.statements == []
|
||||
assert len(redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY]) == 3
|
||||
|
||||
|
||||
def test_leader_folds_every_workers_snapshot_into_one_statement():
|
||||
"""Fifty workers each flushing the same routes must cost the primary one statement, not fifty."""
|
||||
redis = FakeRedis()
|
||||
client = FakePrismaClient()
|
||||
for _ in range(50):
|
||||
acc = GatewayRequestAccumulator()
|
||||
_record(acc, 200)
|
||||
_record(acc, 500, route="/responses")
|
||||
buffer, _ = _buffer(redis, leader=False)
|
||||
asyncio.run(flush_gateway_requests(client, acc, buffer))
|
||||
|
||||
leader_acc = GatewayRequestAccumulator()
|
||||
_record(leader_acc, 200)
|
||||
leader, lock = _buffer(redis, leader=True)
|
||||
asyncio.run(flush_gateway_requests(client, leader_acc, leader))
|
||||
|
||||
assert len(client.db.statements) == 1
|
||||
assert _rows_written(client) == [
|
||||
(_today(), "llm", "/chat/completions", 51, 0),
|
||||
(_today(), "llm", "/responses", 0, 50),
|
||||
]
|
||||
assert redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == []
|
||||
assert lock.held == [GATEWAY_REQUESTS_JOB_NAME]
|
||||
assert lock.released == []
|
||||
|
||||
|
||||
def test_leader_keeps_the_lease_so_staggered_pods_cost_one_statement_per_interval():
|
||||
"""Pods flush on their own clocks; without the lease each one would win the lock in turn and commit alone."""
|
||||
redis = FakeRedis()
|
||||
client = FakePrismaClient()
|
||||
lease = FakeLease()
|
||||
pods = tuple(
|
||||
GatewayRequestRedisBuffer(redis_cache=redis, pod_lock_manager=FakeLeasePodLock(lease, f"pod-{i}")) # pyright: ignore[reportArgumentType] # duck-typed fakes
|
||||
for i in range(4)
|
||||
)
|
||||
|
||||
for _interval in range(3):
|
||||
for pod in pods:
|
||||
acc = GatewayRequestAccumulator()
|
||||
_record(acc, 200)
|
||||
asyncio.run(flush_gateway_requests(client, acc, pod))
|
||||
|
||||
assert lease.holder == "pod-0"
|
||||
assert len(client.db.statements) == 3
|
||||
assert [row[3] for row in _rows_written(client)] == [1, 4, 4]
|
||||
assert len(redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY]) == 3
|
||||
|
||||
|
||||
def test_leader_drains_a_backlog_deeper_than_one_capped_pop():
|
||||
"""More workers than MAX_REDIS_BUFFER_DEQUEUE_COUNT must not leave a growing tail queued behind the cap."""
|
||||
redis = FakeRedis()
|
||||
client = FakePrismaClient()
|
||||
workers = MAX_REDIS_BUFFER_DEQUEUE_COUNT * 2 + 1
|
||||
for _ in range(workers):
|
||||
acc = GatewayRequestAccumulator()
|
||||
_record(acc, 200)
|
||||
buffer, _ = _buffer(redis, leader=False)
|
||||
asyncio.run(flush_gateway_requests(client, acc, buffer))
|
||||
|
||||
leader, _ = _buffer(redis, leader=True)
|
||||
asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), leader))
|
||||
|
||||
assert len(client.db.statements) == 1
|
||||
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", workers, 0)]
|
||||
assert redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == []
|
||||
|
||||
|
||||
def test_leader_with_nothing_buffered_writes_nothing():
|
||||
redis = FakeRedis()
|
||||
client = FakePrismaClient()
|
||||
leader, lock = _buffer(redis, leader=True)
|
||||
|
||||
asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), leader))
|
||||
|
||||
assert client.db.statements == []
|
||||
assert lock.released == []
|
||||
|
||||
|
||||
def test_leader_requeues_to_redis_when_the_database_commit_fails():
|
||||
"""Counts popped from Redis are gone from every worker; a failed commit must put them back."""
|
||||
redis = FakeRedis()
|
||||
acc = GatewayRequestAccumulator()
|
||||
_record(acc, 200)
|
||||
_record(acc, 200)
|
||||
leader, lock = _buffer(redis, leader=True)
|
||||
|
||||
asyncio.run(flush_gateway_requests(ExplodingClient(), acc, leader))
|
||||
|
||||
assert len(redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY]) == 1
|
||||
assert lock.released == []
|
||||
assert acc.drain() == {}
|
||||
|
||||
client = FakePrismaClient()
|
||||
retry, _ = _buffer(redis, leader=True)
|
||||
asyncio.run(flush_gateway_requests(client, GatewayRequestAccumulator(), retry))
|
||||
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)]
|
||||
|
||||
|
||||
class ExplodingRedis(FakeRedis):
|
||||
async def async_rpush(self, key: str, values: list[str]) -> int:
|
||||
raise RuntimeError("redis gone")
|
||||
|
||||
|
||||
class UnreadableRedis(FakeRedis):
|
||||
async def async_lpop(self, key: str, count: int) -> list[str] | None:
|
||||
raise RuntimeError("redis gone mid-flush")
|
||||
|
||||
|
||||
class UnwritableRedis(FakeRedis):
|
||||
"""Pops succeed, pushes fail: a Redis that went read-only between the leader's pop and its re-queue."""
|
||||
|
||||
async def async_rpush(self, key: str, values: list[str]) -> int:
|
||||
raise RuntimeError("redis read-only")
|
||||
|
||||
|
||||
def test_leader_keeps_popped_counts_in_memory_when_both_the_database_and_the_requeue_fail():
|
||||
"""The pop removed the only copy; if Redis will not take it back the leader itself must carry it."""
|
||||
redis = FakeRedis()
|
||||
worker_acc = GatewayRequestAccumulator()
|
||||
_record(worker_acc, 200)
|
||||
_record(worker_acc, 200)
|
||||
worker, _ = _buffer(redis, leader=False)
|
||||
asyncio.run(flush_gateway_requests(FakePrismaClient(), worker_acc, worker))
|
||||
|
||||
degraded = UnwritableRedis()
|
||||
degraded.lists = redis.lists
|
||||
leader_acc = GatewayRequestAccumulator()
|
||||
leader, _ = _buffer(degraded, leader=True)
|
||||
asyncio.run(flush_gateway_requests(ExplodingClient(), leader_acc, leader))
|
||||
assert degraded.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY] == []
|
||||
|
||||
client = FakePrismaClient()
|
||||
retry, _ = _buffer(redis, leader=True)
|
||||
asyncio.run(flush_gateway_requests(client, leader_acc, retry))
|
||||
assert _rows_written(client) == [(_today(), "llm", "/chat/completions", 2, 0)]
|
||||
|
||||
|
||||
def test_leader_whose_redis_read_fails_leaves_the_pushed_rows_for_the_next_flush():
|
||||
"""The scheduler job must not raise, and nothing is popped so nothing needs restoring anywhere."""
|
||||
redis = UnreadableRedis()
|
||||
acc = GatewayRequestAccumulator()
|
||||
_record(acc, 200)
|
||||
client = FakePrismaClient()
|
||||
leader, _ = _buffer(redis, leader=True)
|
||||
|
||||
asyncio.run(flush_gateway_requests(client, acc, leader))
|
||||
|
||||
assert client.db.statements == []
|
||||
assert acc.drain() == {}
|
||||
assert len(redis.lists[REDIS_GATEWAY_REQUESTS_BUFFER_KEY]) == 1
|
||||
|
||||
|
||||
def test_failed_redis_push_keeps_counts_locally_for_the_next_flush():
|
||||
acc = GatewayRequestAccumulator()
|
||||
_record(acc, 200)
|
||||
_record(acc, 500)
|
||||
buffer, lock = _buffer(ExplodingRedis(), leader=True)
|
||||
|
||||
asyncio.run(flush_gateway_requests(FakePrismaClient(), acc, buffer))
|
||||
|
||||
assert lock.held == []
|
||||
assert acc.drain() == {
|
||||
GatewayRequestKey(date=_today(), category="llm", route="/chat/completions"): (
|
||||
GatewayRequestCounts(successful_requests=1, failed_requests=1)
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,13 +4,17 @@ Tests for cost tracking settings management endpoints.
|
|||
Tests the GET and PATCH endpoints for managing cost discount configuration.
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
from pydantic import ValidationError
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm._internal_context import pinned_billing_time
|
||||
from litellm.proxy._types import CostEstimateRequest
|
||||
from litellm.proxy.management_endpoints.cost_tracking_settings import router
|
||||
from litellm.proxy.proxy_server import app
|
||||
|
||||
|
|
@ -789,13 +793,13 @@ INPUT_TOKENS = 1000
|
|||
OUTPUT_TOKENS = 500
|
||||
|
||||
|
||||
def _router_pricing(**pricing: float) -> MagicMock:
|
||||
def _router_pricing(model: str = AN_UNDERLYING_MODEL, **pricing: float) -> MagicMock:
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_model_list.return_value = [
|
||||
{
|
||||
"model_name": AN_ALIAS,
|
||||
"litellm_params": {
|
||||
"model": AN_UNDERLYING_MODEL,
|
||||
"model": model,
|
||||
"custom_llm_provider": "openai",
|
||||
**pricing,
|
||||
},
|
||||
|
|
@ -811,9 +815,7 @@ async def _estimate(mock_router: MagicMock | None, model: str = AN_ALIAS, **over
|
|||
|
||||
request = CostEstimateRequest(
|
||||
model=model,
|
||||
input_tokens=INPUT_TOKENS,
|
||||
output_tokens=OUTPUT_TOKENS,
|
||||
**overrides,
|
||||
**{"input_tokens": INPUT_TOKENS, "output_tokens": OUTPUT_TOKENS, **overrides},
|
||||
)
|
||||
with patch( # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
"litellm.proxy.proxy_server.llm_router", mock_router
|
||||
|
|
@ -909,3 +911,299 @@ class TestEstimateCostPeriodTotals:
|
|||
assert response.cost_per_request == pytest.approx(0.0022)
|
||||
assert response.daily_margin_cost == pytest.approx(0.02)
|
||||
assert response.daily_cost == pytest.approx(0.22)
|
||||
|
||||
|
||||
CACHE_READ_TOKENS = 800
|
||||
CACHE_CREATION_TOKENS = 100
|
||||
REASONING_TOKENS = 200
|
||||
TEXT_INPUT_TOKENS = INPUT_TOKENS - CACHE_READ_TOKENS - CACHE_CREATION_TOKENS
|
||||
TEXT_OUTPUT_TOKENS = OUTPUT_TOKENS - REASONING_TOKENS
|
||||
|
||||
|
||||
async def _estimate_with_cache_and_reasoning(mock_router: MagicMock | None, model: str = AN_ALIAS, **overrides: int):
|
||||
return await _estimate(
|
||||
mock_router,
|
||||
model=model,
|
||||
cache_read_input_tokens=CACHE_READ_TOKENS,
|
||||
cache_creation_input_tokens=CACHE_CREATION_TOKENS,
|
||||
reasoning_tokens=REASONING_TOKENS,
|
||||
**overrides,
|
||||
)
|
||||
|
||||
|
||||
class TestEstimateCostCacheAndReasoningTokens:
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_mapped_model_bills_cache_and_reasoning_tokens_at_their_own_rates(self, monkeypatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
A_MAPPED_MODEL,
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"cache_read_input_token_cost": 3e-7,
|
||||
"cache_creation_input_token_cost": 3.75e-6,
|
||||
"output_cost_per_reasoning_token": 1e-5,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
|
||||
response = await _estimate_with_cache_and_reasoning(None, model=A_MAPPED_MODEL, num_requests_per_day=10)
|
||||
|
||||
assert response.cache_read_cost_per_request == pytest.approx(CACHE_READ_TOKENS * 3e-7)
|
||||
assert response.cache_creation_cost_per_request == pytest.approx(CACHE_CREATION_TOKENS * 3.75e-6)
|
||||
assert response.reasoning_cost_per_request == pytest.approx(REASONING_TOKENS * 1e-5)
|
||||
assert response.input_cost_per_request == pytest.approx(
|
||||
TEXT_INPUT_TOKENS * 3e-6 + CACHE_READ_TOKENS * 3e-7 + CACHE_CREATION_TOKENS * 3.75e-6
|
||||
)
|
||||
assert response.output_cost_per_request == pytest.approx(TEXT_OUTPUT_TOKENS * 15e-6 + REASONING_TOKENS * 1e-5)
|
||||
assert response.cost_per_request == pytest.approx(
|
||||
response.input_cost_per_request + response.output_cost_per_request
|
||||
)
|
||||
assert response.daily_cache_read_cost == pytest.approx(10 * CACHE_READ_TOKENS * 3e-7)
|
||||
assert response.daily_cache_creation_cost == pytest.approx(10 * CACHE_CREATION_TOKENS * 3.75e-6)
|
||||
assert response.daily_reasoning_cost == pytest.approx(10 * REASONING_TOKENS * 1e-5)
|
||||
assert response.monthly_cache_read_cost is None
|
||||
assert response.cache_read_input_token_cost == pytest.approx(3e-7)
|
||||
assert response.cache_creation_input_token_cost == pytest.approx(3.75e-6)
|
||||
assert response.output_cost_per_reasoning_token == pytest.approx(1e-5)
|
||||
assert (
|
||||
response.cache_read_input_tokens,
|
||||
response.cache_creation_input_tokens,
|
||||
response.reasoning_tokens,
|
||||
) == (CACHE_READ_TOKENS, CACHE_CREATION_TOKENS, REASONING_TOKENS)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_model_without_cache_or_reasoning_prices_estimates_what_the_proxy_bills(self, monkeypatch):
|
||||
"""The cost calculator bills cache tokens of a cost-map model without cache prices at zero
|
||||
and its reasoning tokens at the output rate. The estimate reports those effective rates."""
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
A_MAPPED_MODEL,
|
||||
{"input_cost_per_token": 5e-6, "output_cost_per_token": 6e-6, "litellm_provider": "openai", "mode": "chat"},
|
||||
)
|
||||
|
||||
response = await _estimate_with_cache_and_reasoning(None, model=A_MAPPED_MODEL)
|
||||
|
||||
assert response.cache_read_cost_per_request == 0.0
|
||||
assert response.cache_creation_cost_per_request == 0.0
|
||||
assert response.reasoning_cost_per_request == pytest.approx(REASONING_TOKENS * 6e-6)
|
||||
assert response.input_cost_per_request == pytest.approx(TEXT_INPUT_TOKENS * 5e-6)
|
||||
assert response.cost_per_request == pytest.approx(TEXT_INPUT_TOKENS * 5e-6 + OUTPUT_TOKENS * 6e-6)
|
||||
assert response.cache_read_input_token_cost == 0.0
|
||||
assert response.cache_creation_input_token_cost == 0.0
|
||||
assert response.output_cost_per_reasoning_token == pytest.approx(6e-6)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_request_without_cache_or_reasoning_tokens_estimates_as_before(self, monkeypatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
A_MAPPED_MODEL,
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"cache_read_input_token_cost": 3e-7,
|
||||
"cache_creation_input_token_cost": 3.75e-6,
|
||||
"output_cost_per_reasoning_token": 1e-5,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
|
||||
response = await _estimate(None, model=A_MAPPED_MODEL, num_requests_per_day=10)
|
||||
|
||||
assert response.cost_per_request == pytest.approx(INPUT_TOKENS * 3e-6 + OUTPUT_TOKENS * 15e-6)
|
||||
assert response.cache_read_cost_per_request == 0.0
|
||||
assert response.cache_creation_cost_per_request == 0.0
|
||||
assert response.reasoning_cost_per_request == 0.0
|
||||
assert response.daily_cache_read_cost == 0.0
|
||||
assert response.daily_reasoning_cost == 0.0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_custom_priced_deployment_bills_cache_and_reasoning_tokens_from_its_flat_rates(self):
|
||||
response = await _estimate_with_cache_and_reasoning(
|
||||
_router_pricing(input_cost_per_token=1e-6, output_cost_per_token=2e-6, cache_read_input_token_cost=1e-7)
|
||||
)
|
||||
|
||||
assert response.cache_read_cost_per_request == pytest.approx(CACHE_READ_TOKENS * 1e-7)
|
||||
assert response.cache_creation_cost_per_request == pytest.approx(CACHE_CREATION_TOKENS * 1e-6)
|
||||
assert response.reasoning_cost_per_request == pytest.approx(REASONING_TOKENS * 2e-6)
|
||||
assert response.cost_per_request == pytest.approx(
|
||||
TEXT_INPUT_TOKENS * 1e-6 + CACHE_READ_TOKENS * 1e-7 + CACHE_CREATION_TOKENS * 1e-6 + OUTPUT_TOKENS * 2e-6
|
||||
)
|
||||
assert response.cache_read_input_token_cost == pytest.approx(1e-7)
|
||||
assert response.cache_creation_input_token_cost == pytest.approx(1e-6)
|
||||
assert response.output_cost_per_reasoning_token == pytest.approx(2e-6)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_custom_priced_deployment_of_a_mapped_model_inherits_its_built_in_cache_rates(self, monkeypatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
A_MAPPED_MODEL,
|
||||
{
|
||||
"input_cost_per_token": 5e-6,
|
||||
"output_cost_per_token": 6e-6,
|
||||
"cache_read_input_token_cost": 5e-7,
|
||||
"cache_creation_input_token_cost": 6.25e-6,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
|
||||
response = await _estimate_with_cache_and_reasoning(
|
||||
_router_pricing(model=A_MAPPED_MODEL, input_cost_per_token=1e-6, output_cost_per_token=2e-6)
|
||||
)
|
||||
|
||||
assert response.cache_read_cost_per_request == pytest.approx(CACHE_READ_TOKENS * 5e-7)
|
||||
assert response.cache_creation_cost_per_request == pytest.approx(CACHE_CREATION_TOKENS * 6.25e-6)
|
||||
assert response.input_cost_per_request == pytest.approx(
|
||||
TEXT_INPUT_TOKENS * 1e-6 + CACHE_READ_TOKENS * 5e-7 + CACHE_CREATION_TOKENS * 6.25e-6
|
||||
)
|
||||
assert response.cache_read_input_token_cost == pytest.approx(5e-7)
|
||||
assert response.cache_creation_input_token_cost == pytest.approx(6.25e-6)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_tiered_model_reports_the_rates_its_lines_were_billed_at(self, monkeypatch):
|
||||
"""Above a token tier the calculator bills every line at the tier's rate, so the reported
|
||||
rates must be the tier's too: each line equals its token count times the rate next to it."""
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
A_MAPPED_MODEL,
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"cache_read_input_token_cost": 3e-7,
|
||||
"cache_creation_input_token_cost": 3.75e-6,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-6,
|
||||
"output_cost_per_token_above_200k_tokens": 3e-5,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6e-7,
|
||||
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-6,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
|
||||
response = await _estimate(
|
||||
None,
|
||||
model=A_MAPPED_MODEL,
|
||||
input_tokens=250_000,
|
||||
cache_read_input_tokens=200_000,
|
||||
cache_creation_input_tokens=10_000,
|
||||
output_tokens=1_000,
|
||||
reasoning_tokens=200,
|
||||
)
|
||||
|
||||
assert response.input_cost_per_token == pytest.approx(6e-6)
|
||||
assert response.output_cost_per_token == pytest.approx(3e-5)
|
||||
assert response.cache_read_input_token_cost == pytest.approx(6e-7)
|
||||
assert response.cache_creation_input_token_cost == pytest.approx(7.5e-6)
|
||||
assert response.output_cost_per_reasoning_token == pytest.approx(3e-5)
|
||||
assert response.cache_read_cost_per_request == pytest.approx(200_000 * response.cache_read_input_token_cost)
|
||||
assert response.cache_creation_cost_per_request == pytest.approx(
|
||||
10_000 * response.cache_creation_input_token_cost
|
||||
)
|
||||
assert response.reasoning_cost_per_request == pytest.approx(200 * response.output_cost_per_reasoning_token)
|
||||
assert response.input_cost_per_request == pytest.approx(
|
||||
40_000 * response.input_cost_per_token
|
||||
+ response.cache_read_cost_per_request
|
||||
+ response.cache_creation_cost_per_request
|
||||
)
|
||||
assert response.output_cost_per_request == pytest.approx(1_000 * response.output_cost_per_token)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_quote_prices_its_totals_and_its_rates_at_the_same_moment(self, monkeypatch):
|
||||
"""The totals and the reported rates resolve off-peak pricing on separate paths. A quote
|
||||
taken as a window opens must not bill on one side of it and report rates from the other."""
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
A_MAPPED_MODEL,
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"off_peak_pricing": {
|
||||
"hours_utc": "02:00-03:00",
|
||||
"input_cost_per_token": 1e-6,
|
||||
"output_cost_per_token": 5e-6,
|
||||
},
|
||||
"litellm_provider": "openai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
|
||||
with pinned_billing_time(datetime(2026, 1, 1, 2, 30, tzinfo=timezone.utc)):
|
||||
response = await _estimate(None, model=A_MAPPED_MODEL)
|
||||
|
||||
assert response.input_cost_per_token == pytest.approx(1e-6)
|
||||
assert response.output_cost_per_token == pytest.approx(5e-6)
|
||||
assert response.input_cost_per_request == pytest.approx(INPUT_TOKENS * response.input_cost_per_token)
|
||||
assert response.output_cost_per_request == pytest.approx(OUTPUT_TOKENS * response.output_cost_per_token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_unrouted_model_reports_the_rates_of_the_provider_the_calculator_inferred(self, monkeypatch):
|
||||
"""The cost calculator infers a provider this endpoint never resolved, and the provider decides
|
||||
whether a tier threshold is inclusive. xai bills a request sitting exactly on the 200k threshold
|
||||
at the tier rate, so the reported rates have to be the tier's rather than the sub-tier base."""
|
||||
an_xai_model = "xai/tiered-model"
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
an_xai_model,
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"cache_read_input_token_cost": 3e-7,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-6,
|
||||
"output_cost_per_token_above_200k_tokens": 3e-5,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6e-7,
|
||||
"litellm_provider": "xai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
|
||||
response = await _estimate(
|
||||
None,
|
||||
model=an_xai_model,
|
||||
input_tokens=200_000,
|
||||
cache_read_input_tokens=100_000,
|
||||
output_tokens=1_000,
|
||||
)
|
||||
|
||||
assert response.input_cost_per_token == pytest.approx(6e-6)
|
||||
assert response.output_cost_per_token == pytest.approx(3e-5)
|
||||
assert response.cache_read_input_token_cost == pytest.approx(6e-7)
|
||||
assert response.cache_read_cost_per_request == pytest.approx(100_000 * response.cache_read_input_token_cost)
|
||||
assert response.input_cost_per_request == pytest.approx(
|
||||
100_000 * response.input_cost_per_token + response.cache_read_cost_per_request
|
||||
)
|
||||
assert response.output_cost_per_request == pytest.approx(1_000 * response.output_cost_per_token)
|
||||
|
||||
|
||||
class TestCostEstimateRequestTokenSubsets:
|
||||
def test_cache_tokens_beyond_the_input_tokens_are_rejected(self):
|
||||
with pytest.raises(ValidationError, match="cannot exceed input_tokens"):
|
||||
CostEstimateRequest(
|
||||
model=AN_ALIAS,
|
||||
input_tokens=INPUT_TOKENS,
|
||||
output_tokens=OUTPUT_TOKENS,
|
||||
cache_read_input_tokens=INPUT_TOKENS,
|
||||
cache_creation_input_tokens=1,
|
||||
)
|
||||
|
||||
def test_reasoning_tokens_beyond_the_output_tokens_are_rejected(self):
|
||||
with pytest.raises(ValidationError, match="cannot exceed output_tokens"):
|
||||
CostEstimateRequest(
|
||||
model=AN_ALIAS,
|
||||
input_tokens=INPUT_TOKENS,
|
||||
output_tokens=OUTPUT_TOKENS,
|
||||
reasoning_tokens=OUTPUT_TOKENS + 1,
|
||||
)
|
||||
|
||||
def test_the_endpoint_answers_422_when_cache_tokens_exceed_input_tokens(self):
|
||||
response = client.post(
|
||||
"/cost/estimate",
|
||||
headers={"Authorization": "Bearer sk-1234"},
|
||||
json={"model": AN_ALIAS, "input_tokens": 1000, "output_tokens": 100, "cache_read_input_tokens": 8000},
|
||||
)
|
||||
|
||||
assert response.status_code == 422
|
||||
assert "cannot exceed input_tokens" in response.text
|
||||
|
|
|
|||
|
|
@ -17975,6 +17975,32 @@ def test_key_request_blank_organization_id_is_unset():
|
|||
assert UpdateKeyRequest(key="sk-1", organization_id="org-1").organization_id == "org-1"
|
||||
|
||||
|
||||
def test_update_key_request_blank_team_id_is_not_a_team_change():
|
||||
from litellm.proxy._types import UpdateKeyRequest
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
is_different_team,
|
||||
)
|
||||
|
||||
blank = UpdateKeyRequest(key="sk-1", team_id="", key_alias="renamed")
|
||||
assert blank.team_id is None
|
||||
assert "team_id" not in blank.model_dump(exclude_unset=True)
|
||||
assert blank.model_dump(exclude_unset=True) == {"key": "sk-1", "key_alias": "renamed"}
|
||||
assert is_different_team(data=blank, existing_key_row=LiteLLM_VerificationToken(token="hashed")) is False
|
||||
assert (
|
||||
is_different_team(data=blank, existing_key_row=LiteLLM_VerificationToken(token="hashed", team_id="team-1"))
|
||||
is False
|
||||
)
|
||||
assert "team_id" in UpdateKeyRequest(key="sk-1", team_id=None).model_dump(exclude_unset=True)
|
||||
assert UpdateKeyRequest(key="sk-1", team_id="team-1").team_id == "team-1"
|
||||
assert (
|
||||
is_different_team(
|
||||
data=UpdateKeyRequest(key="sk-1", team_id="team-1"),
|
||||
existing_key_row=LiteLLM_VerificationToken(token="hashed"),
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_key_generation_check_blank_team_id_uses_personal_permissions(monkeypatch):
|
||||
"""key_generation_check with team_id="" must take the personal-key path instead
|
||||
of failing the team lookup with "Unable to find team object" (LIT-3925)."""
|
||||
|
|
|
|||
|
|
@ -2205,6 +2205,50 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name):
|
|||
assert update_call_kwargs.get("include", {}).get("object_permission") is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("endpoint_name", ["team_model_add", "team_model_delete"])
|
||||
async def test_team_model_add_delete_keep_model_aliases_in_team_cache(endpoint_name, monkeypatch):
|
||||
"""LIT-5858: Prisma only returns `litellm_model_table` when the `update` asks for it, so the refreshed
|
||||
cache entry lost the team's model aliases and JWT alias requests 403'd until the next DB read."""
|
||||
from litellm.proxy._types import TeamModelAddRequest, TeamModelDeleteRequest
|
||||
from litellm.proxy.auth.team_grants import team_model_aliases
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.team_endpoints import team_model_add, team_model_delete
|
||||
|
||||
columns = {"team_id": "team-1234", "models": ["gpt-4o", "openai/*"]}
|
||||
alias_table = {"id": 1, "model_aliases": '{"fast": "gpt-4o"}', "created_by": "admin", "updated_by": "admin"}
|
||||
|
||||
async def update(where, data, include=None):
|
||||
row = {**columns, "litellm_model_table": alias_table} if (include or {}).get("litellm_model_table") else columns
|
||||
return SimpleNamespace(team_id="team-1234", model_dump=lambda: row)
|
||||
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=SimpleNamespace(model_dump=lambda: columns))
|
||||
prisma_client.db.litellm_teamtable.update = AsyncMock(side_effect=update)
|
||||
prisma_client.db.execute_raw = AsyncMock(return_value=None)
|
||||
cache = UserApiKeyCache()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", None)
|
||||
|
||||
admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin")
|
||||
if endpoint_name == "team_model_add":
|
||||
await team_model_add(
|
||||
data=TeamModelAddRequest(team_id="team-1234", models=["team-byok-1"]),
|
||||
http_request=MagicMock(),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
else:
|
||||
await team_model_delete(
|
||||
data=TeamModelDeleteRequest(team_id="team-1234", models=["openai/*"]),
|
||||
http_request=MagicMock(),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
|
||||
cached_team = await cache.async_get_cache(key="team_id:team-1234", model_type=LiteLLM_TeamTableCachedObj)
|
||||
assert team_model_aliases(cached_team) == {"fast": "gpt-4o"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint_name",
|
||||
|
|
|
|||
|
|
@ -4,11 +4,12 @@ Pins covered:
|
|||
- ``get_current_spend``
|
||||
- ``increment_spend_counters``
|
||||
- ``_reconcile_budget_reservation_for_counter_update``
|
||||
- ``_increment_end_user_and_tag_spend_counters``
|
||||
- ``_increment_org_spend_counter``
|
||||
- ``_init_and_increment_unreserved_spend_counter``
|
||||
- ``_init_and_increment_spend_counter``
|
||||
- ``_init_and_increment_window_spend_counter``
|
||||
- ``_prepare_end_user_and_tag_spend_increments``
|
||||
- ``_prepare_org_spend_increment``
|
||||
- ``_prepare_unreserved_spend_counter_increment``
|
||||
- ``_prepare_spend_counter_increment``
|
||||
- ``_prepare_window_spend_counter_increment``
|
||||
- ``_apply_spend_counter_increments``
|
||||
- ``_ensure_spend_counter_initialized``
|
||||
- ``_get_source_cache_base_spend``
|
||||
- ``_ensure_window_spend_counter_initialized``
|
||||
|
|
@ -48,9 +49,7 @@ def _make_spend_counter_cache(
|
|||
cache.in_memory_cache.delete_cache = MagicMock()
|
||||
if with_redis:
|
||||
cache.redis_cache = MagicMock()
|
||||
cache.redis_cache.async_get_cache = AsyncMock(
|
||||
return_value=redis_get_value, side_effect=redis_get_side_effect
|
||||
)
|
||||
cache.redis_cache.async_get_cache = AsyncMock(return_value=redis_get_value, side_effect=redis_get_side_effect)
|
||||
cache.redis_cache.async_increment = AsyncMock(
|
||||
return_value=redis_increment_value,
|
||||
side_effect=redis_increment_side_effect,
|
||||
|
|
@ -58,6 +57,8 @@ def _make_spend_counter_cache(
|
|||
cache.redis_cache.async_delete_cache = AsyncMock()
|
||||
cache.redis_cache.async_set_cache = AsyncMock()
|
||||
cache.redis_cache.async_set_max = AsyncMock()
|
||||
cache.redis_cache.async_increment_pipeline = AsyncMock(return_value=None)
|
||||
cache.redis_cache.get_ttl = MagicMock(return_value=None)
|
||||
else:
|
||||
cache.redis_cache = None
|
||||
cache.async_increment_cache = AsyncMock(return_value=redis_increment_value)
|
||||
|
|
@ -70,9 +71,7 @@ def _make_spend_counter_cache(
|
|||
|
||||
def _make_user_api_key_cache(get_value=None, get_side_effect=None):
|
||||
cache = MagicMock()
|
||||
cache.async_get_cache = AsyncMock(
|
||||
return_value=get_value, side_effect=get_side_effect
|
||||
)
|
||||
cache.async_get_cache = AsyncMock(return_value=get_value, side_effect=get_side_effect)
|
||||
cache.async_set_cache_pipeline = AsyncMock()
|
||||
return cache
|
||||
|
||||
|
|
@ -109,9 +108,7 @@ async def test_get_current_spend_redis_error_falls_back_to_in_memory(monkeypatch
|
|||
)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
|
||||
result = await ps.get_current_spend(
|
||||
counter_key="spend:key:abc", fallback_spend=99.0
|
||||
)
|
||||
result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=99.0)
|
||||
assert result == 17.0
|
||||
|
||||
|
||||
|
|
@ -136,9 +133,7 @@ async def test_get_current_spend_floors_stale_low_counter_against_db(monkeypatch
|
|||
# the stale counter is repaired up to the authoritative DB value via a
|
||||
# monotonic set-max so other workers read the corrected total, and a
|
||||
# concurrent increment cannot be clobbered
|
||||
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(
|
||||
key="spend:key:abc", value=12.0
|
||||
)
|
||||
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key="spend:key:abc", value=12.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -169,9 +164,7 @@ async def test_get_current_spend_no_floor_without_max_budget(monkeypatch):
|
|||
from_db = AsyncMock(return_value=12.0)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", from_db)
|
||||
|
||||
result = await ps.get_current_spend(
|
||||
counter_key="spend:key:abc", fallback_spend=12.0
|
||||
)
|
||||
result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=12.0)
|
||||
|
||||
assert result == 2.0
|
||||
assert from_db.await_count == 0
|
||||
|
|
@ -210,12 +203,8 @@ async def test_get_current_spend_floor_caches_db_read(monkeypatch):
|
|||
from_db = AsyncMock(return_value=12.0)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", from_db)
|
||||
|
||||
first = await ps.get_current_spend(
|
||||
counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0
|
||||
)
|
||||
second = await ps.get_current_spend(
|
||||
counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0
|
||||
)
|
||||
first = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0)
|
||||
second = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=12.0, max_budget=10.0)
|
||||
|
||||
assert first == 12.0
|
||||
assert second == 12.0
|
||||
|
|
@ -336,9 +325,7 @@ async def test_get_current_spend_floors_window_against_spend_logs(monkeypatch):
|
|||
|
||||
assert result == 15.0
|
||||
assert wfsl.await_count == 1
|
||||
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(
|
||||
key=counter_key, value=15.0
|
||||
)
|
||||
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key=counter_key, value=15.0)
|
||||
|
||||
|
||||
def _make_window_spend_prisma(row=None, spend_logs_total=0.0):
|
||||
|
|
@ -379,9 +366,7 @@ async def test_get_current_spend_floors_window_against_maintained_row(monkeypatc
|
|||
|
||||
assert result == 15.0
|
||||
fake_prisma.db.litellm_spendlogs.group_by.assert_not_awaited()
|
||||
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(
|
||||
key=counter_key, value=15.0
|
||||
)
|
||||
fake_cache.redis_cache.async_set_max.assert_awaited_once_with(key=counter_key, value=15.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -393,9 +378,7 @@ async def test_get_current_spend_floors_window_against_logs_when_row_stale(monke
|
|||
|
||||
window_start = datetime(2026, 1, 8, tzinfo=timezone.utc)
|
||||
fake_prisma = _make_window_spend_prisma(
|
||||
row=SimpleNamespace(
|
||||
window_start=window_start - timedelta(days=7), spend=999.0
|
||||
),
|
||||
row=SimpleNamespace(window_start=window_start - timedelta(days=7), spend=999.0),
|
||||
spend_logs_total=15.0,
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=2.0)
|
||||
|
|
@ -423,21 +406,13 @@ async def test_get_current_spend_fail_closed_rejects_when_unverifiable(monkeypat
|
|||
rather than admitted on an unverifiable budget."""
|
||||
from fastapi import HTTPException
|
||||
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_side_effect=RuntimeError("redis down")
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_side_effect=RuntimeError("redis down"))
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(
|
||||
ps, "general_settings", {"fail_closed_budget_enforcement": True}
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
)
|
||||
monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await ps.get_current_spend(
|
||||
counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0
|
||||
)
|
||||
await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0)
|
||||
assert exc.value.status_code == 503
|
||||
|
||||
|
||||
|
|
@ -445,18 +420,12 @@ async def test_get_current_spend_fail_closed_rejects_when_unverifiable(monkeypat
|
|||
async def test_get_current_spend_fail_closed_off_admits_when_unverifiable(monkeypatch):
|
||||
"""Default (flag off): an unverifiable read keeps the existing behavior and
|
||||
admits using the cached fallback — no new rejection."""
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_side_effect=RuntimeError("redis down")
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_side_effect=RuntimeError("redis down"))
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "general_settings", {})
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
|
||||
|
||||
result = await ps.get_current_spend(
|
||||
counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0
|
||||
)
|
||||
result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0)
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
|
|
@ -466,13 +435,9 @@ async def test_get_current_spend_fail_closed_admits_when_redis_verified(monkeypa
|
|||
authoritative, so an under-budget request is admitted normally."""
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=1.0)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(
|
||||
ps, "general_settings", {"fail_closed_budget_enforcement": True}
|
||||
)
|
||||
monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
|
||||
|
||||
result = await ps.get_current_spend(
|
||||
counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0
|
||||
)
|
||||
result = await ps.get_current_spend(counter_key="spend:key:abc", fallback_spend=1.0, max_budget=10.0)
|
||||
assert result == 1.0
|
||||
|
||||
|
||||
|
|
@ -481,16 +446,10 @@ async def test_get_current_spend_fail_closed_allows_authoritative_fallback(monke
|
|||
"""End-user/tag callers pass fallback_authoritative=True (their spend is
|
||||
loaded fresh from the DB in auth), so fail-closed does not reject them even
|
||||
when the counter path is unreadable."""
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_side_effect=RuntimeError("redis down")
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_side_effect=RuntimeError("redis down"))
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(
|
||||
ps, "general_settings", {"fail_closed_budget_enforcement": True}
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
)
|
||||
monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
|
||||
|
||||
result = await ps.get_current_spend(
|
||||
counter_key="spend:end_user:e1",
|
||||
|
|
@ -508,9 +467,7 @@ async def test_get_current_spend_strict_floors_when_fallback_also_stale(monkeypa
|
|||
re-checks the authoritative DB and enforces against it."""
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=0.00001)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(
|
||||
ps, "general_settings", {"fail_closed_budget_enforcement": True}
|
||||
)
|
||||
monkeypatch.setattr(ps, "general_settings", {"fail_closed_budget_enforcement": True})
|
||||
from_db = AsyncMock(return_value=0.5)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "from_db", from_db)
|
||||
|
||||
|
|
@ -532,9 +489,7 @@ async def test_get_current_spend_strict_floors_when_fallback_also_stale(monkeypa
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_spend_counters_increments_all_buckets(monkeypatch):
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_value=None, redis_increment_value=5.0
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=5.0)
|
||||
fake_user_cache = _make_user_api_key_cache(get_value=None)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
|
|
@ -543,9 +498,7 @@ async def test_increment_spend_counters_increments_all_buckets(monkeypatch):
|
|||
async def _fake_coalesced(**kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced)
|
||||
)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(side_effect=_fake_coalesced))
|
||||
|
||||
await ps.increment_spend_counters(
|
||||
token="hashed-tok",
|
||||
|
|
@ -554,25 +507,36 @@ async def test_increment_spend_counters_increments_all_buckets(monkeypatch):
|
|||
response_cost=5.0,
|
||||
)
|
||||
|
||||
pipeline = fake_cache.redis_cache.async_increment_pipeline
|
||||
pipeline.assert_awaited_once()
|
||||
increment_list = pipeline.await_args.kwargs["increment_list"]
|
||||
assert {op["key"] for op in increment_list} == {
|
||||
"spend:key:hashed-tok",
|
||||
"spend:team:t1",
|
||||
"spend:team_member:u1:t1",
|
||||
"spend:user:u1",
|
||||
}
|
||||
assert all(op["increment_value"] == 5.0 for op in increment_list)
|
||||
observed = {
|
||||
"redis_increment_called": fake_cache.redis_cache.async_increment.called,
|
||||
"increment_calls": fake_cache.redis_cache.async_increment.call_count,
|
||||
"pipeline_calls": pipeline.await_count,
|
||||
"user_cache_used": fake_user_cache.async_get_cache.called,
|
||||
}
|
||||
assert normalize(observed) == {
|
||||
"redis_increment_called": True,
|
||||
"increment_calls": 4,
|
||||
"redis_increment_called": False,
|
||||
"pipeline_calls": 1,
|
||||
"user_cache_used": True,
|
||||
}
|
||||
|
||||
|
||||
class _ConcurrencyProbe:
|
||||
"""Stand-in for redis_cache.async_increment that pins concurrency.
|
||||
"""Stand-in for redis_cache.async_get_cache that pins concurrency.
|
||||
|
||||
Each call registers itself as in-flight and blocks on ``release`` until the
|
||||
test lets it proceed. ``all_arrived`` fires once ``expected`` distinct scope
|
||||
increments are simultaneously suspended here, which can only happen if the
|
||||
per-scope increments are gathered rather than awaited one after another.
|
||||
Each warm-check read registers itself as in-flight and blocks on ``release``
|
||||
until the test lets it proceed. ``all_arrived`` fires once ``expected``
|
||||
distinct scope warm-checks are simultaneously suspended here, which can only
|
||||
happen if the per-scope prepares are gathered rather than awaited one after
|
||||
another.
|
||||
"""
|
||||
|
||||
def __init__(self, expected_concurrency: int):
|
||||
|
|
@ -581,36 +545,45 @@ class _ConcurrencyProbe:
|
|||
self.max_in_flight = 0
|
||||
self.all_arrived = asyncio.Event()
|
||||
self.release = asyncio.Event()
|
||||
self.values: dict[str, float] = {}
|
||||
self.keys: list[str] = []
|
||||
|
||||
async def async_increment(self, *, key, value, refresh_ttl=True):
|
||||
async def async_get_cache(self, *, key, **kwargs):
|
||||
self.in_flight += 1
|
||||
self.max_in_flight = max(self.max_in_flight, self.in_flight)
|
||||
self.keys.append(key)
|
||||
if self.in_flight >= self.expected:
|
||||
self.all_arrived.set()
|
||||
if not self.release.is_set():
|
||||
await self.release.wait()
|
||||
self.in_flight -= 1
|
||||
self.values[key] = self.values.get(key, 0.0) + value
|
||||
return self.values[key]
|
||||
return 1.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_spend_counters_runs_scopes_concurrently(monkeypatch):
|
||||
"""The six independent scopes (key, team, team_member, user, end_user+tags,
|
||||
org) must be incremented concurrently. The probe only fires once all six are
|
||||
suspended in async_increment at the same time, which is impossible if the
|
||||
org) must prepare their increments concurrently. The probe only fires once
|
||||
all eight warm-check reads (one per counter: 6 scopes + 2 tags) are
|
||||
suspended in async_get_cache at the same time, which is impossible if the
|
||||
awaits are chained sequentially."""
|
||||
probe = _ConcurrencyProbe(expected_concurrency=6)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=None)
|
||||
fake_cache.redis_cache.async_increment = probe.async_increment
|
||||
probe = _ConcurrencyProbe(expected_concurrency=8)
|
||||
fake_cache = _make_spend_counter_cache()
|
||||
fake_cache.redis_cache.async_get_cache = probe.async_get_cache
|
||||
recorded: dict[str, float] = {}
|
||||
|
||||
async def _record_pipeline(increment_list, **_):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
recorded[op["key"]] = recorded.get(op["key"], 0.0) + op["increment_value"]
|
||||
results.append(recorded[op["key"]])
|
||||
return results
|
||||
|
||||
fake_cache.redis_cache.async_increment_pipeline = AsyncMock(side_effect=_record_pipeline)
|
||||
fake_user_cache = _make_user_api_key_cache(get_value=None)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
|
||||
|
||||
task = asyncio.create_task(
|
||||
ps.increment_spend_counters(
|
||||
|
|
@ -630,16 +603,16 @@ async def test_increment_spend_counters_runs_scopes_concurrently(monkeypatch):
|
|||
probe.release.set()
|
||||
await task
|
||||
pytest.fail(
|
||||
"scope increments did not run concurrently; sequential awaits "
|
||||
f"detected (peak in-flight was {probe.max_in_flight}, expected 6)"
|
||||
"scope prepares did not run concurrently; sequential awaits "
|
||||
f"detected (peak in-flight was {probe.max_in_flight}, expected 8)"
|
||||
)
|
||||
|
||||
assert probe.in_flight == 6
|
||||
assert probe.max_in_flight == 6
|
||||
assert probe.in_flight == 8
|
||||
assert probe.max_in_flight == 8
|
||||
probe.release.set()
|
||||
await task
|
||||
|
||||
assert probe.values == {
|
||||
assert recorded == {
|
||||
"spend:key:hashed-tok": 5.0,
|
||||
"spend:team:t1": 5.0,
|
||||
"spend:team_member:u1:t1": 5.0,
|
||||
|
|
@ -659,26 +632,25 @@ async def test_increment_spend_counters_skips_reserved_counter_keys(monkeypatch)
|
|||
import litellm.proxy.spend_tracking.budget_reservation as br
|
||||
|
||||
reserved = {"spend:key:hashed-tok", "spend:org:org1"}
|
||||
monkeypatch.setattr(
|
||||
br, "get_reserved_counter_keys", MagicMock(return_value=set(reserved))
|
||||
)
|
||||
monkeypatch.setattr(br, "get_reserved_counter_keys", MagicMock(return_value=set(reserved)))
|
||||
monkeypatch.setattr(br, "reconcile_budget_reservation", AsyncMock())
|
||||
|
||||
recorded: dict[str, float] = {}
|
||||
|
||||
async def _record_increment(*, key, value, refresh_ttl=True):
|
||||
recorded[key] = recorded.get(key, 0.0) + value
|
||||
return recorded[key]
|
||||
async def _record_pipeline(increment_list, **_):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
recorded[op["key"]] = recorded.get(op["key"], 0.0) + op["increment_value"]
|
||||
results.append(recorded[op["key"]])
|
||||
return results
|
||||
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=None)
|
||||
fake_cache.redis_cache.async_increment = _record_increment
|
||||
fake_cache.redis_cache.async_increment_pipeline = AsyncMock(side_effect=_record_pipeline)
|
||||
fake_user_cache = _make_user_api_key_cache(get_value=None)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
|
||||
|
||||
reservation = {"finalized": False}
|
||||
await ps.increment_spend_counters(
|
||||
|
|
@ -708,27 +680,46 @@ async def test_increment_spend_counters_failing_scope_propagates_after_siblings_
|
|||
):
|
||||
"""A failure in one scope must propagate to the caller (so it can invalidate
|
||||
reserved counters) while every other scope still settles rather than being
|
||||
left as an orphaned background task, and the reservation is not finalized."""
|
||||
recorded: dict[str, float] = {}
|
||||
left as an orphaned background task, and the reservation is not finalized.
|
||||
The surviving scopes' increments are still applied in the single pipeline:
|
||||
dropping them would under-count spend, the unsafe direction for budget
|
||||
enforcement."""
|
||||
warmed_keys: list[str] = []
|
||||
|
||||
async def _increment(*, key, value, refresh_ttl=True):
|
||||
async def _warm_check(*, key, **kwargs):
|
||||
warmed_keys.append(key)
|
||||
if key == "spend:team:t1":
|
||||
raise RuntimeError("redis increment failed")
|
||||
recorded[key] = recorded.get(key, 0.0) + value
|
||||
return recorded[key]
|
||||
raise RuntimeError("redis get failed")
|
||||
return 1.0
|
||||
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=None)
|
||||
fake_cache.redis_cache.async_increment = _increment
|
||||
async def _reseed_fails(*, counter_key, **kwargs):
|
||||
if counter_key == "spend:team:t1":
|
||||
raise RuntimeError("reseed failed")
|
||||
|
||||
applied: dict[str, float] = {}
|
||||
|
||||
async def _record_pipeline(increment_list, **_):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
applied[op["key"]] = op["increment_value"]
|
||||
results.append(op["increment_value"])
|
||||
return results
|
||||
|
||||
fake_cache = _make_spend_counter_cache()
|
||||
fake_cache.redis_cache.async_get_cache = AsyncMock(side_effect=_warm_check)
|
||||
fake_cache.redis_cache.async_increment_pipeline = AsyncMock(side_effect=_record_pipeline)
|
||||
fake_user_cache = _make_user_api_key_cache(get_value=None)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
ps.SpendCounterReseed,
|
||||
"coalesced",
|
||||
AsyncMock(side_effect=_reseed_fails),
|
||||
)
|
||||
|
||||
reservation = {"finalized": False}
|
||||
with pytest.raises(RuntimeError, match="redis increment failed"):
|
||||
with pytest.raises(RuntimeError, match="reseed failed"):
|
||||
await ps.increment_spend_counters(
|
||||
token="hashed-tok",
|
||||
team_id="t1",
|
||||
|
|
@ -741,7 +732,19 @@ async def test_increment_spend_counters_failing_scope_propagates_after_siblings_
|
|||
)
|
||||
|
||||
assert reservation["finalized"] is False
|
||||
assert recorded == {
|
||||
# every sibling scope settled (its warm-check ran) before the error propagated
|
||||
assert set(warmed_keys) == {
|
||||
"spend:key:hashed-tok",
|
||||
"spend:team:t1",
|
||||
"spend:team_member:u1:t1",
|
||||
"spend:user:u1",
|
||||
"spend:end_user:eu1",
|
||||
"spend:tag:a",
|
||||
"spend:org:org1",
|
||||
}
|
||||
# the surviving scopes' increments were still applied, in one pipeline call
|
||||
fake_cache.redis_cache.async_increment_pipeline.assert_awaited_once()
|
||||
assert applied == {
|
||||
"spend:key:hashed-tok": 5.0,
|
||||
"spend:team_member:u1:t1": 5.0,
|
||||
"spend:user:u1": 5.0,
|
||||
|
|
@ -749,6 +752,7 @@ async def test_increment_spend_counters_failing_scope_propagates_after_siblings_
|
|||
"spend:tag:a": 5.0,
|
||||
"spend:org:org1": 5.0,
|
||||
}
|
||||
fake_cache.redis_cache.async_increment.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -772,6 +776,108 @@ async def test_increment_spend_counters_zero_cost_is_noop_finalizes_reservation(
|
|||
|
||||
assert reservation == {"finalized": True}
|
||||
assert fake_cache.redis_cache.async_increment.called is False
|
||||
fake_cache.redis_cache.async_increment_pipeline.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_spend_counters_pipelines_all_scopes_in_one_redis_call(
|
||||
monkeypatch,
|
||||
):
|
||||
"""Every scope's increment must go out in a single async_increment_pipeline
|
||||
call, not one INCRBYFLOAT round-trip per scope."""
|
||||
counter_cache = ps.DualCache()
|
||||
fake_redis = AsyncMock()
|
||||
fake_redis.async_get_cache = AsyncMock(return_value=1.0) # counters warm
|
||||
|
||||
async def _pipeline(increment_list, **_):
|
||||
return [1.5] * len(increment_list)
|
||||
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=_pipeline)
|
||||
fake_redis.async_increment = AsyncMock()
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
counter_cache.redis_cache = fake_redis
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", counter_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", ps.DualCache())
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
|
||||
await ps.increment_spend_counters(
|
||||
token="hashed",
|
||||
team_id="team-1",
|
||||
user_id="user-1",
|
||||
response_cost=0.5,
|
||||
org_id="org-1",
|
||||
end_user_id="eu-1",
|
||||
tags=["tag-a", "tag-b"],
|
||||
)
|
||||
|
||||
fake_redis.async_increment_pipeline.assert_awaited_once()
|
||||
assert fake_redis.async_increment.await_count == 0
|
||||
increment_list = fake_redis.async_increment_pipeline.await_args.kwargs["increment_list"]
|
||||
expected_keys = {
|
||||
"spend:key:hashed",
|
||||
"spend:team:team-1",
|
||||
"spend:team_member:user-1:team-1",
|
||||
"spend:user:user-1",
|
||||
"spend:end_user:eu-1",
|
||||
"spend:tag:tag-a",
|
||||
"spend:tag:tag-b",
|
||||
"spend:org:org-1",
|
||||
}
|
||||
assert {op["key"] for op in increment_list} == expected_keys
|
||||
assert all(op["increment_value"] == 0.5 for op in increment_list)
|
||||
for key in expected_keys:
|
||||
assert counter_cache.in_memory_cache.get_cache(key=key) == 1.5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_spend_counters_pipeline_failure_invalidates_all_counters(
|
||||
monkeypatch,
|
||||
):
|
||||
"""A failing pipeline must invalidate every pending counter so the next
|
||||
request reseeds from the DB (which already holds this request's cost)
|
||||
instead of trusting a value the write may have partially applied."""
|
||||
from redis.exceptions import MaxConnectionsError
|
||||
|
||||
counter_cache = ps.DualCache()
|
||||
pending_keys = (
|
||||
"spend:key:hashed",
|
||||
"spend:team:team-1",
|
||||
"spend:team_member:user-1:team-1",
|
||||
"spend:user:user-1",
|
||||
"spend:end_user:eu-1",
|
||||
"spend:tag:tag-a",
|
||||
"spend:tag:tag-b",
|
||||
"spend:org:org-1",
|
||||
)
|
||||
for key in pending_keys:
|
||||
counter_cache.in_memory_cache.set_cache(key=key, value=1.0)
|
||||
fake_redis = AsyncMock()
|
||||
fake_redis.async_get_cache = AsyncMock(return_value=1.0) # counters warm
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=MaxConnectionsError())
|
||||
fake_redis.async_increment = AsyncMock()
|
||||
fake_redis.async_delete_cache = AsyncMock()
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
counter_cache.redis_cache = fake_redis
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", counter_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", ps.DualCache())
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
|
||||
with pytest.raises(MaxConnectionsError):
|
||||
await ps.increment_spend_counters(
|
||||
token="hashed",
|
||||
team_id="team-1",
|
||||
user_id="user-1",
|
||||
response_cost=0.5,
|
||||
org_id="org-1",
|
||||
end_user_id="eu-1",
|
||||
tags=["tag-a", "tag-b"],
|
||||
)
|
||||
|
||||
assert fake_redis.async_increment.await_count == 0
|
||||
deleted_keys = {call.kwargs["key"] for call in fake_redis.async_delete_cache.await_args_list}
|
||||
assert deleted_keys == set(pending_keys)
|
||||
for key in pending_keys:
|
||||
assert counter_cache.in_memory_cache.get_cache(key=key) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -781,9 +887,7 @@ async def test_increment_spend_counters_zero_cost_is_noop_finalizes_reservation(
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reconcile_budget_reservation_for_counter_update_returns_empty_set_when_none():
|
||||
result = await ps._reconcile_budget_reservation_for_counter_update(
|
||||
budget_reservation=None, response_cost=1.0
|
||||
)
|
||||
result = await ps._reconcile_budget_reservation_for_counter_update(budget_reservation=None, response_cost=1.0)
|
||||
assert result == set()
|
||||
|
||||
|
||||
|
|
@ -818,179 +922,151 @@ async def test_reconcile_budget_reservation_for_counter_update_failure_invalidat
|
|||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _increment_end_user_and_tag_spend_counters
|
||||
# _prepare_end_user_and_tag_spend_increments
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_end_user_and_tag_spend_counters_increments_each_unique_tag(
|
||||
async def test_prepare_end_user_and_tag_spend_increments_returns_each_unique_tag(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_value=None, redis_increment_value=3.0
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=1.0)
|
||||
fake_user_cache = _make_user_api_key_cache()
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
|
||||
|
||||
await ps._increment_end_user_and_tag_spend_counters(
|
||||
pending = await ps._prepare_end_user_and_tag_spend_increments(
|
||||
end_user_id="eu1",
|
||||
tags=["a", "b", "a", "", None],
|
||||
response_cost=3.0,
|
||||
reserved_counter_keys=set(),
|
||||
)
|
||||
|
||||
observed = {
|
||||
"increment_calls": fake_cache.redis_cache.async_increment.call_count,
|
||||
"in_memory_set_calls": fake_cache.in_memory_cache.set_cache.call_count,
|
||||
"called": fake_cache.redis_cache.async_increment.called,
|
||||
}
|
||||
assert normalize(observed) == {
|
||||
"increment_calls": 3,
|
||||
"in_memory_set_calls": 3,
|
||||
"called": True,
|
||||
assert {item.counter_key for item in pending} == {
|
||||
"spend:end_user:eu1",
|
||||
"spend:tag:a",
|
||||
"spend:tag:b",
|
||||
}
|
||||
assert all(item.increment == 3.0 for item in pending)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_end_user_and_tag_spend_counters_no_end_user_no_tags_invalid_input_noop(
|
||||
async def test_prepare_end_user_and_tag_spend_increments_no_end_user_no_tags_invalid_input_noop(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_cache = _make_spend_counter_cache()
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
|
||||
await ps._increment_end_user_and_tag_spend_counters(
|
||||
pending = await ps._prepare_end_user_and_tag_spend_increments(
|
||||
end_user_id=None,
|
||||
tags=None,
|
||||
response_cost=1.0,
|
||||
reserved_counter_keys=set(),
|
||||
)
|
||||
|
||||
assert pending == ()
|
||||
assert fake_cache.redis_cache.async_increment.called is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _increment_org_spend_counter
|
||||
# _prepare_org_spend_increment
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_org_spend_counter_increments_when_org_present(monkeypatch):
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_value=None, redis_increment_value=10.0
|
||||
)
|
||||
async def test_prepare_org_spend_increment_returns_pending_when_org_present(monkeypatch):
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=1.0)
|
||||
fake_user_cache = _make_user_api_key_cache()
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
|
||||
|
||||
await ps._increment_org_spend_counter(
|
||||
pending = await ps._prepare_org_spend_increment(
|
||||
org_id="org-1",
|
||||
response_cost=10.0,
|
||||
reserved_counter_keys=set(),
|
||||
)
|
||||
|
||||
observed = {
|
||||
"increment_called": fake_cache.redis_cache.async_increment.called,
|
||||
"increment_calls": fake_cache.redis_cache.async_increment.call_count,
|
||||
"counter_key_arg": fake_cache.redis_cache.async_increment.call_args.kwargs[
|
||||
"key"
|
||||
],
|
||||
}
|
||||
assert normalize(observed) == {
|
||||
"increment_called": True,
|
||||
"increment_calls": 1,
|
||||
"counter_key_arg": "spend:org:org-1",
|
||||
}
|
||||
assert len(pending) == 1
|
||||
assert pending[0].counter_key == "spend:org:org-1"
|
||||
assert pending[0].increment == 10.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_increment_org_spend_counter_no_org_is_noop_invalid_id(monkeypatch):
|
||||
async def test_prepare_org_spend_increment_no_org_is_noop_invalid_id(monkeypatch):
|
||||
fake_cache = _make_spend_counter_cache()
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
|
||||
await ps._increment_org_spend_counter(
|
||||
pending = await ps._prepare_org_spend_increment(
|
||||
org_id=None,
|
||||
response_cost=1.0,
|
||||
reserved_counter_keys=set(),
|
||||
)
|
||||
|
||||
assert pending == ()
|
||||
assert fake_cache.redis_cache.async_increment.called is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _init_and_increment_unreserved_spend_counter
|
||||
# _prepare_unreserved_spend_counter_increment
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_and_increment_unreserved_spend_counter_skips_reserved_keys(
|
||||
async def test_prepare_unreserved_spend_counter_increment_skips_reserved_keys(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_cache = _make_spend_counter_cache()
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
|
||||
await ps._init_and_increment_unreserved_spend_counter(
|
||||
pending = await ps._prepare_unreserved_spend_counter_increment(
|
||||
counter_key="spend:tag:x",
|
||||
source_cache_key="tag:x",
|
||||
increment=1.0,
|
||||
reserved_counter_keys={"spend:tag:x"},
|
||||
)
|
||||
|
||||
assert pending is None
|
||||
assert fake_cache.redis_cache.async_increment.called is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_and_increment_unreserved_spend_counter_proceeds_when_not_reserved(
|
||||
async def test_prepare_unreserved_spend_counter_increment_proceeds_when_not_reserved(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_value=None, redis_increment_value=2.0
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=None)
|
||||
fake_user_cache = _make_user_api_key_cache()
|
||||
reseed = AsyncMock(return_value=None)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", reseed)
|
||||
|
||||
await ps._init_and_increment_unreserved_spend_counter(
|
||||
pending = await ps._prepare_unreserved_spend_counter_increment(
|
||||
counter_key="spend:tag:y",
|
||||
source_cache_key="tag:y",
|
||||
increment=2.0,
|
||||
reserved_counter_keys=set(),
|
||||
)
|
||||
|
||||
observed = {
|
||||
"increment_called": fake_cache.redis_cache.async_increment.called,
|
||||
"redis_get_called": fake_cache.redis_cache.async_get_cache.called,
|
||||
"reseed_consulted": True,
|
||||
}
|
||||
assert observed == {
|
||||
"increment_called": True,
|
||||
"redis_get_called": True,
|
||||
"reseed_consulted": True,
|
||||
}
|
||||
assert pending is not None
|
||||
assert pending.counter_key == "spend:tag:y"
|
||||
assert pending.increment == 2.0
|
||||
assert fake_cache.redis_cache.async_get_cache.called is True
|
||||
assert reseed.called is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _init_and_increment_spend_counter
|
||||
# _prepare_spend_counter_increment
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypatch):
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_value=11.0, redis_increment_value=14.0
|
||||
)
|
||||
async def test_prepare_spend_counter_increment_warm_cache_skips_reseed(monkeypatch):
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=11.0)
|
||||
fake_user_cache = _make_user_api_key_cache()
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
|
|
@ -998,12 +1074,14 @@ async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypa
|
|||
reseed = AsyncMock(return_value=None)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", reseed)
|
||||
|
||||
await ps._init_and_increment_spend_counter(
|
||||
pending = await ps._prepare_spend_counter_increment(
|
||||
counter_key="spend:key:k",
|
||||
source_cache_key="k",
|
||||
increment=3.0,
|
||||
)
|
||||
|
||||
assert pending.counter_key == "spend:key:k"
|
||||
assert pending.increment == 3.0
|
||||
observed = {
|
||||
"reseed_called": reseed.called,
|
||||
"increment_called": fake_cache.redis_cache.async_increment.called,
|
||||
|
|
@ -1011,23 +1089,21 @@ async def test_init_and_increment_spend_counter_warm_cache_skips_reseed(monkeypa
|
|||
}
|
||||
assert normalize(observed) == {
|
||||
"reseed_called": False,
|
||||
"increment_called": True,
|
||||
"increment_called": False,
|
||||
"in_memory_seeded_from_redis": True,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _init_and_increment_window_spend_counter
|
||||
# _prepare_window_spend_counter_increment
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_and_increment_window_spend_counter_increments_when_initialized(
|
||||
async def test_prepare_window_spend_counter_increment_returns_pending_when_initialized(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_value=0.0, redis_increment_value=5.0
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=0.0)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -1036,7 +1112,7 @@ async def test_init_and_increment_window_spend_counter_increments_when_initializ
|
|||
AsyncMock(return_value=0.0),
|
||||
)
|
||||
|
||||
await ps._init_and_increment_window_spend_counter(
|
||||
pending = await ps._prepare_window_spend_counter_increment(
|
||||
counter_key="spend:key:k:window:1d",
|
||||
entity_type="Key",
|
||||
entity_id="k",
|
||||
|
|
@ -1045,26 +1121,19 @@ async def test_init_and_increment_window_spend_counter_increments_when_initializ
|
|||
increment=5.0,
|
||||
)
|
||||
|
||||
observed = {
|
||||
"redis_increment_called": fake_cache.redis_cache.async_increment.called,
|
||||
"increment_calls": fake_cache.redis_cache.async_increment.call_count,
|
||||
"in_memory_set_calls": fake_cache.in_memory_cache.set_cache.call_count,
|
||||
}
|
||||
assert normalize(observed) == {
|
||||
"redis_increment_called": True,
|
||||
"increment_calls": 1,
|
||||
"in_memory_set_calls": 2,
|
||||
}
|
||||
assert pending is not None
|
||||
assert pending.counter_key == "spend:key:k:window:1d"
|
||||
assert pending.increment == 5.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_and_increment_window_spend_counter_missing_window_start_invalid_skips(
|
||||
async def test_prepare_window_spend_counter_increment_missing_window_start_invalid_skips(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_cache = _make_spend_counter_cache()
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
|
||||
await ps._init_and_increment_window_spend_counter(
|
||||
pending = await ps._prepare_window_spend_counter_increment(
|
||||
counter_key="spend:key:k:window:1d",
|
||||
entity_type="Key",
|
||||
entity_id="k",
|
||||
|
|
@ -1073,6 +1142,7 @@ async def test_init_and_increment_window_spend_counter_missing_window_start_inva
|
|||
increment=5.0,
|
||||
)
|
||||
|
||||
assert pending is None
|
||||
assert fake_cache.redis_cache.async_increment.called is False
|
||||
|
||||
|
||||
|
|
@ -1114,16 +1184,12 @@ async def test_ensure_spend_counter_initialized_warm_skips_reseed_and_source(
|
|||
async def test_ensure_spend_counter_initialized_cold_seeds_from_source_cache(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_get_value=None, redis_increment_value=7.0
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_get_value=None, redis_increment_value=7.0)
|
||||
fake_user_cache = _make_user_api_key_cache(get_value={"spend": 7.0})
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
monkeypatch.setattr(ps, "prisma_client", None)
|
||||
monkeypatch.setattr(
|
||||
ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None)
|
||||
)
|
||||
monkeypatch.setattr(ps.SpendCounterReseed, "coalesced", AsyncMock(return_value=None))
|
||||
|
||||
await ps._ensure_spend_counter_initialized(
|
||||
counter_key="spend:user:u",
|
||||
|
|
@ -1163,9 +1229,7 @@ async def test_get_source_cache_base_spend_reads_first_hit_from_list(monkeypatch
|
|||
fake_user_cache.async_get_cache = AsyncMock(side_effect=_get)
|
||||
monkeypatch.setattr(ps, "user_api_key_cache", fake_user_cache)
|
||||
|
||||
result = await ps._get_source_cache_base_spend(
|
||||
source_cache_key=["miss", "hit-obj", "miss2"]
|
||||
)
|
||||
result = await ps._get_source_cache_base_spend(source_cache_key=["miss", "hit-obj", "miss2"])
|
||||
|
||||
observed = {
|
||||
"result": result,
|
||||
|
|
@ -1294,9 +1358,7 @@ async def test_increment_spend_counter_cache_redis_path_returns_new_value(monkey
|
|||
fake_cache = _make_spend_counter_cache(redis_increment_value=44.0)
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
|
||||
result = await ps._increment_spend_counter_cache(
|
||||
counter_key="spend:key:k", increment=4.0
|
||||
)
|
||||
result = await ps._increment_spend_counter_cache(counter_key="spend:key:k", increment=4.0)
|
||||
|
||||
observed = {
|
||||
"result": result,
|
||||
|
|
@ -1314,15 +1376,11 @@ async def test_increment_spend_counter_cache_redis_path_returns_new_value(monkey
|
|||
async def test_increment_spend_counter_cache_redis_error_raises_and_invalidates(
|
||||
monkeypatch,
|
||||
):
|
||||
fake_cache = _make_spend_counter_cache(
|
||||
redis_increment_side_effect=RuntimeError("incr fail")
|
||||
)
|
||||
fake_cache = _make_spend_counter_cache(redis_increment_side_effect=RuntimeError("incr fail"))
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
await ps._increment_spend_counter_cache(
|
||||
counter_key="spend:key:k", increment=1.0
|
||||
)
|
||||
await ps._increment_spend_counter_cache(counter_key="spend:key:k", increment=1.0)
|
||||
|
||||
assert fake_cache.in_memory_cache.delete_cache.called is True
|
||||
assert fake_cache.redis_cache.async_delete_cache.called is True
|
||||
|
|
@ -1343,9 +1401,7 @@ async def test_invalidate_spend_counter_deletes_in_memory_and_redis(monkeypatch)
|
|||
observed = {
|
||||
"in_memory_delete_called": fake_cache.in_memory_cache.delete_cache.called,
|
||||
"redis_delete_called": fake_cache.redis_cache.async_delete_cache.called,
|
||||
"delete_args_key": fake_cache.redis_cache.async_delete_cache.call_args.kwargs[
|
||||
"key"
|
||||
],
|
||||
"delete_args_key": fake_cache.redis_cache.async_delete_cache.call_args.kwargs["key"],
|
||||
}
|
||||
assert normalize(observed) == {
|
||||
"in_memory_delete_called": True,
|
||||
|
|
@ -1357,9 +1413,7 @@ async def test_invalidate_spend_counter_deletes_in_memory_and_redis(monkeypatch)
|
|||
@pytest.mark.asyncio
|
||||
async def test_invalidate_spend_counter_swallows_redis_failure_no_raise(monkeypatch):
|
||||
fake_cache = _make_spend_counter_cache()
|
||||
fake_cache.redis_cache.async_delete_cache = AsyncMock(
|
||||
side_effect=RuntimeError("redis down")
|
||||
)
|
||||
fake_cache.redis_cache.async_delete_cache = AsyncMock(side_effect=RuntimeError("redis down"))
|
||||
monkeypatch.setattr(ps, "spend_counter_cache", fake_cache)
|
||||
|
||||
await ps._invalidate_spend_counter(counter_key="spend:key:k")
|
||||
|
|
|
|||
|
|
@ -41,6 +41,15 @@ class _FlakyRedisCache:
|
|||
self._store[key] = float(value)
|
||||
return True
|
||||
|
||||
async def async_increment_pipeline(self, increment_list, **kwargs):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
results.append(await self.async_increment(op["key"], op["increment_value"]))
|
||||
return results
|
||||
|
||||
def get_ttl(self, **kwargs):
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_direct_increment_runs_when_reservation_reconcile_hits_redis_failure(
|
||||
|
|
|
|||
|
|
@ -2220,6 +2220,15 @@ class _ExpiringRedisCache:
|
|||
async def async_delete_cache(self, key: str, *args: object, **kwargs: object) -> None:
|
||||
self.store.pop(key, None)
|
||||
|
||||
async def async_increment_pipeline(self, increment_list, **kwargs):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
results.append(await self.async_increment(op["key"], op["increment_value"]))
|
||||
return results
|
||||
|
||||
def get_ttl(self, **kwargs) -> None:
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reconcile_after_redis_counter_expiry_keeps_request_cost_enforced(
|
||||
|
|
|
|||
|
|
@ -7749,7 +7749,7 @@ async def test_increment_spend_counters_team_and_member():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss():
|
||||
async def test_prepare_spend_counter_increment_reseeds_from_db_on_counter_miss():
|
||||
"""When the Redis counter is missing, the reseed path reads the
|
||||
authoritative spend from the DB (not a stale cache), so the next
|
||||
increment continues from the correct base value."""
|
||||
|
|
@ -7762,8 +7762,17 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
|
|||
recorded_increments.append({"key": key, "value": value, "ttl": ttl})
|
||||
return value
|
||||
|
||||
async def record_pipeline(increment_list, **kwargs):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
await record_increment(key=op["key"], value=op["increment_value"], ttl=op["ttl"])
|
||||
results.append(op["increment_value"])
|
||||
return results
|
||||
|
||||
fake_redis = AsyncMock()
|
||||
fake_redis.async_increment = AsyncMock(side_effect=record_increment)
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=record_pipeline)
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
fake_redis.async_get_cache = AsyncMock(return_value=None) # counter missing
|
||||
fake_redis.async_set_cache = AsyncMock(return_value=True) # SET NX wins
|
||||
counter_cache.redis_cache = fake_redis
|
||||
|
|
@ -7782,7 +7791,10 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
|
|||
stale_cache.in_memory_cache.set_cache(key="team_id:team-9", value=stale_team)
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
from litellm.proxy.proxy_server import _init_and_increment_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_spend_counter_increment,
|
||||
)
|
||||
|
||||
orig_user, orig_counter, orig_prisma = (
|
||||
ps.user_api_key_cache,
|
||||
|
|
@ -7793,11 +7805,12 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss(
|
|||
ps.spend_counter_cache = counter_cache
|
||||
ps.prisma_client = fake_prisma
|
||||
try:
|
||||
await _init_and_increment_spend_counter(
|
||||
pending = await _prepare_spend_counter_increment(
|
||||
counter_key="spend:team:team-9",
|
||||
source_cache_key="team_id:team-9",
|
||||
increment=1.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,))
|
||||
|
||||
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-9"})
|
||||
# Seed uses SET NX with db_spend (42) — cross-pod safe, no INCR of 42.
|
||||
|
|
@ -7976,7 +7989,10 @@ async def test_reseed_spend_from_db_skips_window_variant_keys():
|
|||
@pytest.mark.asyncio
|
||||
async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_window_spend_counter_increment,
|
||||
)
|
||||
|
||||
counter_cache = DualCache()
|
||||
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
|
|
@ -7992,7 +8008,7 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
|||
ps.spend_counter_cache = counter_cache
|
||||
ps.prisma_client = fake_prisma
|
||||
try:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending = await _prepare_window_spend_counter_increment(
|
||||
counter_key="spend:key:key-window:window:1h",
|
||||
entity_type="Key",
|
||||
entity_id="key-window",
|
||||
|
|
@ -8000,6 +8016,7 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
|||
window_start=window_start,
|
||||
increment=0.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,) if pending is not None else ())
|
||||
|
||||
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
|
||||
by=["api_key"],
|
||||
|
|
@ -8015,7 +8032,10 @@ async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
|||
@pytest.mark.asyncio
|
||||
async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_spend_counter_increment,
|
||||
)
|
||||
|
||||
counter_cache = DualCache()
|
||||
counter_key = "spend:team:team-stale-local"
|
||||
|
|
@ -8037,6 +8057,15 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||||
|
||||
async def redis_increment_pipeline(increment_list, **_):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
results.append(await redis_increment(key=op["key"], value=op["increment_value"]))
|
||||
return results
|
||||
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=redis_increment_pipeline)
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
counter_cache.redis_cache = fake_redis
|
||||
|
||||
db_row = MagicMock()
|
||||
|
|
@ -8055,11 +8084,12 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
ps.prisma_client = fake_prisma
|
||||
ps.user_api_key_cache = DualCache()
|
||||
try:
|
||||
await _init_and_increment_spend_counter(
|
||||
pending = await _prepare_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
source_cache_key="team_id:team-stale-local",
|
||||
increment=1.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,))
|
||||
|
||||
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-stale-local"})
|
||||
# Seed via SET NX (42) + delta via INCRBYFLOAT (1.5) = 43.5.
|
||||
|
|
@ -8074,7 +8104,10 @@ async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
@pytest.mark.asyncio
|
||||
async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_window_spend_counter_increment,
|
||||
)
|
||||
|
||||
counter_cache = DualCache()
|
||||
counter_key = "spend:key:key-window-stale-local:window:1h"
|
||||
|
|
@ -8097,6 +8130,15 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||||
|
||||
async def redis_increment_pipeline(increment_list, **_):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
results.append(await redis_increment(key=op["key"], value=op["increment_value"]))
|
||||
return results
|
||||
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=redis_increment_pipeline)
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
counter_cache.redis_cache = fake_redis
|
||||
|
||||
fake_prisma = MagicMock()
|
||||
|
|
@ -8111,7 +8153,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
ps.spend_counter_cache = counter_cache
|
||||
ps.prisma_client = fake_prisma
|
||||
try:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending = await _prepare_window_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
entity_type="Key",
|
||||
entity_id="key-window-stale-local",
|
||||
|
|
@ -8119,6 +8161,7 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
window_start=window_start,
|
||||
increment=0.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,) if pending is not None else ())
|
||||
|
||||
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
|
||||
by=["api_key"],
|
||||
|
|
@ -8138,7 +8181,10 @@ async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
|||
@pytest.mark.asyncio
|
||||
async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||||
from litellm.proxy.proxy_server import (
|
||||
_apply_spend_counter_increments,
|
||||
_prepare_window_spend_counter_increment,
|
||||
)
|
||||
|
||||
counter_cache = DualCache()
|
||||
counter_key = "spend:key:key-window-concurrent-seed:window:1h"
|
||||
|
|
@ -8161,6 +8207,15 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
|
|||
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache)
|
||||
fake_redis.async_set_cache = AsyncMock(return_value=False)
|
||||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||||
|
||||
async def redis_increment_pipeline(increment_list, **_):
|
||||
results = []
|
||||
for op in increment_list:
|
||||
results.append(await redis_increment(key=op["key"], value=op["increment_value"]))
|
||||
return results
|
||||
|
||||
fake_redis.async_increment_pipeline = AsyncMock(side_effect=redis_increment_pipeline)
|
||||
fake_redis.get_ttl = MagicMock(return_value=None)
|
||||
counter_cache.redis_cache = fake_redis
|
||||
|
||||
fake_prisma = MagicMock()
|
||||
|
|
@ -8175,7 +8230,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
|
|||
ps.spend_counter_cache = counter_cache
|
||||
ps.prisma_client = fake_prisma
|
||||
try:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending = await _prepare_window_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
entity_type="Key",
|
||||
entity_id="key-window-concurrent-seed",
|
||||
|
|
@ -8183,6 +8238,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
|
|||
window_start=window_start,
|
||||
increment=0.5,
|
||||
)
|
||||
await _apply_spend_counter_increments(pending=(pending,) if pending is not None else ())
|
||||
|
||||
fake_redis.async_set_cache.assert_awaited_once_with(
|
||||
key=counter_key,
|
||||
|
|
@ -8199,7 +8255,7 @@ async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed()
|
|||
@pytest.mark.asyncio
|
||||
async def test_window_spend_counter_skips_invalid_window_start():
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||||
from litellm.proxy.proxy_server import _prepare_window_spend_counter_increment
|
||||
|
||||
counter_cache = DualCache()
|
||||
|
||||
|
|
@ -8208,7 +8264,7 @@ async def test_window_spend_counter_skips_invalid_window_start():
|
|||
orig_counter = ps.spend_counter_cache
|
||||
ps.spend_counter_cache = counter_cache
|
||||
try:
|
||||
await _init_and_increment_window_spend_counter(
|
||||
pending = await _prepare_window_spend_counter_increment(
|
||||
counter_key="spend:key:key-invalid-window:window:not-a-duration",
|
||||
entity_type="Key",
|
||||
entity_id="key-invalid-window",
|
||||
|
|
@ -8216,6 +8272,7 @@ async def test_window_spend_counter_skips_invalid_window_start():
|
|||
window_start=None,
|
||||
increment=0.5,
|
||||
)
|
||||
assert pending is None
|
||||
|
||||
assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-invalid-window:window:not-a-duration") is None
|
||||
finally:
|
||||
|
|
@ -8279,6 +8336,9 @@ async def test_increment_spend_counters_finalizes_after_unreserved_increments():
|
|||
async def assert_reservation_not_finalized_yet(**kwargs):
|
||||
assert budget_reservation["finalized"] is False
|
||||
incremented_counters.append(kwargs["counter_key"])
|
||||
return ps._PendingSpendIncrement(
|
||||
counter_key=kwargs["counter_key"], increment=kwargs["increment"]
|
||||
)
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
|
|
@ -8287,7 +8347,7 @@ async def test_increment_spend_counters_finalizes_after_unreserved_increments():
|
|||
ps.user_api_key_cache = DualCache()
|
||||
try:
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server._init_and_increment_spend_counter",
|
||||
"litellm.proxy.proxy_server._prepare_spend_counter_increment",
|
||||
new=AsyncMock(side_effect=assert_reservation_not_finalized_yet),
|
||||
):
|
||||
await increment_spend_counters(
|
||||
|
|
@ -8620,7 +8680,7 @@ async def test_get_current_spend_uses_db_zero_over_stale_fallback():
|
|||
async def test_concurrent_read_and_write_paths_share_one_db_query():
|
||||
"""
|
||||
The read path (`get_current_spend`) and the write path
|
||||
(`_init_and_increment_spend_counter`) both reseed cold counters from
|
||||
(`_prepare_spend_counter_increment`) both reseed cold counters from
|
||||
the DB. They must share the per-counter lock so a concurrent pre-call
|
||||
enforcement read and post-call increment for the same counter collapse
|
||||
to one DB query, not two.
|
||||
|
|
@ -8629,7 +8689,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query():
|
|||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.proxy_server import (
|
||||
_init_and_increment_spend_counter,
|
||||
_prepare_spend_counter_increment,
|
||||
get_current_spend,
|
||||
)
|
||||
|
||||
|
|
@ -8683,7 +8743,7 @@ async def test_concurrent_read_and_write_paths_share_one_db_query():
|
|||
try:
|
||||
results = await _asyncio.gather(
|
||||
get_current_spend(counter_key=counter_key, fallback_spend=0.0),
|
||||
_init_and_increment_spend_counter(
|
||||
_prepare_spend_counter_increment(
|
||||
counter_key=counter_key,
|
||||
source_cache_key="ignored",
|
||||
increment=1.5,
|
||||
|
|
|
|||
|
|
@ -8,11 +8,12 @@ import json
|
|||
from datetime import datetime, timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler
|
||||
from litellm.router import Router
|
||||
from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler, RoutingArgs
|
||||
|
||||
DEPLOYMENT_ID = "9876"
|
||||
KWARGS = {
|
||||
|
|
@ -58,9 +59,9 @@ def test_sync_embedding_latency_is_json_serializable():
|
|||
|
||||
latencies = _recorded_latencies(cache)
|
||||
assert latencies, "expected a latency entry to be recorded"
|
||||
assert all(
|
||||
not isinstance(value, timedelta) for value in latencies
|
||||
), f"raw timedelta leaked into latency list: {latencies}"
|
||||
assert all(not isinstance(value, timedelta) for value in latencies), (
|
||||
f"raw timedelta leaked into latency list: {latencies}"
|
||||
)
|
||||
assert latencies[-1] == pytest.approx(2.0)
|
||||
# the exact failure mode from production: redis cache sync json.dumps
|
||||
json.dumps({"latency": latencies})
|
||||
|
|
@ -84,9 +85,9 @@ async def test_async_embedding_latency_is_json_serializable():
|
|||
|
||||
latencies = _recorded_latencies(cache)
|
||||
assert latencies, "expected a latency entry to be recorded"
|
||||
assert all(
|
||||
not isinstance(value, timedelta) for value in latencies
|
||||
), f"raw timedelta leaked into latency list: {latencies}"
|
||||
assert all(not isinstance(value, timedelta) for value in latencies), (
|
||||
f"raw timedelta leaked into latency list: {latencies}"
|
||||
)
|
||||
assert latencies[-1] == pytest.approx(3.0)
|
||||
json.dumps({"latency": latencies})
|
||||
|
||||
|
|
@ -292,6 +293,85 @@ async def test_streaming_routing_ignores_per_token_ttft_samples_from_older_worke
|
|||
assert picked["model_info"]["id"] == FAST_TTFT_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("sync_mode", [True, False], ids=["sync", "async"])
|
||||
@pytest.mark.parametrize(
|
||||
("ttft_percentile", "first_samples", "second_samples", "expected_id"),
|
||||
[
|
||||
(None, [0.1, 0.1, 1.0], [0.3, 0.3, 0.3], SLOW_TTFT_ID),
|
||||
(0.5, [0.1, 0.1, 1.0], [0.3, 0.3, 0.3], FAST_TTFT_ID),
|
||||
(0.9, [0.1, 0.1, 0.1, 0.1, 1.5], [0.3, 0.3, 0.3, 0.3, 0.3], SLOW_TTFT_ID),
|
||||
],
|
||||
ids=["default_average", "p50", "p90"],
|
||||
)
|
||||
async def test_streaming_ttft_ranking_percentile(
|
||||
sync_mode: bool,
|
||||
ttft_percentile: float | None,
|
||||
first_samples: list[float],
|
||||
second_samples: list[float],
|
||||
expected_id: str,
|
||||
):
|
||||
cache = DualCache()
|
||||
routing_args = {} if ttft_percentile is None else {"ttft_percentile": ttft_percentile}
|
||||
handler = LowestLatencyLoggingHandler(router_cache=cache, routing_args=routing_args)
|
||||
cache.set_cache(
|
||||
key=f"{MODEL_GROUP}_map",
|
||||
value={
|
||||
FAST_TTFT_ID: {"time_to_first_token_seconds": first_samples},
|
||||
SLOW_TTFT_ID: {"time_to_first_token_seconds": second_samples},
|
||||
},
|
||||
)
|
||||
|
||||
if sync_mode:
|
||||
picked = handler.get_available_deployments(
|
||||
model_group=MODEL_GROUP,
|
||||
healthy_deployments=STREAMING_DEPLOYMENTS,
|
||||
request_kwargs={"stream": True, "metadata": {}},
|
||||
)
|
||||
else:
|
||||
picked = await handler.async_get_available_deployments(
|
||||
model_group=MODEL_GROUP,
|
||||
healthy_deployments=STREAMING_DEPLOYMENTS,
|
||||
request_kwargs={"stream": True, "metadata": {}},
|
||||
)
|
||||
|
||||
assert picked is not None
|
||||
assert picked["model_info"]["id"] == expected_id
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ttft_percentile", [0, -0.1, 1.1])
|
||||
def test_ttft_percentile_validation(ttft_percentile: float):
|
||||
with pytest.raises(ValidationError):
|
||||
RoutingArgs(ttft_percentile=ttft_percentile)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ttft_percentile", [0.5, 0.9, 0.95, 1.0])
|
||||
def test_ttft_percentile_accepts_valid_values(ttft_percentile: float):
|
||||
assert RoutingArgs(ttft_percentile=ttft_percentile).ttft_percentile == ttft_percentile
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ttft_percentile_does_not_change_non_streaming_routing():
|
||||
cache = DualCache()
|
||||
handler = LowestLatencyLoggingHandler(router_cache=cache, routing_args={"ttft_percentile": 0.9})
|
||||
cache.set_cache(
|
||||
key=f"{MODEL_GROUP}_map",
|
||||
value={
|
||||
FAST_TTFT_ID: {"latency": [1.0], "time_to_first_token_seconds": [0.1]},
|
||||
SLOW_TTFT_ID: {"latency": [0.2], "time_to_first_token_seconds": [1.5]},
|
||||
},
|
||||
)
|
||||
|
||||
picked = await handler.async_get_available_deployments(
|
||||
model_group=MODEL_GROUP,
|
||||
healthy_deployments=STREAMING_DEPLOYMENTS,
|
||||
request_kwargs={"stream": False, "metadata": {}},
|
||||
)
|
||||
|
||||
assert picked is not None
|
||||
assert picked["model_info"]["id"] == SLOW_TTFT_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"cached_entry",
|
||||
|
|
@ -318,3 +398,81 @@ async def test_async_get_available_deployments_treats_missing_samples_as_zero_la
|
|||
|
||||
assert picked is not None
|
||||
assert picked["model_info"]["id"] == DEPLOYMENT_ID
|
||||
|
||||
|
||||
def _latency_router(routing_strategy_args: dict) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": MODEL_GROUP,
|
||||
"litellm_params": {"model": f"openai/{MODEL_GROUP}", "api_key": "sk-fake"},
|
||||
"model_info": {"id": deployment_id},
|
||||
}
|
||||
for deployment_id in (FAST_TTFT_ID, SLOW_TTFT_ID)
|
||||
],
|
||||
routing_strategy="latency-based-routing",
|
||||
routing_strategy_args=routing_strategy_args,
|
||||
)
|
||||
|
||||
|
||||
def _seed_streaming_ttft(router: Router) -> None:
|
||||
router.cache.set_cache(
|
||||
key=f"{MODEL_GROUP}_map",
|
||||
value={
|
||||
FAST_TTFT_ID: {"time_to_first_token_seconds": [0.1, 0.1, 1.0]},
|
||||
SLOW_TTFT_ID: {"time_to_first_token_seconds": [0.3, 0.3, 0.3]},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _pick_streaming(router: Router) -> str:
|
||||
picked = await router.async_get_available_deployment(
|
||||
model=MODEL_GROUP,
|
||||
request_kwargs={"stream": True, "metadata": {}},
|
||||
)
|
||||
return picked["model_info"]["id"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_routing_strategy_args_update_applies_ttft_percentile():
|
||||
"""A config reload that adds ttft_percentile must reach the live selector,
|
||||
not sit unused until the proxy restarts."""
|
||||
router = _latency_router({"max_latency_list_size": 50})
|
||||
_seed_streaming_ttft(router)
|
||||
|
||||
assert await _pick_streaming(router) == SLOW_TTFT_ID
|
||||
|
||||
router.update_settings(routing_strategy_args={"max_latency_list_size": 50, "ttft_percentile": 0.5})
|
||||
|
||||
assert await _pick_streaming(router) == FAST_TTFT_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_routing_strategy_args_update_keeps_previous_args_when_invalid():
|
||||
router = _latency_router({"ttft_percentile": 0.5})
|
||||
_seed_streaming_ttft(router)
|
||||
|
||||
router.update_settings(routing_strategy_args={"ttft_percentile": 5})
|
||||
|
||||
assert await _pick_streaming(router) == FAST_TTFT_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_routing_strategy_args_update_is_a_noop_without_a_selector():
|
||||
"""simple-shuffle has no selector to re-link, so an args update must leave
|
||||
the router alone instead of blowing up on a missing selector attribute."""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": MODEL_GROUP,
|
||||
"litellm_params": {"model": f"openai/{MODEL_GROUP}", "api_key": "sk-fake"},
|
||||
"model_info": {"id": FAST_TTFT_ID},
|
||||
}
|
||||
],
|
||||
routing_strategy="simple-shuffle",
|
||||
)
|
||||
|
||||
router.update_settings(routing_strategy_args={"ttl": 5})
|
||||
|
||||
assert router.routing_strategy_args == {"ttl": 5}
|
||||
assert await _pick_streaming(router) == FAST_TTFT_ID
|
||||
|
|
|
|||
|
|
@ -16,12 +16,10 @@ The mechanism works without any cache and supports two encoding strategies:
|
|||
"""
|
||||
|
||||
import time
|
||||
from typing import List, Optional
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
|
@ -68,9 +66,7 @@ class TestEncryptedItemIdCodec:
|
|||
def test_roundtrip(self):
|
||||
model_id = "deployment-1"
|
||||
original_item_id = "rs_abc123def456"
|
||||
encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
model_id, original_item_id
|
||||
)
|
||||
encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_item_id)
|
||||
assert encoded.startswith("encitem_")
|
||||
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded)
|
||||
assert decoded is not None
|
||||
|
|
@ -81,9 +77,7 @@ class TestEncryptedItemIdCodec:
|
|||
"""Decoding must succeed even if base64 padding (=) was stripped in transit."""
|
||||
model_id = "gpt-5.1-codex-openai-2"
|
||||
original_item_id = "rs_0efb96cb222403210069a01d5d52588196a9dc394ffdb89d00"
|
||||
encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
model_id, original_item_id
|
||||
)
|
||||
encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_item_id)
|
||||
# Strip any trailing '=' to simulate what happens in transit
|
||||
stripped = encoded.rstrip("=")
|
||||
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(stripped)
|
||||
|
|
@ -100,9 +94,7 @@ class TestEncryptedItemIdCodec:
|
|||
"""item_id values containing ';' must survive the roundtrip."""
|
||||
model_id = "deployment-1"
|
||||
original_item_id = "rs_part1;part2;part3"
|
||||
encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
model_id, original_item_id
|
||||
)
|
||||
encoded = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_item_id)
|
||||
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded)
|
||||
assert decoded is not None
|
||||
assert decoded["item_id"] == original_item_id
|
||||
|
|
@ -118,11 +110,7 @@ class TestUpdateEncryptedContentItemIds:
|
|||
{"id": "rs_xyz", "type": "reasoning", "encrypted_content": "secret"},
|
||||
],
|
||||
}
|
||||
result = (
|
||||
ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(
|
||||
response, model_id
|
||||
)
|
||||
)
|
||||
result = ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(response, model_id)
|
||||
# Plain message item untouched
|
||||
assert result["output"][0]["id"] == "msg_abc"
|
||||
# Reasoning item with encrypted_content gets encoded
|
||||
|
|
@ -133,16 +121,8 @@ class TestUpdateEncryptedContentItemIds:
|
|||
assert decoded["item_id"] == "rs_xyz"
|
||||
|
||||
def test_no_op_when_model_id_is_none(self):
|
||||
response = {
|
||||
"output": [
|
||||
{"id": "rs_xyz", "type": "reasoning", "encrypted_content": "secret"}
|
||||
]
|
||||
}
|
||||
result = (
|
||||
ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(
|
||||
response, None
|
||||
)
|
||||
)
|
||||
response = {"output": [{"id": "rs_xyz", "type": "reasoning", "encrypted_content": "secret"}]}
|
||||
result = ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(response, None)
|
||||
assert result["output"][0]["id"] == "rs_xyz"
|
||||
|
||||
|
||||
|
|
@ -151,9 +131,7 @@ class TestEncryptedContentWrapping:
|
|||
"""Test wrapping encrypted_content with model_id metadata."""
|
||||
model_id = "deployment-1"
|
||||
original_content = "gAAAAABpnW_yEYmSNEyOG_original_encrypted_data"
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
|
||||
original_content, model_id
|
||||
)
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id)
|
||||
assert wrapped.startswith("litellm_enc:")
|
||||
assert wrapped != original_content
|
||||
|
||||
|
|
@ -170,9 +148,7 @@ class TestEncryptedContentWrapping:
|
|||
(
|
||||
model_id,
|
||||
content,
|
||||
) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(
|
||||
plain_content
|
||||
)
|
||||
) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(plain_content)
|
||||
assert model_id is None
|
||||
assert content == plain_content
|
||||
|
||||
|
|
@ -189,11 +165,7 @@ class TestEncryptedContentWrapping:
|
|||
},
|
||||
],
|
||||
}
|
||||
result = (
|
||||
ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(
|
||||
response, model_id
|
||||
)
|
||||
)
|
||||
result = ResponsesAPIRequestUtils._update_encrypted_content_item_ids_in_response(response, model_id)
|
||||
assert result["output"][0].get("encrypted_content") is None
|
||||
wrapped = result["output"][1]["encrypted_content"]
|
||||
assert wrapped.startswith("litellm_enc:")
|
||||
|
|
@ -210,19 +182,13 @@ class TestRestoreEncryptedContentItemIds:
|
|||
def test_restores_encoded_ids(self):
|
||||
model_id = "deployment-1"
|
||||
original_id = "rs_encrypted_item_456"
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
model_id, original_id
|
||||
)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(model_id, original_id)
|
||||
|
||||
request_input = [
|
||||
{"type": "message", "id": "msg_abc123", "role": "assistant"},
|
||||
{"type": "reasoning", "id": encoded_id, "encrypted_content": "secret"},
|
||||
]
|
||||
restored = (
|
||||
ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
|
||||
request_input
|
||||
)
|
||||
)
|
||||
restored = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(request_input)
|
||||
assert restored[0]["id"] == "msg_abc123"
|
||||
assert restored[1]["id"] == original_id
|
||||
|
||||
|
|
@ -230,33 +196,21 @@ class TestRestoreEncryptedContentItemIds:
|
|||
"""Test that wrapped encrypted_content is unwrapped before forwarding."""
|
||||
model_id = "deployment-1"
|
||||
original_content = "gAAAAABpnW_yEYmSNEyOG_original"
|
||||
wrapped_content = (
|
||||
ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
|
||||
original_content, model_id
|
||||
)
|
||||
)
|
||||
wrapped_content = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id)
|
||||
|
||||
request_input = [
|
||||
{"type": "reasoning", "encrypted_content": wrapped_content},
|
||||
]
|
||||
restored = (
|
||||
ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
|
||||
request_input
|
||||
)
|
||||
)
|
||||
restored = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(request_input)
|
||||
assert restored[0]["encrypted_content"] == original_content
|
||||
|
||||
def test_no_op_for_plain_string_input(self):
|
||||
result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
|
||||
"Hello world"
|
||||
)
|
||||
result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input("Hello world")
|
||||
assert result == "Hello world"
|
||||
|
||||
def test_no_op_for_unencoded_ids(self):
|
||||
request_input = [{"type": "message", "id": "msg_plain"}]
|
||||
result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(
|
||||
request_input
|
||||
)
|
||||
result = ResponsesAPIRequestUtils._restore_encrypted_content_item_ids_in_input(request_input)
|
||||
assert result[0]["id"] == "msg_plain"
|
||||
|
||||
|
||||
|
|
@ -283,9 +237,7 @@ async def test_encrypted_content_affinity_tracks_and_routes():
|
|||
"id": "msg_abc123",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "output_text", "text": "Hello!", "annotations": []}
|
||||
],
|
||||
"content": [{"type": "output_text", "text": "Hello!", "annotations": []}],
|
||||
},
|
||||
{
|
||||
"type": "reasoning",
|
||||
|
|
@ -347,9 +299,9 @@ async def test_encrypted_content_affinity_tracks_and_routes():
|
|||
|
||||
# The response must have rewritten the encrypted item's ID to encoded form
|
||||
encoded_item_id = _extract_encoded_item_id(first_response)
|
||||
assert encoded_item_id.startswith(
|
||||
"encitem_"
|
||||
), f"Expected output item ID to be rewritten to encitem_... but got {encoded_item_id!r}"
|
||||
assert encoded_item_id.startswith("encitem_"), (
|
||||
f"Expected output item ID to be rewritten to encitem_... but got {encoded_item_id!r}"
|
||||
)
|
||||
|
||||
# Verify the encoded ID decodes back to the correct deployment + original ID
|
||||
decoded = ResponsesAPIRequestUtils._decode_encrypted_item_id(encoded_item_id)
|
||||
|
|
@ -371,9 +323,9 @@ async def test_encrypted_content_affinity_tracks_and_routes():
|
|||
)
|
||||
second_model_id = second_response._hidden_params["model_id"]
|
||||
|
||||
assert (
|
||||
second_model_id == first_model_id
|
||||
), f"Expected affinity to route to {first_model_id}, but got {second_model_id}"
|
||||
assert second_model_id == first_model_id, (
|
||||
f"Expected affinity to route to {first_model_id}, but got {second_model_id}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -478,9 +430,7 @@ async def test_encrypted_content_affinity_bypasses_rpm_limits():
|
|||
|
||||
# Extract encoded item ID from the first response output
|
||||
encoded_item_id = _extract_encoded_item_id(first_response)
|
||||
assert encoded_item_id.startswith(
|
||||
"encitem_"
|
||||
), f"Expected encitem_... but got {encoded_item_id!r}"
|
||||
assert encoded_item_id.startswith("encitem_"), f"Expected encitem_... but got {encoded_item_id!r}"
|
||||
|
||||
# Follow-up with the encoded item ID — should pin to same deployment
|
||||
second_response = await router.aresponses(
|
||||
|
|
@ -628,17 +578,13 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
|
|||
if hasattr(first_item, "encrypted_content")
|
||||
else first_item.get("encrypted_content")
|
||||
)
|
||||
assert wrapped_content.startswith(
|
||||
"litellm_enc:"
|
||||
), f"Expected wrapped content but got {wrapped_content[:50]}..."
|
||||
assert wrapped_content.startswith("litellm_enc:"), f"Expected wrapped content but got {wrapped_content[:50]}..."
|
||||
|
||||
# Verify we can extract model_id from wrapped content
|
||||
(
|
||||
extracted_model_id,
|
||||
_,
|
||||
) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(
|
||||
wrapped_content
|
||||
)
|
||||
) = ResponsesAPIRequestUtils._unwrap_encrypted_content_with_model_id(wrapped_content)
|
||||
assert extracted_model_id == first_model_id
|
||||
|
||||
# Second request: use wrapped encrypted_content WITHOUT an ID (Codex behavior)
|
||||
|
|
@ -653,9 +599,9 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
|
|||
)
|
||||
second_model_id = second_response._hidden_params["model_id"]
|
||||
|
||||
assert (
|
||||
second_model_id == first_model_id
|
||||
), f"Expected affinity to route to {first_model_id}, but got {second_model_id}"
|
||||
assert second_model_id == first_model_id, (
|
||||
f"Expected affinity to route to {first_model_id}, but got {second_model_id}"
|
||||
)
|
||||
|
||||
|
||||
def test_encrypted_content_wrapping_preserves_original_content():
|
||||
|
|
@ -664,13 +610,9 @@ def test_encrypted_content_wrapping_preserves_original_content():
|
|||
This is critical for streaming responses where content must round-trip correctly.
|
||||
"""
|
||||
model_id = "test-deployment-1"
|
||||
original_encrypted_content = (
|
||||
"gAAAAABpnW_yEYmSNEyOG_streaming_test_content_with_special_chars==+/"
|
||||
)
|
||||
original_encrypted_content = "gAAAAABpnW_yEYmSNEyOG_streaming_test_content_with_special_chars==+/"
|
||||
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
|
||||
original_encrypted_content, model_id
|
||||
)
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_encrypted_content, model_id)
|
||||
|
||||
assert wrapped.startswith("litellm_enc:")
|
||||
assert wrapped != original_encrypted_content
|
||||
|
|
@ -691,9 +633,7 @@ def test_encrypted_content_wrapping_with_multiple_semicolons():
|
|||
model_id = "deployment-with-semicolons"
|
||||
original_content = "gAAAAAB;some;content;with;semicolons"
|
||||
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
|
||||
original_content, model_id
|
||||
)
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id)
|
||||
|
||||
(
|
||||
extracted_model_id,
|
||||
|
|
@ -764,9 +704,7 @@ async def test_encrypted_content_affinity_preserves_litellm_metadata_for_respons
|
|||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
assert (
|
||||
request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] is True
|
||||
)
|
||||
assert request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] is True
|
||||
assert request_kwargs["litellm_metadata"]["model_info"] == {"id": "dep-1"}
|
||||
|
||||
|
||||
|
|
@ -777,9 +715,7 @@ def test_encrypted_content_wrapping_empty_string():
|
|||
model_id = "test-deployment"
|
||||
original_content = ""
|
||||
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(
|
||||
original_content, model_id
|
||||
)
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id(original_content, model_id)
|
||||
|
||||
assert wrapped.startswith("litellm_enc:")
|
||||
|
||||
|
|
@ -1132,9 +1068,7 @@ def test_boundary_key_accepts_pydantic_litellm_params_instance():
|
|||
"api_key": "fake-azure-resource-key-a",
|
||||
}
|
||||
|
||||
pydantic_key = EncryptedContentAffinityCheck._encryption_boundary_key(
|
||||
pydantic_params
|
||||
)
|
||||
pydantic_key = EncryptedContentAffinityCheck._encryption_boundary_key(pydantic_params)
|
||||
plain_key = EncryptedContentAffinityCheck._encryption_boundary_key(plain_params)
|
||||
|
||||
assert pydantic_key is not None
|
||||
|
|
@ -1161,18 +1095,8 @@ def test_boundary_key_rejects_non_dict_like_inputs():
|
|||
for bad in (None, [], "not a dict", 42, object()):
|
||||
assert EncryptedContentAffinityCheck._encryption_boundary_key(bad) is None
|
||||
|
||||
assert (
|
||||
EncryptedContentAffinityCheck._encryption_boundary_key(
|
||||
{"api_base": "", "api_key": "k"}
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert (
|
||||
EncryptedContentAffinityCheck._encryption_boundary_key(
|
||||
{"api_base": "https://x"}
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert EncryptedContentAffinityCheck._encryption_boundary_key({"api_base": "", "api_key": "k"}) is None
|
||||
assert EncryptedContentAffinityCheck._encryption_boundary_key({"api_base": "https://x"}) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -1180,10 +1104,11 @@ def test_boundary_key_rejects_non_dict_like_inputs():
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_originating_mock(api_base: str, api_key: str):
|
||||
def _make_originating_mock(api_base: str, api_key: str, model_name: str = "gpt-5.4"):
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
originating = MagicMock()
|
||||
originating.model_name = model_name
|
||||
originating.litellm_params.model_dump.return_value = {
|
||||
"api_base": api_base,
|
||||
"api_key": api_key,
|
||||
|
|
@ -1192,19 +1117,23 @@ def _make_originating_mock(api_base: str, api_key: str):
|
|||
|
||||
|
||||
def _make_router_mock_with_cooldown(
|
||||
originating, cooldown_entries: Optional[List[tuple]] = None
|
||||
originating,
|
||||
cooldown_entries: list[tuple] | None = None,
|
||||
routed_group_model_ids: list[str] | None = None,
|
||||
):
|
||||
"""
|
||||
Build a MagicMock router whose ``cooldown_cache.async_get_active_cooldowns``
|
||||
returns ``cooldown_entries`` (defaulting to ``[]`` — no active cooldown).
|
||||
returns ``cooldown_entries`` (defaulting to ``[]`` — no active cooldown), and
|
||||
whose ``get_candidate_model_ids_for_route`` returns ``routed_group_model_ids``
|
||||
(the deployment ids the router resolves for the routed model; defaulting to ``[]``
|
||||
— origin absent from the routed group, i.e. a tier change).
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment.return_value = originating
|
||||
mock_router.cooldown_cache.async_get_active_cooldowns = AsyncMock(
|
||||
return_value=list(cooldown_entries or [])
|
||||
)
|
||||
mock_router.cooldown_cache.async_get_active_cooldowns = AsyncMock(return_value=list(cooldown_entries or []))
|
||||
mock_router.get_candidate_model_ids_for_route.return_value = frozenset(routed_group_model_ids or [])
|
||||
return mock_router
|
||||
|
||||
|
||||
|
|
@ -1235,15 +1164,15 @@ async def test_affinity_raises_service_unavailable_when_origin_cooled_for_non_42
|
|||
},
|
||||
)
|
||||
],
|
||||
routed_group_model_ids=["deployment-a-cooled", "deployment-b"],
|
||||
)
|
||||
|
||||
check = EncryptedContentAffinityCheck(router=mock_router)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
"deployment-a-cooled", "rs_test"
|
||||
)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a-cooled", "rs_test")
|
||||
healthy_only_b = [
|
||||
{
|
||||
"model_info": {"id": "deployment-b"},
|
||||
"model_name": "gpt-5.4",
|
||||
"litellm_params": {
|
||||
"api_base": "https://account-b.openai.azure.com/",
|
||||
"api_key": "key-b",
|
||||
|
|
@ -1297,15 +1226,15 @@ async def test_affinity_raises_rate_limit_with_retry_after_when_origin_cooled_fo
|
|||
},
|
||||
)
|
||||
],
|
||||
routed_group_model_ids=["deployment-a-cooled-429", "deployment-b"],
|
||||
)
|
||||
|
||||
check = EncryptedContentAffinityCheck(router=mock_router)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
"deployment-a-cooled-429", "rs_test"
|
||||
)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a-cooled-429", "rs_test")
|
||||
healthy_only_b = [
|
||||
{
|
||||
"model_info": {"id": "deployment-b"},
|
||||
"model_name": "gpt-5.4",
|
||||
"litellm_params": {
|
||||
"api_base": "https://account-b.openai.azure.com/",
|
||||
"api_key": "key-b",
|
||||
|
|
@ -1345,15 +1274,16 @@ async def test_affinity_raises_service_unavailable_when_origin_filtered_without_
|
|||
)
|
||||
|
||||
originating = _make_originating_mock("https://account-a.openai.azure.com/", "key-a")
|
||||
mock_router = _make_router_mock_with_cooldown(originating, cooldown_entries=[])
|
||||
mock_router = _make_router_mock_with_cooldown(
|
||||
originating, cooldown_entries=[], routed_group_model_ids=["deployment-a-filtered", "deployment-b"]
|
||||
)
|
||||
|
||||
check = EncryptedContentAffinityCheck(router=mock_router)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
"deployment-a-filtered", "rs_test"
|
||||
)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a-filtered", "rs_test")
|
||||
healthy_only_b = [
|
||||
{
|
||||
"model_info": {"id": "deployment-b"},
|
||||
"model_name": "gpt-5.4",
|
||||
"litellm_params": {
|
||||
"api_base": "https://account-b.openai.azure.com/",
|
||||
"api_key": "key-b",
|
||||
|
|
@ -1377,15 +1307,18 @@ async def test_affinity_raises_service_unavailable_when_origin_filtered_without_
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_affinity_raises_bad_request_when_origin_removed():
|
||||
async def test_affinity_strips_and_dispatches_when_origin_is_unknown_or_removed():
|
||||
"""
|
||||
Originating deployment was removed from the router config and no boundary
|
||||
peer is available. This is permanent (the stale encrypted_content cannot
|
||||
be honored), so surface a 400 with actionable text.
|
||||
A removed deployment, or a forged/unknown affinity marker, resolves to no
|
||||
originating deployment. It is handled like a cross-group origin: the encrypted
|
||||
reasoning is stripped and the request dispatches with its readable history,
|
||||
rather than returning a distinguishable error. That uniform handling denies an
|
||||
authenticated caller a deployment-id existence oracle, an existing cross-group id
|
||||
and a nonexistent id both strip and proceed, so responses cannot be told apart.
|
||||
The membership lookup is skipped entirely when the origin is unknown.
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
|
||||
EncryptedContentAffinityCheck,
|
||||
)
|
||||
|
|
@ -1394,12 +1327,11 @@ async def test_affinity_raises_bad_request_when_origin_removed():
|
|||
mock_router.get_deployment.return_value = None
|
||||
|
||||
check = EncryptedContentAffinityCheck(router=mock_router)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
"deployment-removed", "rs_test"
|
||||
)
|
||||
healthy_only_b = [
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-removed")
|
||||
routed_pool = [
|
||||
{
|
||||
"model_info": {"id": "deployment-b"},
|
||||
"model_name": "gpt-5.4",
|
||||
"litellm_params": {
|
||||
"api_base": "https://account-b.openai.azure.com/",
|
||||
"api_key": "key-b",
|
||||
|
|
@ -1408,18 +1340,28 @@ async def test_affinity_raises_bad_request_when_origin_removed():
|
|||
}
|
||||
]
|
||||
request_kwargs = {
|
||||
"input": [{"id": encoded_id, "type": "reasoning"}],
|
||||
"litellm_metadata": {},
|
||||
"input": [
|
||||
{"role": "user", "content": "why is the sky blue?"},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"encrypted_content": wrapped,
|
||||
"summary": [{"type": "summary_text", "text": "scattering"}],
|
||||
},
|
||||
{"role": "user", "content": "and sunsets?"},
|
||||
],
|
||||
}
|
||||
|
||||
with pytest.raises(BadRequestError) as excinfo:
|
||||
await check.async_filter_deployments(
|
||||
model="gpt-5.4",
|
||||
healthy_deployments=healthy_only_b,
|
||||
messages=None,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
result = await check.async_filter_deployments(
|
||||
model="gpt-5.4",
|
||||
healthy_deployments=routed_pool,
|
||||
messages=None,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
assert "deployment-removed" not in str(excinfo.value)
|
||||
assert result is routed_pool
|
||||
assert not any(isinstance(item, dict) and item.get("encrypted_content") for item in request_kwargs["input"])
|
||||
mock_router.get_candidate_model_ids_for_route.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1444,9 +1386,7 @@ async def test_affinity_does_not_raise_when_boundary_peer_available():
|
|||
mock_router.get_deployment.return_value = originating
|
||||
|
||||
check = EncryptedContentAffinityCheck(router=mock_router)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
"deployment-a", "rs_test"
|
||||
)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a", "rs_test")
|
||||
peer = {
|
||||
"model_info": {"id": "deployment-a-peer"},
|
||||
"litellm_params": {
|
||||
|
|
@ -1490,9 +1430,7 @@ async def test_model_group_affinity_config_enables_encrypted_content_affinity():
|
|||
},
|
||||
target_deployment,
|
||||
]
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
"deployment-b", "rs_test"
|
||||
)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-b", "rs_test")
|
||||
request_kwargs = {
|
||||
"input": [{"type": "reasoning", "id": encoded_id}],
|
||||
"litellm_metadata": {},
|
||||
|
|
@ -1536,9 +1474,7 @@ async def test_model_group_affinity_config_does_not_disable_global_encrypted_con
|
|||
},
|
||||
target_deployment,
|
||||
]
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
"deployment-b", "rs_test"
|
||||
)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-b", "rs_test")
|
||||
request_kwargs = {
|
||||
"input": [{"type": "reasoning", "id": encoded_id}],
|
||||
"litellm_metadata": {},
|
||||
|
|
@ -1600,15 +1536,9 @@ async def test_model_group_encrypted_content_affinity_overrides_global_deploymen
|
|||
|
||||
try:
|
||||
callbacks = router.optional_callbacks or []
|
||||
deployment_callback = next(
|
||||
cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck)
|
||||
)
|
||||
encrypted_content_callback = next(
|
||||
cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck)
|
||||
)
|
||||
assert callbacks.index(encrypted_content_callback) < callbacks.index(
|
||||
deployment_callback
|
||||
)
|
||||
deployment_callback = next(cb for cb in callbacks if isinstance(cb, DeploymentAffinityCheck))
|
||||
encrypted_content_callback = next(cb for cb in callbacks if isinstance(cb, EncryptedContentAffinityCheck))
|
||||
assert callbacks.index(encrypted_content_callback) < callbacks.index(deployment_callback)
|
||||
assert encrypted_content_callback.enable_global_affinity is False
|
||||
|
||||
cache_key = DeploymentAffinityCheck.get_affinity_cache_key(
|
||||
|
|
@ -1620,9 +1550,7 @@ async def test_model_group_encrypted_content_affinity_overrides_global_deploymen
|
|||
value={"model_id": "deployment-a"},
|
||||
ttl=60,
|
||||
)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id(
|
||||
"deployment-b", "rs_test"
|
||||
)
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-b", "rs_test")
|
||||
request_kwargs = {
|
||||
"input": [
|
||||
{
|
||||
|
|
@ -1643,16 +1571,301 @@ async def test_model_group_encrypted_content_affinity_overrides_global_deploymen
|
|||
)
|
||||
assert after_deployment_affinity == [deployment_a, deployment_b]
|
||||
|
||||
after_encrypted_content_affinity = (
|
||||
await encrypted_content_callback.async_filter_deployments(
|
||||
model=model_group,
|
||||
healthy_deployments=after_deployment_affinity,
|
||||
messages=None,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
after_encrypted_content_affinity = await encrypted_content_callback.async_filter_deployments(
|
||||
model=model_group,
|
||||
healthy_deployments=after_deployment_affinity,
|
||||
messages=None,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
assert after_encrypted_content_affinity == [deployment_b]
|
||||
assert request_kwargs.get("_encrypted_content_affinity_pinned") is True
|
||||
finally:
|
||||
router.discard()
|
||||
|
||||
|
||||
class TestStripEncryptedReasoningFromInput:
|
||||
def test_keeps_summary_and_drops_encrypted_content_and_id(self):
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a")
|
||||
encoded_id = ResponsesAPIRequestUtils._build_encrypted_item_id("deployment-a", "rs_1")
|
||||
request_input = [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"id": encoded_id,
|
||||
"encrypted_content": wrapped,
|
||||
"summary": [{"type": "summary_text", "text": "thought about it"}],
|
||||
},
|
||||
{"type": "reasoning", "id": encoded_id, "encrypted_content": wrapped},
|
||||
{"type": "reasoning", "encrypted_content": wrapped, "summary": []},
|
||||
{"type": "message", "id": "msg_1", "role": "assistant", "content": "hi"},
|
||||
{"role": "user", "content": "second turn"},
|
||||
]
|
||||
ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
|
||||
assert request_input == [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"summary": [{"type": "summary_text", "text": "thought about it"}],
|
||||
},
|
||||
{"type": "message", "id": "msg_1", "role": "assistant", "content": "hi"},
|
||||
{"role": "user", "content": "second turn"},
|
||||
]
|
||||
|
||||
def test_keeps_string_form_summary_when_stripping(self):
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a")
|
||||
request_input = [
|
||||
{"type": "reasoning", "encrypted_content": wrapped, "summary": "plain string thought"},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"encrypted_content": wrapped,
|
||||
"content": [{"type": "output_text", "text": "in content"}],
|
||||
},
|
||||
{"type": "reasoning", "encrypted_content": wrapped, "summary": "", "content": []},
|
||||
]
|
||||
ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
|
||||
assert request_input == [
|
||||
{"type": "reasoning", "summary": "plain string thought"},
|
||||
{"type": "reasoning", "content": [{"type": "output_text", "text": "in content"}]},
|
||||
]
|
||||
|
||||
def test_leaves_input_untouched_when_no_encrypted_reasoning(self):
|
||||
request_input = [
|
||||
{"role": "user", "content": "first turn"},
|
||||
{"type": "reasoning", "summary": [{"type": "summary_text", "text": "no blob here"}]},
|
||||
{"role": "user", "content": "second turn"},
|
||||
]
|
||||
before = [dict(item) for item in request_input]
|
||||
ResponsesAPIRequestUtils.strip_encrypted_reasoning_from_input(request_input)
|
||||
assert request_input == before
|
||||
|
||||
|
||||
def _cross_group_request_kwargs():
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-a")
|
||||
return {
|
||||
"litellm_metadata": {},
|
||||
"input": [
|
||||
{"role": "user", "content": "ZEBRA: why is the sky blue?"},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"encrypted_content": wrapped,
|
||||
"summary": [{"type": "summary_text", "text": "scattering"}],
|
||||
},
|
||||
{"type": "message", "role": "assistant", "content": "Rayleigh scattering."},
|
||||
{"role": "user", "content": "KIWI: and sunsets?"},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_affinity_strips_encrypted_reasoning_when_routed_to_another_model_group():
|
||||
"""
|
||||
An auto-router tier change (or a model switch with no boundary peer): the
|
||||
routed pool holds no deployment of the origin's model group. The origin is
|
||||
healthy, so a 503 would be wrong; the follow-up dispatches to the routed
|
||||
pool with the origin's encrypted reasoning stripped and its summary kept.
|
||||
"""
|
||||
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
|
||||
EncryptedContentAffinityCheck,
|
||||
)
|
||||
|
||||
originating = _make_originating_mock(None, "key-a", model_name="gpt-reasoning-tier")
|
||||
mock_router = _make_router_mock_with_cooldown(
|
||||
originating, cooldown_entries=[], routed_group_model_ids=["deployment-b"]
|
||||
)
|
||||
check = EncryptedContentAffinityCheck(router=mock_router)
|
||||
routed_pool = [
|
||||
{
|
||||
"model_info": {"id": "deployment-b"},
|
||||
"model_name": "gpt-simple-tier",
|
||||
"litellm_params": {
|
||||
"api_base": "https://gateway.example/v1",
|
||||
"api_key": "key-b",
|
||||
"model": "openai/gpt-5-nano",
|
||||
},
|
||||
}
|
||||
]
|
||||
request_kwargs = _cross_group_request_kwargs()
|
||||
original_input = request_kwargs["input"]
|
||||
|
||||
result = await check.async_filter_deployments(
|
||||
model="gpt-simple-tier",
|
||||
healthy_deployments=routed_pool,
|
||||
messages=None,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
assert result is routed_pool
|
||||
assert "_encrypted_content_affinity_pinned" not in request_kwargs
|
||||
assert request_kwargs["litellm_metadata"]["encrypted_content_affinity_enabled"] is True
|
||||
assert request_kwargs["input"] is original_input
|
||||
assert [item.get("type") or item["role"] for item in original_input] == [
|
||||
"user",
|
||||
"reasoning",
|
||||
"message",
|
||||
"user",
|
||||
]
|
||||
assert original_input[1] == {
|
||||
"type": "reasoning",
|
||||
"summary": [{"type": "summary_text", "text": "scattering"}],
|
||||
}
|
||||
assert not any(isinstance(item, dict) and item.get("encrypted_content") for item in original_input)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_affinity_fails_fast_within_the_origins_own_group():
|
||||
"""
|
||||
Negative class for the tier-change discriminator: the routed group IS the
|
||||
origin's group (a same-group cooldown, not a tier change), so even with a
|
||||
healthy non-origin sibling that cannot decrypt the content, the request
|
||||
still fails fast and the encrypted reasoning is left intact rather than
|
||||
stripped. Preserves the LIT-3051 cooldown contract.
|
||||
"""
|
||||
from litellm.exceptions import ServiceUnavailableError
|
||||
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
|
||||
EncryptedContentAffinityCheck,
|
||||
)
|
||||
|
||||
originating = _make_originating_mock(
|
||||
"https://account-a.openai.azure.com/", "key-a", model_name="gpt-reasoning-tier"
|
||||
)
|
||||
mock_router = _make_router_mock_with_cooldown(
|
||||
originating, cooldown_entries=[], routed_group_model_ids=["deployment-a", "deployment-a2"]
|
||||
)
|
||||
check = EncryptedContentAffinityCheck(router=mock_router)
|
||||
sibling_pool = [
|
||||
{
|
||||
"model_info": {"id": "deployment-a2"},
|
||||
"model_name": "gpt-reasoning-tier",
|
||||
"litellm_params": {
|
||||
"api_base": "https://account-a2.openai.azure.com/",
|
||||
"api_key": "key-a2",
|
||||
"model": "azure/gpt-5.4",
|
||||
},
|
||||
}
|
||||
]
|
||||
request_kwargs = _cross_group_request_kwargs()
|
||||
|
||||
with pytest.raises(ServiceUnavailableError):
|
||||
await check.async_filter_deployments(
|
||||
model="gpt-reasoning-tier",
|
||||
healthy_deployments=sibling_pool,
|
||||
messages=None,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
assert request_kwargs["input"][1].get("encrypted_content")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_affinity_does_not_strip_when_group_is_spelled_differently_but_same_by_id():
|
||||
"""
|
||||
The discriminator must key on deployment-id membership, not on the model-group
|
||||
name string. Here the origin's configured group is spelled ``openai/gpt-5.4-mini``
|
||||
while the routed group is the canonical ``gpt-5.4-mini``: same group, different
|
||||
spelling. A name compare (``originating.model_name != model``) would read this as
|
||||
a tier change and strip the reasoning it did not have to. Because the origin's id
|
||||
is a member of the routed group, this is a same-group cooldown instead: the request
|
||||
fails fast and the encrypted reasoning is left intact.
|
||||
"""
|
||||
from litellm.exceptions import ServiceUnavailableError
|
||||
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
|
||||
EncryptedContentAffinityCheck,
|
||||
)
|
||||
|
||||
originating = _make_originating_mock(None, "key-a", model_name="openai/gpt-5.4-mini")
|
||||
mock_router = _make_router_mock_with_cooldown(
|
||||
originating, cooldown_entries=[], routed_group_model_ids=["deployment-mini-a", "deployment-mini-b"]
|
||||
)
|
||||
check = EncryptedContentAffinityCheck(router=mock_router)
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-mini-a")
|
||||
sibling_pool = [
|
||||
{
|
||||
"model_info": {"id": "deployment-mini-b"},
|
||||
"model_name": "gpt-5.4-mini",
|
||||
"litellm_params": {
|
||||
"api_base": "https://gateway.example/v1",
|
||||
"api_key": "key-b",
|
||||
"model": "openai/gpt-5.4-mini",
|
||||
},
|
||||
}
|
||||
]
|
||||
request_kwargs = {
|
||||
"litellm_metadata": {},
|
||||
"input": [
|
||||
{"role": "user", "content": "why is the sky blue?"},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"encrypted_content": wrapped,
|
||||
"summary": [{"type": "summary_text", "text": "scattering"}],
|
||||
},
|
||||
{"role": "user", "content": "and sunsets?"},
|
||||
],
|
||||
}
|
||||
|
||||
with pytest.raises(ServiceUnavailableError):
|
||||
await check.async_filter_deployments(
|
||||
model="gpt-5.4-mini",
|
||||
healthy_deployments=sibling_pool,
|
||||
messages=None,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
assert request_kwargs["input"][1].get("encrypted_content")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_affinity_honors_router_candidate_ids_for_team_and_pattern_routes():
|
||||
"""
|
||||
The exact `model_name` index does not include team-public or pattern routes, so a
|
||||
same-group cooldown reached only through one of those would be misread as a tier change
|
||||
and stripped. The check asks the router for the candidate ids it resolves for the route
|
||||
(`get_candidate_model_ids_for_route`), which covers those paths, rather than the bare
|
||||
index. Here that set marks the origin as a candidate, so the request fails fast with its
|
||||
reasoning intact, and the routed group and team are passed through to the router.
|
||||
"""
|
||||
from litellm.exceptions import ServiceUnavailableError
|
||||
from litellm.router_utils.pre_call_checks.encrypted_content_affinity_check import (
|
||||
EncryptedContentAffinityCheck,
|
||||
)
|
||||
|
||||
originating = _make_originating_mock(None, "key-a", model_name="model_name_teamA_uuid")
|
||||
mock_router = _make_router_mock_with_cooldown(
|
||||
originating, cooldown_entries=[], routed_group_model_ids=["deployment-team-a", "deployment-team-b"]
|
||||
)
|
||||
check = EncryptedContentAffinityCheck(router=mock_router)
|
||||
wrapped = ResponsesAPIRequestUtils._wrap_encrypted_content_with_model_id("gAAAAA-blob", "deployment-team-a")
|
||||
sibling_pool = [
|
||||
{
|
||||
"model_info": {"id": "deployment-team-b"},
|
||||
"model_name": "team-public-model",
|
||||
"litellm_params": {
|
||||
"api_base": "https://gateway.example/v1",
|
||||
"api_key": "key-b",
|
||||
"model": "openai/gpt-5.4-mini",
|
||||
},
|
||||
}
|
||||
]
|
||||
request_kwargs = {
|
||||
"litellm_metadata": {"user_api_key_team_id": "teamA"},
|
||||
"input": [
|
||||
{"role": "user", "content": "why is the sky blue?"},
|
||||
{
|
||||
"type": "reasoning",
|
||||
"encrypted_content": wrapped,
|
||||
"summary": [{"type": "summary_text", "text": "scattering"}],
|
||||
},
|
||||
{"role": "user", "content": "and sunsets?"},
|
||||
],
|
||||
}
|
||||
|
||||
with pytest.raises(ServiceUnavailableError):
|
||||
await check.async_filter_deployments(
|
||||
model="team-public-model",
|
||||
healthy_deployments=sibling_pool,
|
||||
messages=None,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
assert request_kwargs["input"][1].get("encrypted_content")
|
||||
mock_router.get_candidate_model_ids_for_route.assert_called_once_with(model="team-public-model", team_id="teamA")
|
||||
|
|
|
|||
|
|
@ -3736,6 +3736,112 @@ def test_completion_cost_logs_reasoning_and_cache_breakdown(_local_model_cost_ma
|
|||
assert logging_obj.cost_breakdown["cache_read_cost"] == pytest.approx(100 * 3e-08)
|
||||
|
||||
|
||||
def test_completion_cost_logs_the_rates_it_billed_at(monkeypatch):
|
||||
"""A caller reporting the cost lines beside their per-token rates reads both off this one call.
|
||||
completion_cost infers the provider, and xai's inclusive tier thresholds put a request sitting
|
||||
exactly on 200k at the tier rate, which a lookup made without that inferred provider would miss.
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"xai/tiered-model",
|
||||
{
|
||||
"input_cost_per_token": 3e-6,
|
||||
"output_cost_per_token": 15e-6,
|
||||
"cache_read_input_token_cost": 3e-7,
|
||||
"input_cost_per_token_above_200k_tokens": 6e-6,
|
||||
"output_cost_per_token_above_200k_tokens": 3e-5,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 6e-7,
|
||||
"litellm_provider": "xai",
|
||||
"mode": "chat",
|
||||
},
|
||||
)
|
||||
logging_obj = Logging(
|
||||
model="xai/tiered-model",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="billed-rates",
|
||||
function_id="f",
|
||||
)
|
||||
usage = Usage(
|
||||
prompt_tokens=200_000,
|
||||
completion_tokens=1_000,
|
||||
total_tokens=201_000,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=100_000),
|
||||
)
|
||||
|
||||
litellm.completion_cost(
|
||||
completion_response=ModelResponse(model="xai/tiered-model", usage=usage),
|
||||
model="xai/tiered-model",
|
||||
custom_llm_provider=None,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
rates = logging_obj.billed_token_rates
|
||||
assert rates is not None
|
||||
assert rates.input_cost_per_token == pytest.approx(6e-6)
|
||||
assert rates.cache_read_input_token_cost == pytest.approx(6e-7)
|
||||
assert logging_obj.cost_breakdown["cache_read_cost"] == pytest.approx(
|
||||
100_000 * rates.cache_read_input_token_cost
|
||||
)
|
||||
assert logging_obj.cost_breakdown["output_cost"] == pytest.approx(1_000 * rates.output_cost_per_token)
|
||||
|
||||
|
||||
def test_completion_cost_logs_cache_and_reasoning_breakdown_for_custom_pricing():
|
||||
"""
|
||||
A custom-priced deployment bills cache tokens at its custom cache rates, but the
|
||||
breakdown stored for the spend logs carried no cache or reasoning lines for it.
|
||||
"""
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.utils import CompletionTokensDetailsWrapper, CostPerToken
|
||||
|
||||
logging_obj = Logging(
|
||||
model="openai/onprem-model",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="custom-pricing-breakdown",
|
||||
function_id="f",
|
||||
)
|
||||
response = ModelResponse(
|
||||
model="openai/onprem-model",
|
||||
usage=Usage(
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=500,
|
||||
total_tokens=1500,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=800, cache_creation_tokens=100),
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=200),
|
||||
),
|
||||
)
|
||||
|
||||
total = completion_cost(
|
||||
completion_response=response,
|
||||
model="openai/onprem-model",
|
||||
custom_llm_provider="openai",
|
||||
custom_cost_per_token=CostPerToken(
|
||||
input_cost_per_token=1e-6,
|
||||
output_cost_per_token=2e-6,
|
||||
cache_read_input_token_cost=1e-7,
|
||||
cache_creation_input_token_cost=1.25e-6,
|
||||
),
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
assert logging_obj.cost_breakdown is not None
|
||||
assert logging_obj.cost_breakdown["cache_read_cost"] == pytest.approx(800 * 1e-7)
|
||||
assert logging_obj.cost_breakdown["cache_creation_cost"] == pytest.approx(100 * 1.25e-6)
|
||||
assert logging_obj.cost_breakdown["reasoning_cost"] == pytest.approx(200 * 2e-6)
|
||||
assert total == pytest.approx(100 * 1e-6 + 800 * 1e-7 + 100 * 1.25e-6 + 500 * 2e-6)
|
||||
|
||||
|
||||
def test_cost_per_token_per_second_pricing(monkeypatch):
|
||||
"""
|
||||
Models priced by duration (input/output_cost_per_second) with no per-token rates
|
||||
|
|
|
|||
|
|
@ -14728,3 +14728,45 @@ def test_router_stays_quiet_when_a_deployment_drop_params_is_a_flag(value, caplo
|
|||
)
|
||||
|
||||
assert "is not a flag value" not in caplog.text
|
||||
|
||||
|
||||
def test_get_candidate_model_ids_for_route_covers_model_name_and_pattern():
|
||||
"""
|
||||
get_candidate_model_ids_for_route resolves a route the way the router does, so a
|
||||
pre-call check can tell a genuine cross-group route from same-group unavailability.
|
||||
A concrete model group returns its member ids; a wildcard/pattern deployment is
|
||||
included for a concrete model it matches, which the bare model_name index misses.
|
||||
Regression guard for the LIT-7195 tier-change discriminator's team/pattern gaps.
|
||||
"""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "grp",
|
||||
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-a", "api_base": "https://x.invalid"},
|
||||
"model_info": {"id": "dep-a"},
|
||||
},
|
||||
{
|
||||
"model_name": "grp",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "sk-b", "api_base": "https://x.invalid"},
|
||||
"model_info": {"id": "dep-b"},
|
||||
},
|
||||
{
|
||||
"model_name": "openai/*",
|
||||
"litellm_params": {"model": "openai/*", "api_key": "sk-c", "api_base": "https://x.invalid"},
|
||||
"model_info": {"id": "dep-wild"},
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
assert router.get_candidate_model_ids_for_route(model="grp") == frozenset({"dep-a", "dep-b"})
|
||||
assert "dep-wild" in router.get_candidate_model_ids_for_route(model="openai/gpt-4o-some-new-model")
|
||||
|
||||
|
||||
def test_deployment_ids_stringifies_ids_and_skips_entries_without_a_model_info_id():
|
||||
deployments = (
|
||||
{"model_info": {"id": "a"}},
|
||||
{"model_info": {"id": 2}},
|
||||
{"model_info": {}},
|
||||
{"no_model_info": True},
|
||||
)
|
||||
assert Router._deployment_ids(deployments) == frozenset({"a", "2"})
|
||||
|
|
|
|||
|
|
@ -2,10 +2,12 @@ import asyncio
|
|||
import json
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
from jsonschema import validate
|
||||
|
|
@ -2402,6 +2404,28 @@ def test_register_model_with_scientific_notation():
|
|||
_invalidate_model_cost_lowercase_map()
|
||||
|
||||
|
||||
@respx.mock
|
||||
def test_register_model_url_fetch_uses_single_attempt(monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_LOCAL_MODEL_COST_MAP", raising=False)
|
||||
monkeypatch.setattr(litellm, "model_cost", dict(litellm.model_cost))
|
||||
before = dict(litellm.model_cost)
|
||||
threads_before = {thread.name for thread in threading.enumerate()}
|
||||
route = respx.get("https://example.invalid/custom_pricing.json").mock(
|
||||
return_value=httpx.Response(503)
|
||||
)
|
||||
|
||||
litellm.register_model(model_cost="https://example.invalid/custom_pricing.json")
|
||||
|
||||
threads_after = {thread.name for thread in threading.enumerate()}
|
||||
assert route.call_count == 1
|
||||
assert not (threads_after - threads_before) & {"litellm-model-cost-map-retry"}
|
||||
assert not any(
|
||||
thread.name == "litellm-model-cost-map-retry" and thread.is_alive()
|
||||
for thread in threading.enumerate()
|
||||
)
|
||||
assert litellm.model_cost.keys() >= before.keys()
|
||||
|
||||
|
||||
def test_register_model_openrouter_without_slash():
|
||||
"""
|
||||
Test that register_model handles openrouter models without '/' in the name.
|
||||
|
|
|
|||
|
|
@ -89,6 +89,65 @@ describe("PassThroughEndpointsTable", () => {
|
|||
expect(onDeleteClick).toHaveBeenCalledWith("ep-1");
|
||||
});
|
||||
|
||||
it("should disable edit and delete for config-defined endpoints", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onEndpointClick = vi.fn();
|
||||
const onDeleteClick = vi.fn();
|
||||
const configEndpoint: passThroughItem = {
|
||||
id: "ep-config",
|
||||
path: "/from-config",
|
||||
target: "https://config.example.com",
|
||||
headers: {},
|
||||
is_from_config: true,
|
||||
};
|
||||
render(
|
||||
<PassThroughEndpointsTable
|
||||
{...defaultProps}
|
||||
endpoints={[configEndpoint]}
|
||||
onEndpointClick={onEndpointClick}
|
||||
onDeleteClick={onDeleteClick}
|
||||
/>,
|
||||
);
|
||||
|
||||
await user.click(screen.getByTestId("endpoint-actions-ep-config"));
|
||||
const editItem = await screen.findByTestId("endpoint-action-edit");
|
||||
const deleteItem = await screen.findByTestId("endpoint-action-delete");
|
||||
|
||||
expect(editItem).toHaveAttribute("data-disabled");
|
||||
expect(deleteItem).toHaveAttribute("data-disabled");
|
||||
expect(screen.getByTestId("endpoint-config-hint")).toHaveTextContent(
|
||||
"This endpoint is defined in the config file and cannot be edited or deleted on the dashboard.",
|
||||
);
|
||||
|
||||
await user.click(editItem);
|
||||
await user.click(deleteItem);
|
||||
|
||||
expect(onEndpointClick).not.toHaveBeenCalled();
|
||||
expect(onDeleteClick).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should not show the config hint for DB endpoints", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<PassThroughEndpointsTable {...defaultProps} />);
|
||||
|
||||
await user.click(screen.getByTestId("endpoint-actions-ep-1"));
|
||||
await screen.findByTestId("endpoint-action-delete");
|
||||
expect(screen.queryByTestId("endpoint-config-hint")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should label endpoint source as Config or DB", () => {
|
||||
const configEndpoint: passThroughItem = {
|
||||
id: "ep-config",
|
||||
path: "/from-config",
|
||||
target: "https://config.example.com",
|
||||
headers: {},
|
||||
is_from_config: true,
|
||||
};
|
||||
render(<PassThroughEndpointsTable {...defaultProps} endpoints={[...endpoints, configEndpoint]} />);
|
||||
expect(screen.getByText("Config")).toBeInTheDocument();
|
||||
expect(screen.getAllByText("DB")).toHaveLength(2);
|
||||
});
|
||||
|
||||
it("should disable edit and delete for endpoints without an id", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onEndpointClick = vi.fn();
|
||||
|
|
|
|||
|
|
@ -18,6 +18,9 @@ import { cn } from "@/lib/cva.config";
|
|||
|
||||
import type { passThroughItem } from "./PassThroughSettings";
|
||||
|
||||
const CONFIG_ENDPOINT_HINT =
|
||||
"This endpoint is defined in the config file and cannot be edited or deleted on the dashboard.";
|
||||
|
||||
function HeaderWithTooltip({ title, tooltip }: { title: string; tooltip: string }) {
|
||||
return (
|
||||
<div className="flex items-center gap-1">
|
||||
|
|
@ -73,6 +76,7 @@ interface EndpointRowActionsProps {
|
|||
|
||||
function EndpointRowActions({ endpoint, onEndpointClick, onDeleteClick }: EndpointRowActionsProps) {
|
||||
const endpointId = endpoint.id;
|
||||
const isFromConfig = endpoint.is_from_config ?? false;
|
||||
return (
|
||||
<DropdownMenu>
|
||||
<DropdownMenuTrigger
|
||||
|
|
@ -85,8 +89,8 @@ function EndpointRowActions({ endpoint, onEndpointClick, onDeleteClick }: Endpoi
|
|||
<DropdownMenuContent align="end" className="w-52">
|
||||
<DropdownMenuItem
|
||||
data-testid="endpoint-action-edit"
|
||||
disabled={!endpointId}
|
||||
onClick={() => endpointId && onEndpointClick(endpointId)}
|
||||
disabled={isFromConfig || !endpointId}
|
||||
onClick={() => !isFromConfig && endpointId && onEndpointClick(endpointId)}
|
||||
>
|
||||
<Pencil />
|
||||
Edit
|
||||
|
|
@ -95,12 +99,17 @@ function EndpointRowActions({ endpoint, onEndpointClick, onDeleteClick }: Endpoi
|
|||
<DropdownMenuItem
|
||||
variant="destructive"
|
||||
data-testid="endpoint-action-delete"
|
||||
disabled={!endpointId}
|
||||
onClick={() => endpointId && onDeleteClick(endpointId)}
|
||||
disabled={isFromConfig || !endpointId}
|
||||
onClick={() => !isFromConfig && endpointId && onDeleteClick(endpointId)}
|
||||
>
|
||||
<Trash2 />
|
||||
Delete
|
||||
</DropdownMenuItem>
|
||||
{isFromConfig && (
|
||||
<div data-testid="endpoint-config-hint" className="px-2 py-1.5 text-xs text-muted-foreground">
|
||||
{CONFIG_ENDPOINT_HINT}
|
||||
</div>
|
||||
)}
|
||||
</DropdownMenuContent>
|
||||
</DropdownMenu>
|
||||
);
|
||||
|
|
@ -124,7 +133,9 @@ export const getPassThroughEndpointsTableColumns = ({
|
|||
enableSorting: false,
|
||||
cell: ({ row }) => {
|
||||
const endpointId = row.original.id;
|
||||
if (!endpointId) return <span className="font-mono text-xs text-muted-foreground">—</span>;
|
||||
if (!endpointId || row.original.is_from_config) {
|
||||
return <span className="font-mono text-xs text-muted-foreground">—</span>;
|
||||
}
|
||||
return (
|
||||
<IdentityCell
|
||||
title={endpointId}
|
||||
|
|
@ -134,6 +145,17 @@ export const getPassThroughEndpointsTableColumns = ({
|
|||
);
|
||||
},
|
||||
},
|
||||
{
|
||||
id: "source",
|
||||
meta: { title: "Source", skeleton: "badge" },
|
||||
header: "Source",
|
||||
size: 100,
|
||||
enableSorting: false,
|
||||
cell: ({ row }) => {
|
||||
const isFromConfig = row.original.is_from_config ?? false;
|
||||
return <StatusBadge tone={isFromConfig ? "neutral" : "info"} label={isFromConfig ? "Config" : "DB"} />;
|
||||
},
|
||||
},
|
||||
{
|
||||
id: "path",
|
||||
accessorKey: "path",
|
||||
|
|
|
|||
|
|
@ -1,4 +1,13 @@
|
|||
import React, { useState, useEffect } from "react";
|
||||
import {
|
||||
AlertDialog,
|
||||
AlertDialogCancel,
|
||||
AlertDialogContent,
|
||||
AlertDialogDescription,
|
||||
AlertDialogFooter,
|
||||
AlertDialogHeader,
|
||||
AlertDialogTitle,
|
||||
} from "@/components/ui/alert-dialog";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { deletePassThroughEndpointsCall, getPassThroughEndpointsCall } from "../networking";
|
||||
import AddPassThroughEndpoint from "../add_pass_through";
|
||||
|
|
@ -25,6 +34,7 @@ export interface passThroughItem {
|
|||
methods?: string[];
|
||||
guardrails?: Record<string, { request_fields?: string[]; response_fields?: string[] } | null>;
|
||||
default_query_params?: Record<string, string>;
|
||||
is_from_config?: boolean;
|
||||
}
|
||||
|
||||
const PassThroughSettings: React.FC<PassThroughSettingsProps> = ({ accessToken, userRole, userID, premiumUser }) => {
|
||||
|
|
@ -133,42 +143,22 @@ const PassThroughSettings: React.FC<PassThroughSettingsProps> = ({ accessToken,
|
|||
onDeleteClick={handleDelete}
|
||||
/>
|
||||
|
||||
{isDeleteModalOpen && (
|
||||
<div className="fixed z-overlay inset-0 overflow-y-auto">
|
||||
<div className="flex items-end justify-center min-h-screen pt-4 px-4 pb-20 text-center sm:block sm:p-0">
|
||||
<div className="fixed inset-0 transition-opacity" aria-hidden="true">
|
||||
<div className="absolute inset-0 bg-gray-500 opacity-75"></div>
|
||||
</div>
|
||||
|
||||
<span className="hidden sm:inline-block sm:align-middle sm:h-screen" aria-hidden="true">
|
||||
​
|
||||
</span>
|
||||
|
||||
<div className="inline-block align-bottom bg-card rounded-lg text-left overflow-hidden shadow-xl transform transition-all sm:my-8 sm:align-middle sm:max-w-lg sm:w-full">
|
||||
<div className="bg-card px-4 pt-5 pb-4 sm:p-6 sm:pb-4">
|
||||
<div className="sm:flex sm:items-start">
|
||||
<div className="mt-3 text-center sm:mt-0 sm:ml-4 sm:text-left">
|
||||
<h3 className="text-lg leading-6 font-medium text-foreground">Delete Pass-Through Endpoint</h3>
|
||||
<div className="mt-2">
|
||||
<p className="text-sm text-muted-foreground">
|
||||
Are you sure you want to delete this pass-through endpoint? This action cannot be undone.
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<div className="bg-muted px-4 py-3 sm:px-6 sm:flex sm:flex-row-reverse">
|
||||
<Button variant="destructive" onClick={confirmDelete} className="ml-2">
|
||||
Delete
|
||||
</Button>
|
||||
<Button variant="outline" onClick={cancelDelete}>
|
||||
Cancel
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
<AlertDialog open={isDeleteModalOpen} onOpenChange={(open) => !open && cancelDelete()}>
|
||||
<AlertDialogContent>
|
||||
<AlertDialogHeader>
|
||||
<AlertDialogTitle>Delete Pass-Through Endpoint</AlertDialogTitle>
|
||||
<AlertDialogDescription>
|
||||
Are you sure you want to delete this pass-through endpoint? This action cannot be undone.
|
||||
</AlertDialogDescription>
|
||||
</AlertDialogHeader>
|
||||
<AlertDialogFooter>
|
||||
<AlertDialogCancel>Cancel</AlertDialogCancel>
|
||||
<Button variant="destructive" onClick={confirmDelete}>
|
||||
Delete
|
||||
</Button>
|
||||
</AlertDialogFooter>
|
||||
</AlertDialogContent>
|
||||
</AlertDialog>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
113
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
113
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -3318,11 +3318,14 @@ export interface paths {
|
|||
* - model: Model name (e.g., "gpt-4", "claude-3-opus")
|
||||
* - input_tokens: Expected input tokens per request
|
||||
* - output_tokens: Expected output tokens per request
|
||||
* - cache_read_input_tokens: Cache-read tokens per request, counted within input_tokens (optional)
|
||||
* - cache_creation_input_tokens: Cache-write tokens per request, counted within input_tokens (optional)
|
||||
* - reasoning_tokens: Reasoning tokens per request, counted within output_tokens (optional)
|
||||
* - num_requests_per_day: Number of requests per day (optional)
|
||||
* - num_requests_per_month: Number of requests per month (optional)
|
||||
*
|
||||
* Returns cost breakdown including:
|
||||
* - Per-request costs (input, output, margin)
|
||||
* - Per-request costs (input, output, margin, plus the cache-read, cache-write and reasoning shares)
|
||||
* - Daily costs (if num_requests_per_day provided)
|
||||
* - Monthly costs (if num_requests_per_month provided)
|
||||
*
|
||||
|
|
@ -3331,7 +3334,9 @@ export interface paths {
|
|||
* {
|
||||
* "model": "gpt-4",
|
||||
* "input_tokens": 1000,
|
||||
* "cache_read_input_tokens": 800,
|
||||
* "output_tokens": 500,
|
||||
* "reasoning_tokens": 200,
|
||||
* "num_requests_per_day": 100,
|
||||
* "num_requests_per_month": 3000
|
||||
* }
|
||||
|
|
@ -26468,6 +26473,18 @@ export interface components {
|
|||
* @description Request body for /cost/estimate endpoint.
|
||||
*/
|
||||
CostEstimateRequest: {
|
||||
/**
|
||||
* Cache Creation Input Tokens
|
||||
* @description Input tokens written to the prompt cache; counted within input_tokens
|
||||
* @default 0
|
||||
*/
|
||||
cache_creation_input_tokens: number;
|
||||
/**
|
||||
* Cache Read Input Tokens
|
||||
* @description Input tokens read from the prompt cache; counted within input_tokens
|
||||
* @default 0
|
||||
*/
|
||||
cache_read_input_tokens: number;
|
||||
/**
|
||||
* Input Tokens
|
||||
* @description Expected input tokens per request
|
||||
|
|
@ -26493,17 +26510,65 @@ export interface components {
|
|||
* @description Expected output tokens per request
|
||||
*/
|
||||
output_tokens: number;
|
||||
/**
|
||||
* Reasoning Tokens
|
||||
* @description Reasoning tokens the model emits; counted within output_tokens
|
||||
* @default 0
|
||||
*/
|
||||
reasoning_tokens: number;
|
||||
};
|
||||
/**
|
||||
* CostEstimateResponse
|
||||
* @description Response body for /cost/estimate endpoint.
|
||||
*/
|
||||
CostEstimateResponse: {
|
||||
/**
|
||||
* Cache Creation Cost Per Request
|
||||
* @description Cache-write share of input_cost_per_request
|
||||
* @default 0
|
||||
*/
|
||||
cache_creation_cost_per_request: number;
|
||||
/**
|
||||
* Cache Creation Input Token Cost
|
||||
* @description Rate billed per cache-write token
|
||||
*/
|
||||
cache_creation_input_token_cost?: number | null;
|
||||
/**
|
||||
* Cache Creation Input Tokens
|
||||
* @default 0
|
||||
*/
|
||||
cache_creation_input_tokens: number;
|
||||
/**
|
||||
* Cache Read Cost Per Request
|
||||
* @description Cache-read share of input_cost_per_request
|
||||
* @default 0
|
||||
*/
|
||||
cache_read_cost_per_request: number;
|
||||
/**
|
||||
* Cache Read Input Token Cost
|
||||
* @description Rate billed per cache-read token
|
||||
*/
|
||||
cache_read_input_token_cost?: number | null;
|
||||
/**
|
||||
* Cache Read Input Tokens
|
||||
* @default 0
|
||||
*/
|
||||
cache_read_input_tokens: number;
|
||||
/**
|
||||
* Cost Per Request
|
||||
* @description Total cost per request (includes margin)
|
||||
*/
|
||||
cost_per_request: number;
|
||||
/**
|
||||
* Daily Cache Creation Cost
|
||||
* @description Cache-write share of daily_input_cost
|
||||
*/
|
||||
daily_cache_creation_cost?: number | null;
|
||||
/**
|
||||
* Daily Cache Read Cost
|
||||
* @description Cache-read share of daily_input_cost
|
||||
*/
|
||||
daily_cache_read_cost?: number | null;
|
||||
/**
|
||||
* Daily Cost
|
||||
* @description Total daily cost (includes margin)
|
||||
|
|
@ -26524,12 +26589,20 @@ export interface components {
|
|||
* @description Daily output token cost
|
||||
*/
|
||||
daily_output_cost?: number | null;
|
||||
/**
|
||||
* Daily Reasoning Cost
|
||||
* @description Reasoning share of daily_output_cost
|
||||
*/
|
||||
daily_reasoning_cost?: number | null;
|
||||
/**
|
||||
* Input Cost Per Request
|
||||
* @description Input token cost per request (before margin)
|
||||
*/
|
||||
input_cost_per_request: number;
|
||||
/** Input Cost Per Token */
|
||||
/**
|
||||
* Input Cost Per Token
|
||||
* @description Rate billed per input token
|
||||
*/
|
||||
input_cost_per_token?: number | null;
|
||||
/** Input Tokens */
|
||||
input_tokens: number;
|
||||
|
|
@ -26541,6 +26614,16 @@ export interface components {
|
|||
margin_cost_per_request: number;
|
||||
/** Model */
|
||||
model: string;
|
||||
/**
|
||||
* Monthly Cache Creation Cost
|
||||
* @description Cache-write share of monthly_input_cost
|
||||
*/
|
||||
monthly_cache_creation_cost?: number | null;
|
||||
/**
|
||||
* Monthly Cache Read Cost
|
||||
* @description Cache-read share of monthly_input_cost
|
||||
*/
|
||||
monthly_cache_read_cost?: number | null;
|
||||
/**
|
||||
* Monthly Cost
|
||||
* @description Total monthly cost (includes margin)
|
||||
|
|
@ -26561,21 +26644,45 @@ export interface components {
|
|||
* @description Monthly output token cost
|
||||
*/
|
||||
monthly_output_cost?: number | null;
|
||||
/**
|
||||
* Monthly Reasoning Cost
|
||||
* @description Reasoning share of monthly_output_cost
|
||||
*/
|
||||
monthly_reasoning_cost?: number | null;
|
||||
/** Num Requests Per Day */
|
||||
num_requests_per_day?: number | null;
|
||||
/** Num Requests Per Month */
|
||||
num_requests_per_month?: number | null;
|
||||
/**
|
||||
* Output Cost Per Reasoning Token
|
||||
* @description Rate billed per reasoning token
|
||||
*/
|
||||
output_cost_per_reasoning_token?: number | null;
|
||||
/**
|
||||
* Output Cost Per Request
|
||||
* @description Output token cost per request (before margin)
|
||||
*/
|
||||
output_cost_per_request: number;
|
||||
/** Output Cost Per Token */
|
||||
/**
|
||||
* Output Cost Per Token
|
||||
* @description Rate billed per output token
|
||||
*/
|
||||
output_cost_per_token?: number | null;
|
||||
/** Output Tokens */
|
||||
output_tokens: number;
|
||||
/** Provider */
|
||||
provider?: string | null;
|
||||
/**
|
||||
* Reasoning Cost Per Request
|
||||
* @description Reasoning share of output_cost_per_request
|
||||
* @default 0
|
||||
*/
|
||||
reasoning_cost_per_request: number;
|
||||
/**
|
||||
* Reasoning Tokens
|
||||
* @default 0
|
||||
*/
|
||||
reasoning_tokens: number;
|
||||
};
|
||||
/** CreateCredentialItem */
|
||||
CreateCredentialItem: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue