mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge remote-tracking branch 'origin/main' into litellm_per_turn_control_beta
This commit is contained in:
commit
ce40b5773d
143 changed files with 10386 additions and 1527 deletions
3
.github/pull_request_template.md
vendored
3
.github/pull_request_template.md
vendored
|
|
@ -101,7 +101,8 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac
|
|||
For bug fixes: Before shows the reproduction, After shows the same steps passing
|
||||
For new features: Before shows the capability missing, After shows it working end-to-end
|
||||
If the change applies to all three LLM endpoints (/v1/responses, /v1/chat/completions, /v1/messages), make each endpoint its own case, not just one
|
||||
For UI changes: before/after screenshots under the same headings -->
|
||||
For UI changes: before/after screenshots under the same headings
|
||||
If the main use case runs through a coding tool like Claude Code or Codex, drive that tool interactively the way the user does (never `claude -p`, `codex exec`, or curl on its own) and embed before/after screenshots of its pane under the same headings; curl replays and headless runs can follow as extra cases, never as the only proof -->
|
||||
|
||||
## Type
|
||||
|
||||
|
|
|
|||
|
|
@ -538,7 +538,7 @@ context_window_fallbacks: Optional[List] = None
|
|||
content_policy_fallbacks: Optional[List] = None
|
||||
allowed_fails: int = 3
|
||||
allow_dynamic_callback_disabling: bool = True
|
||||
num_retries_per_request: Optional[int] = None # cap on Router retries of one model group; resets per fallback hop
|
||||
num_retries_per_request: Optional[int] = None # for the request overall (incl. fallbacks + model retries)
|
||||
####### SECRET MANAGERS #####################
|
||||
secret_manager_client: Optional[Any] = (
|
||||
None # list of instantiated key management clients - e.g. azure kv, infisical, etc.
|
||||
|
|
|
|||
125
litellm/caching/affinity_cache.py
Normal file
125
litellm/caching/affinity_cache.py
Normal file
|
|
@ -0,0 +1,125 @@
|
|||
"""Atomic affinity claims shared by deployment and tier-model selection."""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import (
|
||||
Final,
|
||||
cast, # noqa: TID251 # Redis script results are narrowed only to object, then validated
|
||||
)
|
||||
|
||||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
|
||||
_PIN_JSON_ADAPTER: Final = TypeAdapter[JsonValue](JsonValue)
|
||||
|
||||
_CLAIM_PIN_SCRIPT: Final = """
|
||||
local current = redis.call('GET', KEYS[1])
|
||||
if current == false then
|
||||
redis.call('SET', KEYS[1], ARGV[1], 'EX', ARGV[2])
|
||||
return ARGV[1]
|
||||
end
|
||||
if ARGV[3] then
|
||||
local decoded, stored = pcall(cjson.decode, current)
|
||||
if decoded and type(stored) == 'table' then
|
||||
for _, eligible in ipairs(cjson.decode(ARGV[3])) do
|
||||
local matches = true
|
||||
for key, value in pairs(eligible) do
|
||||
if stored[key] ~= value then matches = false; break end
|
||||
end
|
||||
for key, _ in pairs(stored) do
|
||||
if eligible[key] == nil then matches = false; break end
|
||||
end
|
||||
if matches then
|
||||
redis.call('EXPIRE', KEYS[1], ARGV[2])
|
||||
return current
|
||||
end
|
||||
end
|
||||
end
|
||||
redis.call('SET', KEYS[1], ARGV[1], 'EX', ARGV[2])
|
||||
return ARGV[1]
|
||||
end
|
||||
if current == ARGV[1] then
|
||||
redis.call('EXPIRE', KEYS[1], ARGV[2])
|
||||
end
|
||||
return current
|
||||
"""
|
||||
|
||||
|
||||
def set_local_affinity_pin(cache: DualCache, cache_key: str, value: object, ttl_seconds: int) -> None:
|
||||
"""Replace the entry because InMemoryCache.set_cache preserves a live key's expiry."""
|
||||
cache.in_memory_cache.delete_cache(cache_key)
|
||||
cache.in_memory_cache.set_cache(cache_key, value, ttl=ttl_seconds)
|
||||
|
||||
|
||||
def _legacy_pin_matches(stored: object, pin_value: Mapping[str, str]) -> bool:
|
||||
if isinstance(stored, dict):
|
||||
return all(stored.get(key) is not None and str(stored[key]) == value for key, value in pin_value.items())
|
||||
return isinstance(stored, str) and len(pin_value) == 1 and stored in pin_value.values()
|
||||
|
||||
|
||||
def claim_affinity_pin_in_memory(
|
||||
cache: DualCache,
|
||||
cache_key: str,
|
||||
pin_value: Mapping[str, str],
|
||||
ttl_seconds: int,
|
||||
*,
|
||||
eligible_values: tuple[Mapping[str, str], ...] | None = None,
|
||||
) -> object:
|
||||
"""No await between read and write, so same-loop claims agree during a Redis outage."""
|
||||
existing: Final[object] = cache.in_memory_cache.get_cache(cache_key)
|
||||
if existing is not None and eligible_values is None:
|
||||
if _legacy_pin_matches(existing, pin_value):
|
||||
set_local_affinity_pin(cache, cache_key, pin_value, ttl_seconds)
|
||||
return existing
|
||||
winner: Final = existing if existing is not None and existing in (eligible_values or ()) else pin_value
|
||||
set_local_affinity_pin(cache, cache_key, winner, ttl_seconds)
|
||||
return winner
|
||||
|
||||
|
||||
def _decode_pin(value: str) -> object:
|
||||
try:
|
||||
return _PIN_JSON_ADAPTER.validate_json(value)
|
||||
except ValidationError:
|
||||
return value
|
||||
|
||||
|
||||
async def claim_affinity_pin(
|
||||
cache: DualCache,
|
||||
cache_key: str,
|
||||
pin_value: Mapping[str, str],
|
||||
ttl_seconds: int,
|
||||
*,
|
||||
eligible_values: tuple[Mapping[str, str], ...] | None = None,
|
||||
) -> object:
|
||||
"""Return the authoritative first writer, replacing it only when it becomes ineligible.
|
||||
|
||||
Eligible claims refresh the returned winner. Legacy deployment claims only refresh
|
||||
a matching candidate. Resolve Redis per call because the proxy attaches it lazily.
|
||||
"""
|
||||
redis_cache: Final = cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
claim_script: Final = redis_cache.async_register_script(_CLAIM_PIN_SCRIPT)
|
||||
args: Final = (
|
||||
json.dumps(dict(pin_value)), # mutable-ok: JSON serialization requires dict, not a generic Mapping
|
||||
int(ttl_seconds),
|
||||
*(
|
||||
(json.dumps(tuple(dict(value) for value in eligible_values)),) # mutable-ok: JSON requires dict
|
||||
if eligible_values is not None
|
||||
else ()
|
||||
),
|
||||
)
|
||||
raw: Final = cast( # cast-ok: Redis scripts return heterogeneous values; only object is asserted here
|
||||
object, await claim_script(keys=(cache_key,), args=args)
|
||||
)
|
||||
decoded: Final = raw.decode("utf-8") if isinstance(raw, bytes) else raw
|
||||
if not isinstance(decoded, str):
|
||||
return pin_value
|
||||
winner: Final = _decode_pin(decoded)
|
||||
set_local_affinity_pin(cache, cache_key, winner, ttl_seconds)
|
||||
return winner
|
||||
except Exception as error: # noqa: BLE001 # Redis/Lua faults retain same-pod affinity through local claims
|
||||
verbose_router_logger.debug("Affinity cache: Redis claim failed, using pod-local claim. error=%s", error)
|
||||
return claim_affinity_pin_in_memory(cache, cache_key, pin_value, ttl_seconds, eligible_values=eligible_values)
|
||||
|
|
@ -1976,6 +1976,8 @@ BROWSER_SECURITY_HEADERS: Final[frozenset[str]] = frozenset(
|
|||
|
||||
UNSAFE_PROXY_RESPONSE_HEADERS: Final[frozenset[str]] = HTTP_FRAMING_HEADERS | BROWSER_SECURITY_HEADERS
|
||||
|
||||
STRINGIFIED_NONE: Final[str] = "None"
|
||||
|
||||
# A retrieved response replays the usage of the call that created it, so pricing these
|
||||
# read/management routes like inference bills the same tokens twice.
|
||||
NON_INFERENCE_CALL_TYPES: Final[frozenset[str]] = frozenset(
|
||||
|
|
|
|||
|
|
@ -309,8 +309,8 @@ def max_retries_per_request_hit(kwargs: Mapping[str, object], num_retries_per_re
|
|||
metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs))
|
||||
if not isinstance(metadata, Mapping):
|
||||
return False
|
||||
attempted_retries: Final = metadata.get("attempted_retries")
|
||||
return type(attempted_retries) is int and 0 < attempted_retries and num_retries_per_request <= attempted_retries
|
||||
retry_count: Final = metadata.get("request_retry_count")
|
||||
return type(retry_count) is int and 0 < retry_count and num_retries_per_request <= retry_count
|
||||
|
||||
|
||||
def get_or_create_metadata_bucket(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
import inspect
|
||||
import json
|
||||
import re
|
||||
import traceback
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Protocol, cast
|
||||
|
||||
import httpx
|
||||
|
|
@ -202,11 +205,17 @@ def _get_response_headers(original_exception: Exception) -> httpx.Headers | None
|
|||
return _response_headers
|
||||
|
||||
|
||||
def _accepted_init_kwargs(exception_class: type[Exception], candidates: Mapping[str, object]) -> Mapping[str, object]:
|
||||
accepted: Final = inspect.signature(exception_class).parameters
|
||||
return MappingProxyType({name: value for name, value in candidates.items() if name in accepted})
|
||||
|
||||
|
||||
def extract_and_raise_litellm_exception(
|
||||
response: Any | None,
|
||||
error_str: str,
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
body: object | None = None,
|
||||
):
|
||||
"""
|
||||
Covers scenario where litellm sdk calling proxy.
|
||||
|
|
@ -216,32 +225,19 @@ def extract_and_raise_litellm_exception(
|
|||
Relevant Issue: https://github.com/BerriAI/litellm/issues/7259
|
||||
"""
|
||||
pattern: Final = r"litellm\.\w+Error"
|
||||
|
||||
# Search for the exception in the error string
|
||||
match: Final = re.search(pattern, error_str)
|
||||
|
||||
# Extract the exception if found
|
||||
if match:
|
||||
exception_name = match.group(0)
|
||||
exception_name = exception_name.strip().replace("litellm.", "")
|
||||
raised_exception_obj: Final = getattr(litellm, exception_name, None)
|
||||
if raised_exception_obj:
|
||||
# Try with response parameter first, fall back to without it
|
||||
# Some exceptions (e.g., APIConnectionError) don't accept response param
|
||||
try:
|
||||
raise raised_exception_obj(
|
||||
message=error_str,
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=response,
|
||||
)
|
||||
except TypeError:
|
||||
# Exception doesn't accept response parameter
|
||||
raise raised_exception_obj(
|
||||
message=error_str,
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
)
|
||||
if match is None:
|
||||
return
|
||||
exception_name: Final = match.group(0).removeprefix("litellm.")
|
||||
raised_exception_obj: Final = getattr(litellm, exception_name, None)
|
||||
if not raised_exception_obj:
|
||||
return
|
||||
raise raised_exception_obj(
|
||||
message=error_str,
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
**_accepted_init_kwargs(raised_exception_obj, MappingProxyType({"response": response, "body": body})),
|
||||
)
|
||||
|
||||
|
||||
class _ProviderHTTPException(Protocol):
|
||||
|
|
@ -254,6 +250,23 @@ class _ProviderHTTPException(Protocol):
|
|||
llm_provider: str
|
||||
|
||||
|
||||
def _litellm_proxy_response(
|
||||
original_exception: _ProviderHTTPException, custom_llm_provider: str
|
||||
) -> httpx.Response | None:
|
||||
response: Final = getattr(original_exception, "response", None)
|
||||
if custom_llm_provider != "litellm_proxy" or not isinstance(response, httpx.Response) or response.headers:
|
||||
return response
|
||||
headers: Final = getattr(original_exception, "headers", None)
|
||||
if not isinstance(headers, Mapping) or not headers:
|
||||
return response
|
||||
pairs: Final = headers.multi_items() if isinstance(headers, httpx.Headers) else headers.items()
|
||||
return httpx.Response(
|
||||
status_code=response.status_code,
|
||||
headers=[(str(k), str(v)) for k, v in pairs],
|
||||
request=getattr(original_exception, "request", None),
|
||||
)
|
||||
|
||||
|
||||
def _map_openai_exception(
|
||||
*,
|
||||
model: str,
|
||||
|
|
@ -264,6 +277,7 @@ def _map_openai_exception(
|
|||
exception_provider: str,
|
||||
extra_information: str,
|
||||
) -> None:
|
||||
response: Final = _litellm_proxy_response(original_exception, custom_llm_provider)
|
||||
# custom_llm_provider is openai, make it OpenAI
|
||||
message = get_error_message(error_obj=original_exception)
|
||||
if message is None:
|
||||
|
|
@ -292,14 +306,14 @@ def _map_openai_exception(
|
|||
message=f"RateLimitError: {exception_provider} - {message}",
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
)
|
||||
elif ExceptionCheckers.is_error_str_context_window_exceeded(error_str):
|
||||
raise ContextWindowExceededError(
|
||||
message=f"ContextWindowExceededError: {exception_provider} - {message}",
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif "invalid_request_error" in error_str and "model_not_found" in error_str:
|
||||
|
|
@ -307,7 +321,7 @@ def _map_openai_exception(
|
|||
message=f"{exception_provider} - {message}",
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif "A timeout occurred" in error_str:
|
||||
|
|
@ -326,8 +340,9 @@ def _map_openai_exception(
|
|||
message=f"ContentPolicyViolationError: {exception_provider} - {message}",
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
body=getattr(original_exception, "body", None),
|
||||
)
|
||||
elif "invalid_encrypted_content" in error_str or "could not be verified" in error_str:
|
||||
helpful_message: Final = (
|
||||
|
|
@ -345,7 +360,7 @@ def _map_openai_exception(
|
|||
message=helpful_message,
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
body=getattr(original_exception, "body", None),
|
||||
)
|
||||
|
|
@ -354,7 +369,7 @@ def _map_openai_exception(
|
|||
message=f"{exception_provider} - {message}",
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
body=getattr(original_exception, "body", None),
|
||||
)
|
||||
|
|
@ -372,7 +387,7 @@ def _map_openai_exception(
|
|||
message=f"RateLimitError: {exception_provider} - {message}",
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif (
|
||||
|
|
@ -383,7 +398,7 @@ def _map_openai_exception(
|
|||
message=f"AuthenticationError: {exception_provider} - {message}",
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif "Mistral API raised a streaming error" in error_str:
|
||||
|
|
@ -402,15 +417,16 @@ def _map_openai_exception(
|
|||
message=f"{exception_provider} - {message}",
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
body=getattr(original_exception, "body", None),
|
||||
)
|
||||
elif original_exception.status_code == 401:
|
||||
raise AuthenticationError(
|
||||
message=f"AuthenticationError: {exception_provider} - {message}",
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif original_exception.status_code == 404:
|
||||
|
|
@ -418,7 +434,7 @@ def _map_openai_exception(
|
|||
message=f"NotFoundError: {exception_provider} - {message}",
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif original_exception.status_code == 408:
|
||||
|
|
@ -433,7 +449,7 @@ def _map_openai_exception(
|
|||
message=f"{exception_provider} - {message}",
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
body=getattr(original_exception, "body", None),
|
||||
)
|
||||
|
|
@ -442,7 +458,7 @@ def _map_openai_exception(
|
|||
message=f"RateLimitError: {exception_provider} - {message}",
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif original_exception.status_code == 500:
|
||||
|
|
@ -450,7 +466,7 @@ def _map_openai_exception(
|
|||
message=f"InternalServerError: {exception_provider} - {message}",
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif original_exception.status_code == 502:
|
||||
|
|
@ -458,7 +474,7 @@ def _map_openai_exception(
|
|||
message=f"BadGatewayError: {exception_provider} - {message}",
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif original_exception.status_code == 503:
|
||||
|
|
@ -466,7 +482,7 @@ def _map_openai_exception(
|
|||
message=f"ServiceUnavailableError: {exception_provider} - {message}",
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=response,
|
||||
litellm_debug_info=extra_information,
|
||||
)
|
||||
elif original_exception.status_code == 504: # gateway timeout error
|
||||
|
|
@ -2423,10 +2439,11 @@ def exception_type(
|
|||
custom_llm_provider == "litellm_proxy"
|
||||
): # handle special case where calling litellm proxy + exception str contains error message
|
||||
extract_and_raise_litellm_exception(
|
||||
response=getattr(original_exception, "response", None),
|
||||
response=_litellm_proxy_response(mappable_exception, custom_llm_provider),
|
||||
error_str=error_str,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
body=getattr(original_exception, "body", None),
|
||||
)
|
||||
if (
|
||||
custom_llm_provider == "openai"
|
||||
|
|
|
|||
|
|
@ -484,11 +484,12 @@ def apply_off_peak_pricing(model_info: ModelInfo, current_time: datetime | None,
|
|||
def _apply_off_peak_to_base_costs(
|
||||
model_info: ModelInfo,
|
||||
current_time: datetime | None,
|
||||
base_costs: tuple[float, float, float, float, float],
|
||||
base_costs: tuple[float, float, float, float | None, float],
|
||||
) -> tuple[float, float, float, float, float]:
|
||||
"""Apply off-peak rates to an already-resolved set of base costs, whichever pricing path
|
||||
produced them. The one-hour cache-creation rate passes through untouched, since
|
||||
off_peak_pricing has no field for it, and reasoning is left to _resolve_billed_reasoning_rate.
|
||||
produced them. off_peak_pricing has no field for the one-hour cache-creation rate, so a
|
||||
present one passes through untouched and an absent one resolves to the applied
|
||||
cache-creation rate. Reasoning is left to _resolve_billed_reasoning_rate.
|
||||
"""
|
||||
prompt, completion, cache_creation, cache_creation_above_1hr, cache_read = base_costs
|
||||
rates: Final = apply_off_peak_pricing(
|
||||
|
|
@ -506,7 +507,7 @@ def _apply_off_peak_to_base_costs(
|
|||
rates.input_rate,
|
||||
rates.output_rate,
|
||||
rates.cache_creation_rate,
|
||||
cache_creation_above_1hr,
|
||||
rates.cache_creation_rate if cache_creation_above_1hr is None else cache_creation_above_1hr,
|
||||
rates.cache_read_rate,
|
||||
)
|
||||
|
||||
|
|
@ -532,6 +533,11 @@ def _get_token_base_cost(
|
|||
`missing_cache_read_uses_input` resolves an absent cache-read rate to the resolved
|
||||
input rate instead of 0.0; an explicit 0.0 rate stays a real price either way.
|
||||
|
||||
An absent cache-creation rate always resolves to the resolved input rate, the way the
|
||||
tiered table and custom deployment pricing already do, since a provider that publishes
|
||||
no write price bills cache writes as ordinary input. An absent 1h write rate resolves
|
||||
to the cache-creation rate, off-peak included. An explicit 0.0 stays a real price for both.
|
||||
|
||||
Returns:
|
||||
Tuple[float, float, float, float] - (prompt_cost, completion_cost, cache_creation_cost, cache_read_cost)
|
||||
"""
|
||||
|
|
@ -554,10 +560,9 @@ def _get_token_base_cost(
|
|||
output_image_cost: Final = _get_cost_per_unit(model_info, "output_cost_per_image_token", None)
|
||||
if output_image_cost is not None:
|
||||
completion_base_cost = cast(float, output_image_cost)
|
||||
cache_creation_cost = cast(float, _get_cost_per_unit(model_info, cache_creation_cost_key))
|
||||
cache_creation_cost_above_1hr = cast(
|
||||
float,
|
||||
_get_cost_per_unit(model_info, "cache_creation_input_token_cost_above_1hr"),
|
||||
cache_creation_cost = _get_cost_per_unit(model_info, cache_creation_cost_key, default_value=None)
|
||||
cache_creation_cost_above_1hr = _get_cost_per_unit(
|
||||
model_info, "cache_creation_input_token_cost_above_1hr", default_value=None
|
||||
)
|
||||
cache_read_cost = _get_cost_per_unit(model_info, cache_read_cost_key, default_value=None)
|
||||
|
||||
|
|
@ -639,22 +644,10 @@ def _get_token_base_cost(
|
|||
else f"cache_read_input_token_cost_above_{threshold_str}_tokens"
|
||||
)
|
||||
|
||||
cache_creation_cost = cast(
|
||||
float,
|
||||
_get_cost_per_unit(
|
||||
model_info,
|
||||
cache_creation_tiered_key,
|
||||
cache_creation_cost,
|
||||
),
|
||||
)
|
||||
cache_creation_cost = _get_cost_per_unit(model_info, cache_creation_tiered_key, cache_creation_cost)
|
||||
|
||||
cache_creation_cost_above_1hr = cast(
|
||||
float,
|
||||
_get_cost_per_unit(
|
||||
model_info,
|
||||
cache_creation_1hr_tiered_key,
|
||||
cache_creation_cost_above_1hr,
|
||||
),
|
||||
cache_creation_cost_above_1hr = _get_cost_per_unit(
|
||||
model_info, cache_creation_1hr_tiered_key, cache_creation_cost_above_1hr
|
||||
)
|
||||
|
||||
cache_read_cost = _get_cost_per_unit(model_info, cache_read_tiered_key, cache_read_cost)
|
||||
|
|
@ -665,16 +658,16 @@ def _get_token_base_cost(
|
|||
except Exception:
|
||||
continue
|
||||
|
||||
input_rate_for_missing_cache_rates: Final = _off_peak_rate(
|
||||
_open_off_peak_block(model_info, current_time) or MappingProxyType({}),
|
||||
"input_cost_per_token",
|
||||
prompt_base_cost,
|
||||
)
|
||||
if cache_read_cost is None:
|
||||
cache_read_cost = (
|
||||
_off_peak_rate(
|
||||
_open_off_peak_block(model_info, current_time) or MappingProxyType({}),
|
||||
"input_cost_per_token",
|
||||
prompt_base_cost,
|
||||
)
|
||||
if missing_cache_read_uses_input
|
||||
else 0.0
|
||||
)
|
||||
cache_read_cost = input_rate_for_missing_cache_rates if missing_cache_read_uses_input else 0.0
|
||||
resolved_cache_creation_cost: Final = (
|
||||
input_rate_for_missing_cache_rates if cache_creation_cost is None else cache_creation_cost
|
||||
)
|
||||
|
||||
return _apply_off_peak_to_base_costs(
|
||||
model_info,
|
||||
|
|
@ -682,7 +675,7 @@ def _get_token_base_cost(
|
|||
(
|
||||
prompt_base_cost,
|
||||
completion_base_cost,
|
||||
cache_creation_cost,
|
||||
resolved_cache_creation_cost,
|
||||
cache_creation_cost_above_1hr,
|
||||
cache_read_cost,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from itertools import chain, repeat
|
|||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, cast, overload, runtime_checkable
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict, assert_never
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -44,6 +45,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
|||
scoped_structured_message_indices,
|
||||
stream_item_field,
|
||||
stream_item_fingerprint,
|
||||
unappliable_request_rewrite,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
|
||||
AnthropicPassthroughLoggingHandler,
|
||||
|
|
@ -103,9 +105,24 @@ class ToolResultBlockTextTarget:
|
|||
block_idx: int
|
||||
|
||||
|
||||
InputWriteBackTarget = (
|
||||
MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget
|
||||
)
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SystemStringTarget:
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SystemBlockTextTarget:
|
||||
block_idx: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ToolUseInputTarget:
|
||||
msg_idx: int
|
||||
content_idx: int
|
||||
|
||||
|
||||
MessageTextTarget = MessageContentTarget | ContentBlockTextTarget | ToolResultStringTarget | ToolResultBlockTextTarget
|
||||
InputWriteBackTarget = SystemStringTarget | SystemBlockTextTarget | MessageTextTarget
|
||||
|
||||
|
||||
def _as_str_mapping(value: Mapping[str, object]) -> Mapping[str, object]:
|
||||
|
|
@ -146,10 +163,17 @@ class ScannedText:
|
|||
target: InputWriteBackTarget
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ScannedToolCall:
|
||||
tool_call: ChatCompletionToolCallChunk
|
||||
target: ToolUseInputTarget
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ExtractedInput:
|
||||
scanned: tuple[ScannedText, ...]
|
||||
images: tuple[str, ...]
|
||||
tool_calls: tuple[ScannedToolCall, ...] = ()
|
||||
|
||||
|
||||
EMPTY_EXTRACTED_INPUT: Final = ExtractedInput(scanned=(), images=())
|
||||
|
|
@ -161,6 +185,74 @@ class _ToolCallShape:
|
|||
arguments: str
|
||||
|
||||
|
||||
def _is_client_tool_use(block: Mapping[str, object]) -> bool:
|
||||
return (
|
||||
block.get("type") == "tool_use"
|
||||
and isinstance(block.get("id"), str)
|
||||
and isinstance(block.get("name"), str)
|
||||
and isinstance(block.get("input"), dict)
|
||||
)
|
||||
|
||||
|
||||
def _write_back_system_block(system: object, block_idx: int, response: str) -> None:
|
||||
if not isinstance(system, list):
|
||||
return
|
||||
text_blocks: Final = tuple(block for block in system if isinstance(block, dict) and block.get("type") == "text")
|
||||
if block_idx < len(text_blocks):
|
||||
text_blocks[block_idx]["text"] = (
|
||||
response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
|
||||
|
||||
def _write_back_message_text(message: _WritableMessage, target: MessageTextTarget, response: str) -> None:
|
||||
content: Final = message.get("content", None)
|
||||
if content is None:
|
||||
return
|
||||
match target:
|
||||
case MessageContentTarget():
|
||||
if isinstance(content, str):
|
||||
message["content"] = response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
case ContentBlockTextTarget(content_idx=content_idx):
|
||||
if isinstance(content, list):
|
||||
content[content_idx]["text"] = (
|
||||
response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case ToolResultStringTarget(content_idx=content_idx):
|
||||
if isinstance(content, list):
|
||||
content[content_idx]["content"] = (
|
||||
response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx):
|
||||
if isinstance(content, list):
|
||||
content[content_idx]["content"][block_idx]["text"] = (
|
||||
response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case _:
|
||||
assert_never(target)
|
||||
|
||||
|
||||
_TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def _rewritten_tool_use_input(arguments: str) -> Mapping[str, object] | None:
|
||||
try:
|
||||
return _TOOL_USE_INPUT_ADAPTER.validate_json(arguments)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _write_back_tool_use(
|
||||
message: _WritableMessage, target: ToolUseInputTarget, shape: _ToolCallShape, rewritten_input: Mapping[str, object]
|
||||
) -> None:
|
||||
content: Final = message.get("content", None)
|
||||
block: Final = content[target.content_idx] if isinstance(content, list) else None
|
||||
if not isinstance(block, dict):
|
||||
return
|
||||
block["input"] = rewritten_input # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
if shape.name is not None and shape.name != block.get("name"):
|
||||
block["name"] = shape.name # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _SSEFieldRewrite:
|
||||
"""One field of one nested section of a buffered SSE event, rewritten."""
|
||||
|
|
@ -452,9 +544,8 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
skip_tool: Final = effective_skip_tool_message_for_guardrail(guardrail_to_apply)
|
||||
scan_only_tool_results: Final = effective_scan_only_tool_results_for_guardrail(guardrail_to_apply)
|
||||
|
||||
# Exclude only the trusted top-level prompt. In-sequence system entries are untrusted
|
||||
# and must stay aligned with texts_to_check for positional masking. When the top-level
|
||||
# prompt is included, the pre-existing count mismatch disables positional masking.
|
||||
# The top-level prompt is translated on its own below so it can be hoisted in front of
|
||||
# any mid-turn system entries and scanned first, aligned with that structured position.
|
||||
translation_source: Final = { # mutable-ok: API message payload
|
||||
key: value for key, value in data.items() if key != "system"
|
||||
}
|
||||
|
|
@ -490,7 +581,12 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
]
|
||||
)
|
||||
|
||||
# Step 1: Extract all text content and images
|
||||
# Step 1: Extract all text content, images, and tool calls
|
||||
top_level_system_scanned: Final = (
|
||||
()
|
||||
if hoisted_system_message is None or scan_only_tool_results
|
||||
else self._extract_top_level_system_text(hoisted_system_message)
|
||||
)
|
||||
extracted: Final = tuple(
|
||||
self._extract_input_text_and_images(
|
||||
message=message,
|
||||
|
|
@ -501,17 +597,27 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
)
|
||||
for msg_idx, message in enumerate(messages)
|
||||
)
|
||||
scanned: Final = tuple(item for one_message in extracted for item in one_message.scanned)
|
||||
scanned: Final = (
|
||||
*top_level_system_scanned,
|
||||
*(item for one_message in extracted for item in one_message.scanned),
|
||||
)
|
||||
texts_to_check: Final = [item.text for item in scanned] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
|
||||
images_to_check: Final = [
|
||||
image for one_message in extracted for image in one_message.images
|
||||
] # mutable-ok: GenericGuardrailAPIInputs takes list[str]
|
||||
scanned_tool_calls: Final = tuple(item for one_message in extracted for item in one_message.tool_calls)
|
||||
tool_calls_to_check: Final = [
|
||||
item.tool_call for item in scanned_tool_calls
|
||||
] # mutable-ok: GenericGuardrailAPIInputs takes list[ChatCompletionToolCallChunk]
|
||||
pre_guardrail_tool_calls: Final = _tool_call_shapes(tool_calls_to_check)
|
||||
|
||||
# Step 2: Apply guardrail to all texts in batch
|
||||
if texts_to_check:
|
||||
# Step 2: Apply guardrail to all texts and tool calls in batch
|
||||
if texts_to_check or tool_calls_to_check:
|
||||
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
if tool_calls_to_check:
|
||||
inputs["tool_calls"] = tool_calls_to_check
|
||||
if tools_to_check:
|
||||
inputs["tools"] = tools_to_check
|
||||
original_structured_messages: Final = structured_messages
|
||||
|
|
@ -570,9 +676,18 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
preserve_system_messages=has_midturn_system_message,
|
||||
)
|
||||
else:
|
||||
if guardrailed_texts and len(guardrailed_texts) != len(scanned):
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
self._apply_guardrail_tool_calls_to_input(
|
||||
messages=messages,
|
||||
scanned_tool_calls=scanned_tool_calls,
|
||||
pre_guardrail_tool_calls=pre_guardrail_tool_calls,
|
||||
returned_tool_calls=guardrailed_inputs.get("tool_calls"),
|
||||
guardrail_name=guardrail_to_apply.guardrail_name,
|
||||
)
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=messages,
|
||||
data=data,
|
||||
responses=guardrailed_texts,
|
||||
scanned=scanned,
|
||||
)
|
||||
|
|
@ -598,6 +713,19 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
hoisted: Final = probe.get("messages") or [] # mutable-ok: API message payload
|
||||
return hoisted[0] if hoisted else None
|
||||
|
||||
@staticmethod
|
||||
def _extract_top_level_system_text(hoisted_system_message: AllMessageValues) -> tuple[ScannedText, ...]:
|
||||
content: Final = hoisted_system_message.get("content")
|
||||
if isinstance(content, str):
|
||||
return (ScannedText(content, SystemStringTarget()),)
|
||||
if not isinstance(content, list):
|
||||
return ()
|
||||
return tuple(
|
||||
ScannedText(text_str, SystemBlockTextTarget(block_idx))
|
||||
for block_idx, block in enumerate(content)
|
||||
if isinstance(block, dict) and isinstance(text_str := block.get("text"), str)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _openai_system_message_to_anthropic(
|
||||
message: Mapping[str, object],
|
||||
|
|
@ -852,9 +980,25 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
for content_idx, content_item in enumerate(content)
|
||||
if isinstance(content_item, dict)
|
||||
)
|
||||
tool_use_blocks: Final = (
|
||||
()
|
||||
if scan_only_tool_results
|
||||
else tuple(
|
||||
(content_idx, content_item)
|
||||
for content_idx, content_item in enumerate(content)
|
||||
if isinstance(content_item, dict) and _is_client_tool_use(content_item)
|
||||
)
|
||||
)
|
||||
return ExtractedInput(
|
||||
scanned=tuple(item for block in blocks for item in block.scanned),
|
||||
images=tuple(image for block in blocks for image in block.images),
|
||||
tool_calls=tuple(
|
||||
ScannedToolCall(
|
||||
tool_call=AnthropicConfig.convert_tool_use_to_openai_format(content_item, tool_call_idx),
|
||||
target=ToolUseInputTarget(msg_idx, content_idx),
|
||||
)
|
||||
for tool_call_idx, (content_idx, content_item) in enumerate(tool_use_blocks)
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
|
|
@ -940,43 +1084,59 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
async def _apply_guardrail_responses_to_input(
|
||||
self,
|
||||
messages: Sequence[_WritableMessage],
|
||||
responses: list[str],
|
||||
data: dict[str, object], # mutable-ok: API message payload
|
||||
responses: Sequence[str],
|
||||
scanned: tuple[ScannedText, ...],
|
||||
) -> None:
|
||||
"""
|
||||
Apply guardrail responses back to input messages.
|
||||
Apply guardrail responses back to the top-level system prompt and the input messages.
|
||||
"""
|
||||
raw_messages: Final = data.get("messages")
|
||||
messages: Final[Sequence[_WritableMessage]] = raw_messages if isinstance(raw_messages, list) else ()
|
||||
for item, guardrail_response in zip(scanned, responses):
|
||||
target = item.target
|
||||
message = messages[target.msg_idx]
|
||||
content = message.get("content", None)
|
||||
if content is None:
|
||||
continue
|
||||
|
||||
match target:
|
||||
case MessageContentTarget():
|
||||
if isinstance(content, str):
|
||||
message["content"] = (
|
||||
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case ContentBlockTextTarget(content_idx=content_idx):
|
||||
if isinstance(content, list):
|
||||
content[content_idx]["text"] = (
|
||||
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case ToolResultStringTarget(content_idx=content_idx):
|
||||
if isinstance(content, list):
|
||||
content[content_idx]["content"] = (
|
||||
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case ToolResultBlockTextTarget(content_idx=content_idx, block_idx=block_idx):
|
||||
if isinstance(content, list):
|
||||
content[content_idx]["content"][block_idx]["text"] = (
|
||||
match item.target:
|
||||
case SystemStringTarget():
|
||||
if isinstance(data.get("system"), str):
|
||||
data["system"] = (
|
||||
guardrail_response # mutable-ok: guardrails rewrite the caller's request payload in place
|
||||
)
|
||||
case SystemBlockTextTarget(block_idx=block_idx):
|
||||
_write_back_system_block(data.get("system"), block_idx, guardrail_response)
|
||||
case (
|
||||
MessageContentTarget()
|
||||
| ContentBlockTextTarget()
|
||||
| ToolResultStringTarget()
|
||||
| ToolResultBlockTextTarget() as message_target
|
||||
):
|
||||
_write_back_message_text(messages[message_target.msg_idx], message_target, guardrail_response)
|
||||
case _:
|
||||
assert_never(target)
|
||||
assert_never(item.target)
|
||||
|
||||
@staticmethod
|
||||
def _apply_guardrail_tool_calls_to_input(
|
||||
messages: Sequence[_WritableMessage],
|
||||
scanned_tool_calls: tuple[ScannedToolCall, ...],
|
||||
pre_guardrail_tool_calls: tuple[_ToolCallShape, ...],
|
||||
returned_tool_calls: Sequence[object] | None,
|
||||
guardrail_name: str | None,
|
||||
) -> None:
|
||||
post_guardrail_tool_calls: Final = _tool_call_shapes(
|
||||
returned_tool_calls
|
||||
if returned_tool_calls is not None and len(returned_tool_calls) == len(pre_guardrail_tool_calls)
|
||||
else tuple(item.tool_call for item in scanned_tool_calls)
|
||||
)
|
||||
rewritten: Final = tuple(
|
||||
(item, after, _rewritten_tool_use_input(after.arguments))
|
||||
for item, before, after in zip(scanned_tool_calls, pre_guardrail_tool_calls, post_guardrail_tool_calls)
|
||||
if before != after
|
||||
)
|
||||
applicable: Final = tuple(
|
||||
(item, after, rewritten_input) for item, after, rewritten_input in rewritten if rewritten_input is not None
|
||||
)
|
||||
if len(applicable) != len(rewritten):
|
||||
raise unappliable_request_rewrite(guardrail_name)
|
||||
for item, after, rewritten_input in applicable:
|
||||
_write_back_tool_use(messages[item.target.msg_idx], item.target, after, rewritten_input)
|
||||
|
||||
async def process_output_response(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Callable, Iterator, Sequence
|
||||
from typing import Final, TypeVar
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from typing import Final, TypeVar, cast # noqa: TID251 # a rebuilt chat row has no typed constructor across roles
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
|
@ -364,3 +364,67 @@ def merge_guardrailed_scoped_messages(
|
|||
yield from appended
|
||||
|
||||
return list(_merged())
|
||||
|
||||
|
||||
def _content_part_text(part: object) -> str | None:
|
||||
if not isinstance(part, Mapping):
|
||||
return None
|
||||
text: Final = part.get("text")
|
||||
return text if isinstance(text, str) else None
|
||||
|
||||
|
||||
def message_slot_texts(message: Mapping[str, object]) -> tuple[str, ...]:
|
||||
content: Final = message.get("content")
|
||||
if isinstance(content, str):
|
||||
return (content,)
|
||||
if isinstance(content, list):
|
||||
return tuple(text for part in content if (text := _content_part_text(part)) is not None)
|
||||
return ()
|
||||
|
||||
|
||||
def message_text_slot_count(message: AllMessageValues) -> int:
|
||||
return len(message_slot_texts(message))
|
||||
|
||||
|
||||
def _part_with_text(part: object, text: str) -> object:
|
||||
if not isinstance(part, Mapping):
|
||||
return part
|
||||
return {**part, "text": text} # mutable-ok: content parts stay JSON-plain dicts
|
||||
|
||||
|
||||
def _content_with_slot_texts(content: Sequence[object], texts: Sequence[str]) -> Sequence[object]:
|
||||
remaining_texts: Final = iter(texts)
|
||||
return [ # mutable-ok: message content stays a JSON list
|
||||
_part_with_text(part, next(remaining_texts)) if _content_part_text(part) is not None else part
|
||||
for part in content
|
||||
]
|
||||
|
||||
|
||||
def message_with_slot_texts(message: AllMessageValues, texts: Sequence[str]) -> AllMessageValues | None:
|
||||
"""Swap one rewritten text into each text slot of a chat row, in order.
|
||||
|
||||
A slot is a string ``content`` or one list part carrying a string ``text``;
|
||||
images and other parts ride along untouched. Returns None unless the counts
|
||||
line up exactly, so a rewrite never lands on the wrong slot.
|
||||
"""
|
||||
if message_text_slot_count(message) != len(texts):
|
||||
return None
|
||||
content: Final = message.get("content")
|
||||
if not isinstance(content, (str, list)):
|
||||
return message
|
||||
rewritten_content: Final = texts[0] if isinstance(content, str) else _content_with_slot_texts(content, texts)
|
||||
rewritten: Final = {**message, "content": rewritten_content} # mutable-ok: chat rows stay JSON-plain dicts
|
||||
return cast("AllMessageValues", rewritten) # cast-ok: the same row with only its text slots swapped
|
||||
|
||||
|
||||
class UnappliableRequestRewrite(Exception):
|
||||
def __init__(self, guardrail_name: str) -> None:
|
||||
super().__init__(
|
||||
f"Guardrail '{guardrail_name}' rewrote the request in a way this endpoint cannot apply, "
|
||||
"so the request was rejected rather than sent unrewritten"
|
||||
)
|
||||
self.guardrail_name: Final = guardrail_name
|
||||
|
||||
|
||||
def unappliable_request_rewrite(guardrail_name: str | None) -> UnappliableRequestRewrite:
|
||||
return UnappliableRequestRewrite(guardrail_name or "unknown")
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from concurrent.futures import ThreadPoolExecutor
|
|||
from datetime import datetime
|
||||
from functools import partial
|
||||
from threading import Lock
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload
|
||||
|
||||
import httpx
|
||||
|
|
@ -96,6 +97,77 @@ def _assume_role_params(
|
|||
)
|
||||
|
||||
|
||||
_SecureTransportBool = TypedDict("_SecureTransportBool", {"aws:SecureTransport": ReadOnly[Literal["true"]]})
|
||||
|
||||
|
||||
class _SecureTransportCondition(TypedDict):
|
||||
Bool: ReadOnly[_SecureTransportBool]
|
||||
|
||||
|
||||
class _SessionPolicyStatement(TypedDict):
|
||||
Sid: ReadOnly[str]
|
||||
Effect: ReadOnly[Literal["Allow"]]
|
||||
Action: ReadOnly[tuple[str, ...]]
|
||||
Resource: ReadOnly[Literal["*"]]
|
||||
Condition: ReadOnly[_SecureTransportCondition]
|
||||
|
||||
|
||||
class WebIdentitySessionPolicy(TypedDict):
|
||||
Version: ReadOnly[Literal["2012-10-17"]]
|
||||
Statement: ReadOnly[tuple[_SessionPolicyStatement, ...]]
|
||||
|
||||
|
||||
_WEB_IDENTITY_SESSION_POLICY_ACTIONS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
|
||||
{
|
||||
"BedrockLiteLLM": (
|
||||
"bedrock:InvokeModel",
|
||||
"bedrock:InvokeModelWithResponseStream",
|
||||
"bedrock:CountTokens",
|
||||
"bedrock:Rerank",
|
||||
"bedrock:Retrieve",
|
||||
"bedrock:ListKnowledgeBases",
|
||||
"bedrock:InvokeAgent",
|
||||
"bedrock:ApplyGuardrail",
|
||||
"bedrock:GetGuardrail",
|
||||
"bedrock:ListGuardrails",
|
||||
),
|
||||
"BedrockAgentCoreLiteLLM": (
|
||||
"bedrock-agentcore:InvokeAgentRuntime",
|
||||
"bedrock-agentcore:InvokeAgentRuntimeForUser",
|
||||
"bedrock-agentcore:InvokeGateway",
|
||||
),
|
||||
"ClaudePlatformLiteLLM": (
|
||||
"aws-external-anthropic:CreateInference",
|
||||
"aws-external-anthropic:CreateBatchInference",
|
||||
"aws-external-anthropic:CancelBatchInference",
|
||||
"aws-external-anthropic:DeleteBatchInference",
|
||||
"aws-external-anthropic:CountTokens",
|
||||
"aws-external-anthropic:Get*",
|
||||
"aws-external-anthropic:List*",
|
||||
),
|
||||
"BedrockMantleLiteLLM": ("bedrock-mantle:CreateInference",),
|
||||
}
|
||||
)
|
||||
|
||||
_SECURE_TRANSPORT_ONLY: Final = _SecureTransportCondition(Bool=_SecureTransportBool({"aws:SecureTransport": "true"}))
|
||||
|
||||
|
||||
def build_web_identity_session_policy() -> WebIdentitySessionPolicy:
|
||||
return WebIdentitySessionPolicy(
|
||||
Version="2012-10-17",
|
||||
Statement=tuple(
|
||||
_SessionPolicyStatement(
|
||||
Sid=sid,
|
||||
Effect="Allow",
|
||||
Action=actions,
|
||||
Resource="*",
|
||||
Condition=_SECURE_TRANSPORT_ONLY,
|
||||
)
|
||||
for sid, actions in _WEB_IDENTITY_SESSION_POLICY_ACTIONS.items()
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class BedrockRequestTarget(BaseModel):
|
||||
aws_region_name: str
|
||||
aws_bedrock_runtime_endpoint: str | None
|
||||
|
|
@ -940,60 +1012,12 @@ class BaseAWSLLM(SignsRequestsWithAWS):
|
|||
# auth only (static creds + IRSA take other code paths).
|
||||
# https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html
|
||||
# https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html
|
||||
bedrock_session_policy: Final = {
|
||||
"Version": "2012-10-17",
|
||||
"Statement": [
|
||||
{
|
||||
"Sid": "BedrockLiteLLM",
|
||||
"Effect": "Allow",
|
||||
"Action": [
|
||||
"bedrock:InvokeModel",
|
||||
"bedrock:InvokeModelWithResponseStream",
|
||||
"bedrock:CountTokens",
|
||||
"bedrock:ApplyGuardrail",
|
||||
"bedrock:GetGuardrail",
|
||||
"bedrock:ListGuardrails",
|
||||
],
|
||||
"Resource": "*",
|
||||
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
|
||||
},
|
||||
# Claude Platform on AWS (added by #27678 for the
|
||||
# ``bedrock/claude_platform/<model>`` route) lives under
|
||||
# a separate IAM action namespace; without these entries
|
||||
# the OIDC path 403s on every claude_platform request
|
||||
# even with a fully permissive identity policy (#30200).
|
||||
{
|
||||
"Sid": "ClaudePlatformLiteLLM",
|
||||
"Effect": "Allow",
|
||||
"Action": [
|
||||
"aws-external-anthropic:CreateInference",
|
||||
"aws-external-anthropic:CreateBatchInference",
|
||||
"aws-external-anthropic:CancelBatchInference",
|
||||
"aws-external-anthropic:DeleteBatchInference",
|
||||
"aws-external-anthropic:CountTokens",
|
||||
"aws-external-anthropic:Get*",
|
||||
"aws-external-anthropic:List*",
|
||||
],
|
||||
"Resource": "*",
|
||||
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
|
||||
},
|
||||
{
|
||||
"Sid": "BedrockMantleLiteLLM",
|
||||
"Effect": "Allow",
|
||||
"Action": [
|
||||
"bedrock-mantle:CreateInference",
|
||||
],
|
||||
"Resource": "*",
|
||||
"Condition": {"Bool": {"aws:SecureTransport": "true"}},
|
||||
},
|
||||
],
|
||||
}
|
||||
assume_role_params: Final = {
|
||||
"RoleArn": aws_role_name,
|
||||
"RoleSessionName": aws_session_name,
|
||||
"WebIdentityToken": oidc_token,
|
||||
"DurationSeconds": 3600,
|
||||
"Policy": json.dumps(bedrock_session_policy, separators=(",", ":")),
|
||||
"Policy": json.dumps(build_web_identity_session_policy(), separators=(",", ":")),
|
||||
}
|
||||
|
||||
# Add ExternalId parameter if provided
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ BaseAWSLLM._sign_request after the request body is finalized.
|
|||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, cast # noqa: TID251 # map_openai_params returns the filtered params as a bare dict
|
||||
|
||||
import httpx
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
|
@ -32,6 +32,7 @@ from litellm.llms.bedrock_mantle.common_utils import (
|
|||
BedrockMantleAuthMixin,
|
||||
)
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.responses.additional_tools import HoistedAdditionalTools, hoist_additional_tools
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseInputParam,
|
||||
|
|
@ -58,8 +59,6 @@ _BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES: Final = frozenset(
|
|||
_BEDROCK_MANTLE_SUPPORTED_SERVICE_TIERS: Final = frozenset({"auto", "default"})
|
||||
_BEDROCK_MANTLE_OPENAI_PATH_SUPPORTED_REASONING_SUMMARIES: Final = frozenset({"auto"})
|
||||
|
||||
_CODEX_ADDITIONAL_TOOLS_INPUT_ITEM_TYPE: Final = "additional_tools"
|
||||
|
||||
_CODEX_AGENT_MESSAGE_INPUT_ITEM_TYPE: Final = "agent_message"
|
||||
_CODEX_CONTEXT_COMPACTION_INPUT_ITEM_TYPE: Final = "context_compaction"
|
||||
_CODEX_LOCAL_SHELL_CALL_INPUT_ITEM_TYPE: Final = "local_shell_call"
|
||||
|
|
@ -233,17 +232,14 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
|
|||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
remaining_input, hoisted_tools = self._hoist_codex_additional_tools(input)
|
||||
normalized_input: Final = self._normalize_codex_input_items(remaining_input)
|
||||
params: Final = cast( # cast-ok: the base signature leaves the params dict untyped
|
||||
"ResponsesAPIOptionalRequestParams", response_api_optional_request_params
|
||||
)
|
||||
hoisted: Final = hoist_additional_tools(input, params.get("tools"))
|
||||
normalized_input: Final = self._normalize_codex_input_items(hoisted.input)
|
||||
request_params: Final = (
|
||||
{
|
||||
**response_api_optional_request_params,
|
||||
"tools": [
|
||||
*(response_api_optional_request_params.get("tools") or []),
|
||||
*hoisted_tools,
|
||||
],
|
||||
}
|
||||
if hoisted_tools
|
||||
self._params_with_hoisted_tools(params, hoisted)
|
||||
if hoisted.hoisted
|
||||
else response_api_optional_request_params
|
||||
)
|
||||
return super().transform_responses_api_request(
|
||||
|
|
@ -254,41 +250,14 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_codex_additional_tools_item(item: Any) -> bool:
|
||||
return isinstance(item, dict) and item.get("type") == _CODEX_ADDITIONAL_TOOLS_INPUT_ITEM_TYPE
|
||||
|
||||
@staticmethod
|
||||
def _tools_of_additional_tools_item(item: "dict[str, Any]") -> "list[Any]":
|
||||
tools: Final = item.get("tools")
|
||||
return tools if isinstance(tools, list) else []
|
||||
|
||||
@classmethod
|
||||
def _hoist_codex_additional_tools(
|
||||
cls,
|
||||
input: "str | ResponseInputParam",
|
||||
) -> "tuple[str | ResponseInputParam, list[Any]]":
|
||||
"""Codex's "responses lite" wire mode ships tool definitions inside
|
||||
`input` as {"type": "additional_tools", "role": "developer",
|
||||
"tools": [...]} items. api.openai.com accepts that item type; Mantle
|
||||
rejects the whole request with 400 "Invalid 'input': value did not
|
||||
match any expected variant" but accepts the same tools at the top
|
||||
level, so move them there and strip the items from `input`.
|
||||
"""
|
||||
if not isinstance(input, list):
|
||||
return input, []
|
||||
additional_tools_items: Final = [item for item in input if cls._is_codex_additional_tools_item(item)]
|
||||
if not additional_tools_items:
|
||||
return input, []
|
||||
remaining_input: Final = [item for item in input if not cls._is_codex_additional_tools_item(item)]
|
||||
hoisted_tools = [tool for item in additional_tools_items for tool in cls._tools_of_additional_tools_item(item)]
|
||||
verbose_logger.debug(
|
||||
"Bedrock Mantle Responses API: hoisting %d tool(s) out of %d 'additional_tools' input item(s) "
|
||||
"into the top-level tools param (Mantle rejects that input item type).",
|
||||
len(hoisted_tools),
|
||||
len(additional_tools_items),
|
||||
)
|
||||
return remaining_input, cls._filter_unsupported_tools(hoisted_tools)
|
||||
def _params_with_hoisted_tools(
|
||||
cls, params: Mapping[str, object], hoisted: HoistedAdditionalTools
|
||||
) -> dict[str, object]:
|
||||
supported_tools: Final = cls._filter_unsupported_tools(list(hoisted.tools))
|
||||
if supported_tools:
|
||||
return {**params, "tools": supported_tools}
|
||||
return {key: value for key, value in params.items() if key != "tools"}
|
||||
|
||||
@staticmethod
|
||||
def _agent_message_text(item: "Mapping[str, object]") -> str:
|
||||
|
|
|
|||
|
|
@ -115,6 +115,19 @@ def _parse_setup(session_configuration_request: str) -> BidiGenerateContentSetup
|
|||
return envelope.get("setup", empty_setup)
|
||||
|
||||
|
||||
def _grounding_metadata_from_frame(frame: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
|
||||
"""Read ``serverContent.groundingMetadata`` off the frame that carries the turn's usage.
|
||||
|
||||
Live reports grounding in the server frames rather than in ``usageMetadata``, and it emits both
|
||||
on the same frame, so the per-query charge is countable at the point usage is built.
|
||||
"""
|
||||
server_content: Final = frame.get("serverContent")
|
||||
if not isinstance(server_content, Mapping):
|
||||
return ()
|
||||
metadata: Final = server_content.get("groundingMetadata")
|
||||
return (metadata,) if isinstance(metadata, Mapping) else ()
|
||||
|
||||
|
||||
# Google bills Live transcription at an estimated 25 audio tokens/sec of input and
|
||||
# 175 text tokens/min of output (ai.google.dev/gemini-api/docs/pricing).
|
||||
GEMINI_LIVE_TRANSCRIBE_AUDIO_TOKENS_PER_SECOND: Final = 25
|
||||
|
|
@ -323,7 +336,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
)
|
||||
elif key == "input_audio_transcription" and value is not None:
|
||||
optional_params["inputAudioTranscription"] = {}
|
||||
elif key == "turn_detection":
|
||||
elif key == "turn_detection" and value is not None:
|
||||
value_typed = cast(OpenAIRealtimeTurnDetection, value)
|
||||
if (
|
||||
isinstance(value_typed, dict)
|
||||
|
|
@ -1049,6 +1062,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
|
|||
{**cast(dict, message), "usageMetadata": resolved_usage_metadata},
|
||||
),
|
||||
)
|
||||
grounding_metadata: Final = _grounding_metadata_from_frame(message)
|
||||
if grounding_metadata:
|
||||
VertexGeminiConfig._set_grounding_usage_counters( # pyright: ignore[reportPrivateUsage] # shared with the chat path; no public alias exists yet
|
||||
_chat_completion_usage, grounding_metadata
|
||||
)
|
||||
else:
|
||||
_chat_completion_usage = get_empty_usage()
|
||||
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
|||
stream_item_field,
|
||||
stream_item_fingerprint,
|
||||
stream_item_items,
|
||||
unappliable_request_rewrite,
|
||||
)
|
||||
from litellm.main import stream_chunk_builder
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
|
||||
|
|
@ -196,6 +197,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
else:
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
if guardrailed_texts and texts_to_check:
|
||||
if len(guardrailed_texts) != len(text_task_mappings):
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
await self._apply_guardrail_responses_to_input_texts(
|
||||
messages=messages,
|
||||
responses=guardrailed_texts,
|
||||
|
|
@ -210,6 +213,17 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
task_mappings=tool_call_task_mappings,
|
||||
)
|
||||
|
||||
elif (
|
||||
not images_to_check
|
||||
and not guardrail_to_apply.records_own_guardrail_information
|
||||
and (not_run_reason := self._not_run_reason(messages)) is not None
|
||||
):
|
||||
guardrail_to_apply.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response=not_run_reason,
|
||||
request_data=data,
|
||||
guardrail_status="not_run",
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"OpenAI Chat Completions: Processed input messages: %s",
|
||||
data.get("messages"),
|
||||
|
|
@ -217,6 +231,28 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def _not_run_reason(
|
||||
self,
|
||||
messages: Sequence[dict[str, Any]], # mutable-ok: raw request messages consumed by _extract_inputs
|
||||
) -> str | None:
|
||||
"""Why nothing was scanned, or None when the only unscoped content is images, which this handler never scans."""
|
||||
texts: Final[list[str]] = [] # mutable-ok: filled by _extract_inputs
|
||||
images: Final[list[str]] = [] # mutable-ok: filled by _extract_inputs
|
||||
tool_calls: Final[list[ChatCompletionToolParam]] = [] # mutable-ok: filled by _extract_inputs
|
||||
for msg_idx, message in enumerate(messages):
|
||||
self._extract_inputs(
|
||||
message=message,
|
||||
msg_idx=msg_idx,
|
||||
texts_to_check=texts,
|
||||
images_to_check=images,
|
||||
tool_calls_to_check=tool_calls,
|
||||
text_task_mappings=[], # mutable-ok: required by _extract_inputs, unused here
|
||||
tool_call_task_mappings=[], # mutable-ok: required by _extract_inputs, unused here
|
||||
)
|
||||
if texts or tool_calls:
|
||||
return "no scannable content after message scoping"
|
||||
return None if images else "no scannable content"
|
||||
|
||||
def extract_request_tool_names(self, data: dict) -> list[str]:
|
||||
"""Extract tool names from OpenAI chat completions request (tools[].function.name, functions[].name)."""
|
||||
names: Final[list[str]] = []
|
||||
|
|
|
|||
|
|
@ -56,6 +56,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
|
|||
stream_item_field,
|
||||
stream_item_fingerprint,
|
||||
stream_item_items,
|
||||
unappliable_request_rewrite,
|
||||
)
|
||||
from litellm.llms.openai.responses.guardrail_translation.tool_merge import merge_guardrailed_tools
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
|
|
@ -495,13 +496,13 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
data["instructions"] = written_back.instructions # rebind-ok: data is an out-param
|
||||
elif isinstance(input_data, str):
|
||||
guardrailed_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(guardrailed_texts) > 1:
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data # rebind-ok: data is an out-param
|
||||
else:
|
||||
rewritten_texts: Final = guardrailed_inputs.get("texts") or ()
|
||||
if len(rewritten_texts) != len(extracted.task_mappings):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite
|
||||
|
||||
raise UnappliableRequestRewrite(guardrail_to_apply.guardrail_name or "unknown")
|
||||
raise unappliable_request_rewrite(guardrail_to_apply.guardrail_name)
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
responses=rewritten_texts,
|
||||
|
|
|
|||
|
|
@ -6,8 +6,10 @@ from typing import Final, TypeAlias
|
|||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.responses.litellm_completion_transformation.custom_tools import custom_tool_grammar_suffix
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
NAMESPACE_DESCRIPTION_SEPARATOR,
|
||||
NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS,
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
|
|
@ -34,8 +36,8 @@ def _validated_tools(values: Iterable[object]) -> tuple[Tool, ...]:
|
|||
return tuple(tool for tool in validated if tool is not None)
|
||||
|
||||
|
||||
def _is_function(tool: Tool) -> bool:
|
||||
return tool.get("type") == "function"
|
||||
def _has_chat_tool(member: Tool) -> bool:
|
||||
return member.get("type") in NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS
|
||||
|
||||
|
||||
def _chat_tool_key(tool: Tool) -> str:
|
||||
|
|
@ -67,18 +69,19 @@ def _function_fields(tool: Tool) -> Tool:
|
|||
return function if function is not None else MappingProxyType({})
|
||||
|
||||
|
||||
def _without_namespace_prefix(key: str, value: object, prefix: str) -> object:
|
||||
if key != "description" or not isinstance(value, str) or not value.startswith(prefix):
|
||||
def _member_description(key: str, value: object, prefix: str, suffix: str) -> object:
|
||||
if key != "description" or not isinstance(value, str):
|
||||
return value
|
||||
return value[len(prefix) :]
|
||||
return value.replace(prefix, "", 1).replace(suffix, "", 1)
|
||||
|
||||
|
||||
def _rebuilt_member(member: Tool, flattened: Tool, guardrailed: Tool, namespace_description: str) -> Tool:
|
||||
flattened_function: Final = _function_fields(flattened)
|
||||
prefix: Final = f"{namespace_description}{NAMESPACE_DESCRIPTION_SEPARATOR}" if namespace_description else ""
|
||||
suffix: Final = custom_tool_grammar_suffix(member.get("format")) if member.get("type") == "custom" else ""
|
||||
changed_function: Final = MappingProxyType(
|
||||
{
|
||||
key: _without_namespace_prefix(key, value, prefix)
|
||||
key: _member_description(key, value, prefix, suffix)
|
||||
for key, value in _function_fields(guardrailed).items()
|
||||
if flattened_function.get(key) != value
|
||||
}
|
||||
|
|
@ -93,8 +96,8 @@ def _rebuilt_member(member: Tool, flattened: Tool, guardrailed: Tool, namespace_
|
|||
return {**member, **changed_extras, **changed_function} # mutable-ok: json.dumps rejects MappingProxyType
|
||||
|
||||
|
||||
def _rebuilt_function_members(
|
||||
function_members: Sequence[Tool],
|
||||
def _rebuilt_flattened_members(
|
||||
flattened_members: Sequence[Tool],
|
||||
flattened_group: Sequence[Tool],
|
||||
group_keys: Sequence[IndexedKey],
|
||||
guardrailed_by_key: Mapping[IndexedKey, Tool],
|
||||
|
|
@ -106,7 +109,7 @@ def _rebuilt_function_members(
|
|||
else member
|
||||
if guardrailed_by_key[key] == flattened
|
||||
else _rebuilt_member(member, flattened, guardrailed_by_key[key], namespace_description)
|
||||
for member, flattened, key in zip(function_members, flattened_group, group_keys)
|
||||
for member, flattened, key in zip(flattened_members, flattened_group, group_keys)
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -118,9 +121,9 @@ def _rebuilt_namespace(
|
|||
guardrailed_by_key: Mapping[IndexedKey, Tool],
|
||||
) -> tuple[Tool, ...]:
|
||||
namespace_description: Final = str(original.get("description") or "")
|
||||
rebuilt_functions: Final = iter(
|
||||
_rebuilt_function_members(
|
||||
tuple(member for member in members if _is_function(member)),
|
||||
rebuilt_flattened: Final = iter(
|
||||
_rebuilt_flattened_members(
|
||||
tuple(member for member in members if _has_chat_tool(member)),
|
||||
flattened_group,
|
||||
group_keys,
|
||||
guardrailed_by_key,
|
||||
|
|
@ -129,7 +132,7 @@ def _rebuilt_namespace(
|
|||
)
|
||||
rebuilt_members: Final = tuple(
|
||||
rebuilt
|
||||
for rebuilt in (next(rebuilt_functions) if _is_function(member) else member for member in members)
|
||||
for rebuilt in (next(rebuilt_flattened) if _has_chat_tool(member) else member for member in members)
|
||||
if rebuilt is not None
|
||||
)
|
||||
if not rebuilt_members:
|
||||
|
|
@ -149,7 +152,7 @@ def _merged_original(
|
|||
if guardrailed_group == tuple(flattened_group):
|
||||
return (original,)
|
||||
members: Final = _namespace_members(original) if original.get("type") == "namespace" else ()
|
||||
if members and sum(map(_is_function, members)) == len(flattened_group):
|
||||
if members and sum(map(_has_chat_tool, members)) == len(flattened_group):
|
||||
return _rebuilt_namespace(original, members, flattened_group, group_keys, guardrailed_by_key)
|
||||
if not guardrailed_group:
|
||||
return ()
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ These are the canonical credential types for the proxy. They live in the model
|
|||
layer; ``litellm.types.utils`` re-exports them for backwards compatibility.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
|
||||
from pydantic import BaseModel, model_validator
|
||||
|
||||
|
||||
|
|
@ -27,3 +29,10 @@ class CreateCredentialItem(CredentialBase):
|
|||
if not values.get("credential_values") and not values.get("model_id"):
|
||||
raise ValueError("Either credential_values or model_id must be set")
|
||||
return values
|
||||
|
||||
|
||||
class UpdateCredentialItem(BaseModel):
|
||||
credential_name: str
|
||||
credential_info: Mapping[str, object]
|
||||
credential_values: Mapping[str, object] | None = None
|
||||
model_id: str | None = None
|
||||
|
|
|
|||
|
|
@ -4,11 +4,12 @@ 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
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, NamedTuple
|
||||
|
||||
import httpx
|
||||
from pydantic import (
|
||||
BaseModel,
|
||||
BeforeValidator,
|
||||
ConfigDict,
|
||||
Field,
|
||||
Json,
|
||||
|
|
@ -47,6 +48,7 @@ from litellm.types.proxy.carried_budget_state import (
|
|||
)
|
||||
from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry
|
||||
from litellm.types.router import RouterErrors, UpdateRouterConfig
|
||||
from litellm.types.router_weights import validate_router_settings_dict
|
||||
from litellm.types.secret_managers.main import KeyManagementSystem
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
|
|
@ -656,6 +658,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
[
|
||||
# user
|
||||
"/user/new",
|
||||
"/management/v1/users/bulk",
|
||||
"/user/update",
|
||||
"/user/bulk_update",
|
||||
"/user/delete",
|
||||
|
|
@ -1981,8 +1984,14 @@ class OrgMember(MemberBase):
|
|||
|
||||
from litellm.models.team import TeamBase as TeamBase # noqa: E402
|
||||
|
||||
RouterSettingsDict = Annotated[
|
||||
dict[str, object],
|
||||
BeforeValidator(validate_router_settings_dict, json_schema_input_type=UpdateRouterConfig),
|
||||
]
|
||||
|
||||
|
||||
class NewTeamRequest(TeamBase):
|
||||
router_settings: RouterSettingsDict | None = None
|
||||
model_aliases: dict | None = None
|
||||
tags: list | None = None
|
||||
guardrails: list[str] | None = None
|
||||
|
|
@ -2080,7 +2089,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
|
|||
allowed_vector_store_indexes: list[AllowedVectorStoreIndexItem] | None = None
|
||||
enforced_batch_output_expires_after: dict | None = None
|
||||
enforced_file_expires_after: dict | None = None
|
||||
router_settings: dict | None = None
|
||||
router_settings: RouterSettingsDict | None = None
|
||||
access_group_ids: list[str] | None = None
|
||||
budget_limits: list[BudgetLimitEntry] | None = None # multiple concurrent budget windows
|
||||
default_team_member_models: list[str] | None = None # default allowed_models seeded onto new team members
|
||||
|
|
|
|||
|
|
@ -249,6 +249,7 @@ async def authenticate_user(
|
|||
|
||||
if os.getenv("DATABASE_URL") is not None:
|
||||
response = await generate_key_helper_fn(
|
||||
llm_router=None,
|
||||
request_type="key",
|
||||
**{
|
||||
"user_role": LitellmUserRoles.PROXY_ADMIN,
|
||||
|
|
@ -324,6 +325,7 @@ async def authenticate_user(
|
|||
await _rehash_password_if_needed(_user_row.user_id, password, _password)
|
||||
if os.getenv("DATABASE_URL") is not None:
|
||||
response = await generate_key_helper_fn(
|
||||
llm_router=None,
|
||||
request_type="key",
|
||||
**{
|
||||
"user_role": user_role,
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ _PROXY_ADMIN_VIEW_ONLY_BLOCKED_ROUTES: Final = frozenset(
|
|||
[
|
||||
# user
|
||||
"/user/new",
|
||||
"/management/v1/users/bulk",
|
||||
"/user/delete",
|
||||
"/user/bulk_update",
|
||||
# team
|
||||
|
|
@ -758,6 +759,7 @@ class RouteChecks:
|
|||
_ADMIN_VIEWER_BLOCKED_WRITE_ROUTES = frozenset(
|
||||
[
|
||||
"/user/new",
|
||||
"/management/v1/users/bulk",
|
||||
"/user/delete",
|
||||
"/user/bulk_update",
|
||||
"/team/new",
|
||||
|
|
|
|||
|
|
@ -878,6 +878,7 @@ async def _auto_register_jwt_mapping(
|
|||
# the NOT NULL @id constraint. Every successful key-creation caller (e.g.
|
||||
# /key/generate) passes table_name="key" explicitly.
|
||||
key_data: Final = await generate_key_helper_fn(
|
||||
llm_router=None,
|
||||
request_type="key",
|
||||
table_name="key",
|
||||
team_id=team_id,
|
||||
|
|
|
|||
|
|
@ -580,8 +580,8 @@ What the command changed is recorded in `~/.litellm/claude_configure_state.json`
|
|||
`lite configure claude`, `lite login --config-claude`, `lite up` and `lite autoroute up` also install a status line (`~/.litellm/statusline.py`, registered as `statusLine` in `~/.claude/settings.json` unless you already run one) that shows which model the auto-router actually served the last turn and, once the proxy has recorded the session, what the session cost against the router's savings baseline:
|
||||
|
||||
```
|
||||
claude-auto · Routed to: claude-haiku-4-5 -63% vs Claude Opus 5
|
||||
LiteLLM ████████░░░░░░░░░░░░░░░░ $0.14
|
||||
Routed to: claude-haiku-4-5 -63% vs Claude Opus 5
|
||||
claude-auto ████████░░░░░░░░░░░░░░░░ $0.14
|
||||
Claude Opus 5 ████████████████████████ $0.38
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ import os
|
|||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
import unicodedata
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from collections.abc import Callable, Mapping
|
||||
|
|
@ -42,7 +43,6 @@ FETCH_TIMEOUT_SECONDS: Final = 3
|
|||
BAR_WIDTH: Final = 24
|
||||
BAR_FULL: Final = "\u2588"
|
||||
BAR_EMPTY: Final = "\u2591"
|
||||
SEPARATOR: Final = " \u00b7 "
|
||||
TRANSCRIPT_SCAN_LIMIT_BYTES: Final = 4 * 1024 * 1024
|
||||
CLAUDE_BASE_URL_ENV_KEYS: Final = ("ANTHROPIC_BASE_URL",)
|
||||
CLAUDE_API_KEY_ENV_KEYS: Final = ("ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_KEY")
|
||||
|
|
@ -50,7 +50,6 @@ CODEX_BASE_URL_ENV_KEYS: Final = ("OPENAI_BASE_URL",)
|
|||
CODEX_API_KEY_ENV_KEYS: Final = ("OPENAI_API_KEY",)
|
||||
CODEX_STOP_EVENT: Final = "Stop"
|
||||
SYNTHETIC_MODEL: Final = "<synthetic>"
|
||||
LITELLM_LABEL: Final = "LiteLLM"
|
||||
RESET: Final = "\033[0m"
|
||||
BOLD: Final = "\033[1m"
|
||||
DIM: Final = "\033[90m"
|
||||
|
|
@ -302,31 +301,37 @@ def _bar(fraction: float, color: str, width: int, use_color: bool) -> str:
|
|||
return f"{color}{BAR_FULL * filled}{DIM}{BAR_EMPTY * (width - filled)}{RESET}"
|
||||
|
||||
|
||||
def _display_width(label: str) -> int:
|
||||
return sum(
|
||||
2 if unicodedata.east_asian_width(character) in ("W", "F") else 1
|
||||
for character in label
|
||||
if unicodedata.category(character) not in ("Mn", "Me")
|
||||
)
|
||||
|
||||
|
||||
def render(model: str, session: Session | None, config_dir: Path, use_color: bool, bar_width: int = BAR_WIDTH) -> str:
|
||||
def paint(code: str, text: str) -> str:
|
||||
return f"{code}{text}{RESET}" if use_color else text
|
||||
|
||||
routed: Final = paint(BOLD, f"Routed to: {model}")
|
||||
if session is None:
|
||||
if session is None or session.baseline_model is None or session.baseline_spend <= 0:
|
||||
return routed
|
||||
header: Final = f"{session.router_name}{SEPARATOR}{routed}"
|
||||
if session.baseline_model is None or session.baseline_spend <= 0:
|
||||
return header
|
||||
reference: Final = baseline_label(session.baseline_model, config_dir)
|
||||
pct: Final = (session.baseline_spend - session.spend) / session.baseline_spend * 100
|
||||
delta: Final = paint(LITELLM_COLOR, f"{'-' if pct >= 0 else '+'}{abs(round(pct))}% vs {reference}")
|
||||
peak: Final = max(session.spend, session.baseline_spend)
|
||||
label_width: Final = max(len(LITELLM_LABEL), len(reference))
|
||||
label_width: Final = max(_display_width(session.router_name), _display_width(reference))
|
||||
rows: Final = (
|
||||
(LITELLM_LABEL, session.spend, LITELLM_COLOR),
|
||||
(session.router_name, session.spend, LITELLM_COLOR),
|
||||
(reference, session.baseline_spend, BASELINE_COLOR),
|
||||
)
|
||||
lines: Final = (
|
||||
f"{paint(DIM, label.ljust(label_width))} {_bar(amount / peak, color, bar_width, use_color)} "
|
||||
f"{paint(DIM, label + ' ' * (label_width - _display_width(label)))} "
|
||||
f"{_bar(amount / peak, color, bar_width, use_color)} "
|
||||
f"{paint(DIM, f'${amount:.2f}')}"
|
||||
for label, amount, color in rows
|
||||
)
|
||||
return "\n".join((f"{header} {delta}", *lines))
|
||||
return "\n".join((f"{routed} {delta}", *lines))
|
||||
|
||||
|
||||
def color_enabled(env: Mapping[str, str]) -> bool:
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ import httpx
|
|||
import orjson
|
||||
from fastapi import HTTPException, Request, status
|
||||
from fastapi.responses import JSONResponse, Response, StreamingResponse
|
||||
from pydantic import ValidationError
|
||||
from starlette.types import Receive, Scope, Send
|
||||
|
||||
import litellm
|
||||
|
|
@ -76,6 +77,7 @@ from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_di
|
|||
from litellm.router_utils.common_utils import resolve_model_group_alias
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.router import RouterRateLimitError
|
||||
from litellm.types.router_weights import validate_router_weights
|
||||
|
||||
_LateResponseT = TypeVar("_LateResponseT", bound=Response)
|
||||
_LlmCallT = TypeVar("_LlmCallT")
|
||||
|
|
@ -1939,6 +1941,13 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# This avoids expensive Router instantiation on each request
|
||||
if router_settings is not None:
|
||||
self.data["router_settings_override"] = router_settings
|
||||
try:
|
||||
self.data["_router_weights"] = validate_router_weights(router_settings.get("weights"))
|
||||
except ValidationError:
|
||||
self.data["_router_weights"] = None
|
||||
verbose_proxy_logger.warning(
|
||||
"Ignoring invalid saved router weights; update team/key router_settings"
|
||||
)
|
||||
alias_target: Final = await _resolve_per_request_model_group_alias(
|
||||
requested_model=self.data.get("model"),
|
||||
router_settings=router_settings,
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ from typing import Final
|
|||
|
||||
from fastapi import status
|
||||
|
||||
from litellm.constants import STRINGIFIED_NONE
|
||||
|
||||
_OPENAI_ERROR_TYPE_BY_STATUS: Final[Mapping[int, str]] = MappingProxyType(
|
||||
{
|
||||
status.HTTP_401_UNAUTHORIZED: "authentication_error",
|
||||
|
|
@ -35,7 +37,7 @@ def openai_error_type(exc: object, status_code: int) -> str:
|
|||
"""OpenAI types ``error.type`` as a required string, so an exception carrying none
|
||||
falls back to the type its status code stands for."""
|
||||
carried: Final = attribute_of(exc, "type")
|
||||
if isinstance(carried, str):
|
||||
if isinstance(carried, str) and carried != STRINGIFIED_NONE:
|
||||
return carried
|
||||
mapped: Final = _OPENAI_ERROR_TYPE_BY_STATUS.get(status_code)
|
||||
if mapped is not None:
|
||||
|
|
@ -49,4 +51,4 @@ def openai_error_param(exc: object) -> str | None:
|
|||
"""OpenAI types ``error.param`` as nullable, so an exception carrying none
|
||||
serializes as JSON ``null``."""
|
||||
carried: Final = attribute_of(exc, "param")
|
||||
return carried if isinstance(carried, str) else None
|
||||
return carried if isinstance(carried, str) and carried != STRINGIFIED_NONE else None
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ class ComplianceChecker:
|
|||
|
||||
def __init__(self, data: ComplianceCheckRequest):
|
||||
self.data = data
|
||||
self.guardrails = data.guardrail_information or []
|
||||
self.guardrails = tuple(g for g in data.guardrail_information or () if g.get("guardrail_status") != "not_run")
|
||||
|
||||
def _get_guardrails_by_mode(self, mode: str) -> list[dict]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -2,25 +2,31 @@
|
|||
CRUD endpoints for storing reusable credentials.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import (
|
||||
Annotated,
|
||||
Final,
|
||||
cast, # noqa: TID251 # jsonify_object in proxy/utils.py is annotated with a bare dict
|
||||
)
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.litellm_core_utils.litellm_logging import _get_masked_values
|
||||
from litellm.models.credentials import UpdateCredentialItem
|
||||
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.proxy.utils import handle_exception_on_proxy, jsonify_object
|
||||
from litellm.repositories.base_repository import is_unique_violation
|
||||
from litellm.repositories.credentials_repository import CredentialsRepository
|
||||
from litellm.types.utils import CreateCredentialItem, CredentialItem
|
||||
|
||||
router: Final = APIRouter()
|
||||
_CREDENTIAL_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
class CredentialHelperUtils:
|
||||
|
|
@ -40,6 +46,33 @@ class CredentialHelperUtils:
|
|||
)
|
||||
|
||||
|
||||
def _credential_exists_detail(credential_name: str) -> str:
|
||||
return (
|
||||
f"Credential '{credential_name}' already exists. "
|
||||
f"Update it with PATCH /credentials/{credential_name}, or delete it first."
|
||||
)
|
||||
|
||||
|
||||
def get_llm_router() -> litellm.Router | None:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
return llm_router
|
||||
|
||||
|
||||
def _resolve_deployment_credentials(llm_router: litellm.Router | None, model_id: str) -> Mapping[str, object]:
|
||||
if llm_router is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail="LLM router not found. Please ensure you have a valid router instance.",
|
||||
)
|
||||
if llm_router.get_deployment(model_id) is None:
|
||||
raise HTTPException(status_code=404, detail="Model not found")
|
||||
credential_values: Final = llm_router.get_deployment_credentials(model_id)
|
||||
if credential_values is None:
|
||||
raise HTTPException(status_code=404, detail="Model not found")
|
||||
return _CREDENTIAL_DICT_ADAPTER.validate_python(credential_values)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/credentials",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
|
|
@ -50,13 +83,14 @@ async def create_credential(
|
|||
fastapi_response: Response,
|
||||
credential: CreateCredentialItem,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
llm_router: Annotated[litellm.Router | None, Depends(get_llm_router)] = None,
|
||||
):
|
||||
"""
|
||||
[BETA] endpoint. This might change unexpectedly.
|
||||
Stores credential in DB.
|
||||
Reloads credentials in memory.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import llm_router, prisma_client
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
try:
|
||||
if prisma_client is None:
|
||||
|
|
@ -64,29 +98,19 @@ async def create_credential(
|
|||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
if credential.model_id:
|
||||
if llm_router is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail="LLM router not found. Please ensure you have a valid router instance.",
|
||||
)
|
||||
# get model from router
|
||||
model: Final = llm_router.get_deployment(credential.model_id)
|
||||
if model is None:
|
||||
raise HTTPException(status_code=404, detail="Model not found")
|
||||
credential_values: Final = llm_router.get_deployment_credentials(credential.model_id)
|
||||
if credential_values is None:
|
||||
raise HTTPException(status_code=404, detail="Model not found")
|
||||
credential.credential_values = credential_values
|
||||
|
||||
if credential.credential_values is None:
|
||||
credential_values: Final = (
|
||||
_resolve_deployment_credentials(llm_router, credential.model_id)
|
||||
if credential.model_id
|
||||
else credential.credential_values
|
||||
)
|
||||
if credential_values is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Credential values are required. Unable to infer credential values from model ID.",
|
||||
)
|
||||
processed_credential: Final = CredentialItem(
|
||||
credential_name=credential.credential_name,
|
||||
credential_values=credential.credential_values,
|
||||
credential_values=_CREDENTIAL_DICT_ADAPTER.validate_python(credential_values),
|
||||
credential_info=credential.credential_info,
|
||||
)
|
||||
encrypted_credential: Final = CredentialHelperUtils.encrypt_credential_values(processed_credential)
|
||||
|
|
@ -94,13 +118,18 @@ async def create_credential(
|
|||
credentials_dict_jsonified: Final = cast( # cast-ok: deep-copies a model_dump, so keys are str
|
||||
"dict[str, object]", jsonify_object(credentials_dict)
|
||||
)
|
||||
await CredentialsRepository(prisma_client).create(
|
||||
data={
|
||||
**credentials_dict_jsonified,
|
||||
"created_by": user_api_key_dict.user_id,
|
||||
"updated_by": user_api_key_dict.user_id,
|
||||
}
|
||||
)
|
||||
try:
|
||||
await CredentialsRepository(prisma_client).create(
|
||||
data={
|
||||
**credentials_dict_jsonified,
|
||||
"created_by": user_api_key_dict.user_id,
|
||||
"updated_by": user_api_key_dict.user_id,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
if not is_unique_violation(e):
|
||||
raise
|
||||
raise HTTPException(status_code=409, detail=_credential_exists_detail(credential.credential_name))
|
||||
|
||||
## ADD TO LITELLM ##
|
||||
CredentialAccessor.upsert_credentials([processed_credential])
|
||||
|
|
@ -300,9 +329,10 @@ def update_db_credential(
|
|||
async def update_credential(
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
credential: CredentialItem,
|
||||
credential: UpdateCredentialItem,
|
||||
credential_name: str = Path(..., description="The credential name, percent-decoded; may contain slashes"),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
llm_router: Annotated[litellm.Router | None, Depends(get_llm_router)] = None,
|
||||
):
|
||||
"""
|
||||
[BETA] endpoint. This might change unexpectedly.
|
||||
|
|
@ -319,7 +349,16 @@ async def update_credential(
|
|||
db_credential: Final = await credentials_repository.find_by_name(credential_name)
|
||||
if db_credential is None:
|
||||
raise HTTPException(status_code=404, detail="Credential not found in DB.")
|
||||
merged_credential: Final = update_db_credential(db_credential, credential)
|
||||
patch: Final = CredentialItem(
|
||||
credential_name=credential.credential_name,
|
||||
credential_info=_CREDENTIAL_DICT_ADAPTER.validate_python(credential.credential_info),
|
||||
credential_values=_CREDENTIAL_DICT_ADAPTER.validate_python(
|
||||
_resolve_deployment_credentials(llm_router, credential.model_id)
|
||||
if credential.model_id
|
||||
else credential.credential_values or {}
|
||||
),
|
||||
)
|
||||
merged_credential: Final = update_db_credential(db_credential, patch)
|
||||
credential_object_jsonified: Final = cast( # cast-ok: deep-copies a model_dump, so keys are str
|
||||
"dict[str, object]", jsonify_object(merged_credential.model_dump())
|
||||
)
|
||||
|
|
@ -341,11 +380,11 @@ async def update_credential(
|
|||
|
||||
if existing_in_memory is not None:
|
||||
in_memory_values: Final = dict(existing_in_memory.credential_values or {})
|
||||
if credential.credential_values:
|
||||
in_memory_values.update(credential.credential_values)
|
||||
if patch.credential_values:
|
||||
in_memory_values.update(patch.credential_values)
|
||||
in_memory_info: Final = dict(existing_in_memory.credential_info or {})
|
||||
if credential.credential_info:
|
||||
in_memory_info.update(credential.credential_info)
|
||||
if patch.credential_info:
|
||||
in_memory_info.update(patch.credential_info)
|
||||
updated_in_memory: Final = CredentialItem(
|
||||
credential_name=new_name,
|
||||
credential_values=in_memory_values,
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@
|
|||
|
||||
import fnmatch
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
|
||||
import httpx
|
||||
|
|
@ -24,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import ChatCompletionToolParam
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GenericGuardrailAPIMetadata,
|
||||
GenericGuardrailAPIRequest,
|
||||
|
|
@ -150,6 +150,26 @@ def _extract_inbound_headers(
|
|||
return None
|
||||
|
||||
|
||||
def _structured_rows_to_write_back(
|
||||
original_rows: Sequence[AllMessageValues] | None,
|
||||
shown_rows: Sequence[AllMessageValues] | None,
|
||||
returned_rows: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...] | None:
|
||||
"""The request model drops row keys its message types do not declare, so a
|
||||
row the server echoes back verbatim is restored to the original row object.
|
||||
A server that echoes every row back unchanged has not rewritten anything
|
||||
per row, so its answer is read from texts, as it was before rows could be
|
||||
returned at all."""
|
||||
if original_rows is None or shown_rows is None or len(returned_rows) != len(original_rows):
|
||||
return tuple(returned_rows)
|
||||
if all(returned == shown for shown, returned in zip(shown_rows, returned_rows)):
|
||||
return None
|
||||
return tuple(
|
||||
original if returned == shown else returned
|
||||
for original, shown, returned in zip(original_rows, shown_rows, returned_rows)
|
||||
)
|
||||
|
||||
|
||||
class GenericGuardrailAPI(CustomGuardrail):
|
||||
"""
|
||||
Generic Guardrail API integration for LiteLLM.
|
||||
|
|
@ -322,6 +342,8 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
texts: list,
|
||||
images: list[str] | None,
|
||||
tools: list[ChatCompletionToolParam] | None,
|
||||
structured_messages: Sequence[AllMessageValues] | None,
|
||||
shown_messages: Sequence[AllMessageValues] | None,
|
||||
guardrail_response: GenericGuardrailAPIResponse,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
# Action is NONE or no modifications needed
|
||||
|
|
@ -336,6 +358,13 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
return_inputs["tools"] = guardrail_response.tools
|
||||
elif tools:
|
||||
return_inputs["tools"] = tools
|
||||
rows_to_write_back: Final = (
|
||||
_structured_rows_to_write_back(structured_messages, shown_messages, guardrail_response.structured_messages)
|
||||
if guardrail_response.structured_messages
|
||||
else None
|
||||
)
|
||||
if rows_to_write_back is not None:
|
||||
return_inputs["structured_messages"] = list(rows_to_write_back) # mutable-ok: guardrail inputs take a list
|
||||
if guardrail_response.stream_holdback_chars is not None:
|
||||
return_inputs["stream_holdback_chars"] = guardrail_response.stream_holdback_chars
|
||||
return return_inputs
|
||||
|
|
@ -473,6 +502,8 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
texts=texts,
|
||||
images=images,
|
||||
tools=tools,
|
||||
structured_messages=structured_messages,
|
||||
shown_messages=guardrail_request.structured_messages,
|
||||
guardrail_response=guardrail_response,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1600,8 +1600,8 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
Args:
|
||||
texts: Flattened text entries from the framework.
|
||||
messages: Original request messages (request_data["messages"]),
|
||||
NOT structured_messages (which may have injected system content).
|
||||
messages: The structured messages the framework flattened into ``texts``,
|
||||
hoisted top-level system prompt included, so positions line up.
|
||||
|
||||
Returns a set of scannable indices, or None on count mismatch or no user/developer
|
||||
message (safety fallback to existing role-filter behavior).
|
||||
|
|
@ -1788,15 +1788,10 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
structured_messages: Final = inputs.get("structured_messages")
|
||||
if structured_messages:
|
||||
# For Anthropic /v1/messages: default to latest-user-only scanning.
|
||||
# Uses request_data["messages"] (original format), NOT structured_messages
|
||||
# (which has injected system content from adapter translation).
|
||||
if self._use_latest_user_only(request_data, logging_obj):
|
||||
original_messages: Final = request_data.get("messages")
|
||||
if original_messages:
|
||||
scannable_indices = self._get_latest_user_text_indices(texts, original_messages)
|
||||
scannable_indices = self._get_latest_user_text_indices(texts, structured_messages)
|
||||
# Fall through to existing role filtering if:
|
||||
# - not Anthropic, OR flag explicitly False, OR
|
||||
# - no original messages, OR
|
||||
# - latest-user extraction returned None (no user / count mismatch)
|
||||
if scannable_indices is None:
|
||||
scannable_indices = self._get_scannable_text_indices(texts, structured_messages)
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import asyncio
|
|||
import base64
|
||||
import os
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional
|
||||
|
||||
import httpx
|
||||
|
|
@ -14,11 +15,13 @@ from litellm.integrations.custom_guardrail import (
|
|||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts, message_with_slot_texts
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -28,12 +31,36 @@ if TYPE_CHECKING:
|
|||
|
||||
_SANITIZE_FILE_FAIL_OPEN_TIMEOUT_SECONDS: Final = 30.0
|
||||
_SANITIZE_FILE_QUEUED_STATUSES: Final = frozenset({"created", "in progress"})
|
||||
_PROTECT_ROLES: Final = frozenset({"system", "user", "assistant"})
|
||||
|
||||
|
||||
class PromptSecurityGuardrailMissingSecrets(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _inputs_with_structured_messages(
|
||||
inputs: GenericGuardrailAPIInputs, rewritten_messages: Sequence[AllMessageValues] | None
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
if rewritten_messages is None:
|
||||
return inputs
|
||||
patched: Final[GenericGuardrailAPIInputs] = {
|
||||
**inputs,
|
||||
"structured_messages": list(rewritten_messages), # mutable-ok: the TypedDict field is declared as a list
|
||||
}
|
||||
return patched
|
||||
|
||||
|
||||
def _inputs_with_modifications(
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
modified_texts: list[str],
|
||||
rewritten_messages: Sequence[AllMessageValues] | None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
if not modified_texts:
|
||||
return _inputs_with_structured_messages(inputs, rewritten_messages)
|
||||
with_texts: Final[GenericGuardrailAPIInputs] = {**inputs, "texts": modified_texts}
|
||||
return _inputs_with_structured_messages(with_texts, rewritten_messages)
|
||||
|
||||
|
||||
class _ProtectVerdict(TypedDict, total=False):
|
||||
"""One side (``prompt`` or ``response``) of an ``/api/protect`` verdict."""
|
||||
|
||||
|
|
@ -276,14 +303,39 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
detail="Blocked by Prompt Security, Violations: " + ", ".join(violations),
|
||||
)
|
||||
elif action == "modify":
|
||||
# Extract modified texts from modified_messages
|
||||
modified_messages: Final = result.get("modified_messages", [])
|
||||
modified_texts: Final = self._extract_texts_from_messages(modified_messages)
|
||||
if modified_texts:
|
||||
inputs["texts"] = modified_texts
|
||||
return _inputs_with_modifications(
|
||||
inputs,
|
||||
self._extract_texts_from_messages(modified_messages),
|
||||
self._structured_messages_with_modifications(structured_messages, modified_messages),
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
def _is_sent_to_protect(self, message: Mapping[str, object]) -> bool:
|
||||
return self.check_tool_results or message.get("role") in _PROTECT_ROLES
|
||||
|
||||
def _structured_messages_with_modifications(
|
||||
self,
|
||||
structured_messages: Sequence[AllMessageValues],
|
||||
modified_messages: Sequence[Mapping[str, object]],
|
||||
) -> tuple[AllMessageValues, ...] | None:
|
||||
sent_indices: Final = tuple(
|
||||
index for index, message in enumerate(structured_messages) if self._is_sent_to_protect(message)
|
||||
)
|
||||
if not sent_indices or len(sent_indices) != len(modified_messages):
|
||||
return None
|
||||
rewritten: Final = tuple(
|
||||
message_with_slot_texts(structured_messages[index], self._extract_texts_from_messages((modified,)))
|
||||
for index, modified in zip(sent_indices, modified_messages)
|
||||
)
|
||||
replacements: Final = MappingProxyType(
|
||||
{index: message for index, message in zip(sent_indices, rewritten) if message is not None}
|
||||
)
|
||||
if len(replacements) != len(sent_indices):
|
||||
return None
|
||||
return tuple(replacements.get(index, message) for index, message in enumerate(structured_messages))
|
||||
|
||||
async def _apply_guardrail_on_response(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
|
|
@ -347,19 +399,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
return inputs
|
||||
|
||||
def _extract_texts_from_messages(self, messages: Sequence[Mapping[str, object]]) -> list[str]:
|
||||
"""Extract text content from messages."""
|
||||
texts: Final = []
|
||||
for message in messages:
|
||||
content = message.get("content")
|
||||
if isinstance(content, str):
|
||||
texts.append(content)
|
||||
elif isinstance(content, list):
|
||||
for item in content:
|
||||
if isinstance(item, dict) and item.get("type") == "text":
|
||||
text = item.get("text")
|
||||
if text:
|
||||
texts.append(text)
|
||||
return texts
|
||||
return [text for message in messages for text in message_slot_texts(message)]
|
||||
|
||||
async def _process_standalone_images(self, images: list[str], user_api_key_alias: str | None) -> None:
|
||||
"""Process standalone images from inputs (data URLs)."""
|
||||
|
|
@ -681,14 +721,13 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
|
||||
This allows checking tool results for indirect prompt injection when enabled.
|
||||
"""
|
||||
supported_roles: Final = ["system", "user", "assistant"]
|
||||
filtered_messages: Final = []
|
||||
transformed_count = 0
|
||||
filtered_count = 0
|
||||
|
||||
for message in messages:
|
||||
role = message.get("role", "")
|
||||
if role in supported_roles:
|
||||
if role in _PROTECT_ROLES:
|
||||
filtered_messages.append(message)
|
||||
else:
|
||||
if self.check_tool_results:
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@ if TYPE_CHECKING:
|
|||
router: Final = APIRouter()
|
||||
|
||||
_EMPTY_UNITS: Final[Mapping[str, int]] = MappingProxyType({})
|
||||
_ACTION_SEVERITY: Final[Mapping[str, int]] = MappingProxyType({"passed": 0, "flagged": 1, "blocked": 2})
|
||||
_ACTION_SEVERITY: Final[Mapping[str, int]] = MappingProxyType({"not_run": 0, "passed": 1, "flagged": 2, "blocked": 3})
|
||||
|
||||
_T = TypeVar("_T")
|
||||
|
||||
|
|
@ -325,7 +325,7 @@ class UsageDetailResponse(BaseModel):
|
|||
class UsageLogEntry(BaseModel):
|
||||
id: str
|
||||
timestamp: str
|
||||
action: str # blocked | passed | flagged
|
||||
action: str # blocked | passed | flagged | not_run
|
||||
score: float | None
|
||||
latency_ms: float | None
|
||||
model: str | None
|
||||
|
|
|
|||
|
|
@ -193,10 +193,12 @@ async def _upsert_rows_with_retry(
|
|||
|
||||
|
||||
def guardrail_status_to_action(status: str | None) -> str:
|
||||
"""Map StandardLogging guardrail_status to blocked/passed/flagged."""
|
||||
"""Map StandardLogging guardrail_status to blocked/passed/flagged/not_run."""
|
||||
if not status:
|
||||
return "passed"
|
||||
s: Final = (status or "").lower()
|
||||
if s == "not_run":
|
||||
return "not_run"
|
||||
if "intervened" in s or "block" in s:
|
||||
return "blocked"
|
||||
if "flagged" in s or "fail" in s or "error" in s:
|
||||
|
|
@ -354,37 +356,49 @@ async def process_spend_logs_guardrail_usage(
|
|||
"flagged_count": 0,
|
||||
}
|
||||
)
|
||||
index_rows: Final[list[dict[str, object]]] = []
|
||||
index_rows_by_key: Final[dict[tuple[str, str], dict[str, object]]] = {}
|
||||
|
||||
for payload in logs_to_process:
|
||||
request_id = payload.get("request_id")
|
||||
start_time = _parse_payload_start_time(payload)
|
||||
if not request_id or start_time is None:
|
||||
if not isinstance(request_id, str) or not request_id or start_time is None:
|
||||
continue
|
||||
date_key = _date_str(start_time)
|
||||
|
||||
for entry in _parse_guardrail_info_from_payload(payload):
|
||||
guardrail_id = entry.get("guardrail_id") or entry.get("guardrail_name") or ""
|
||||
if not guardrail_id:
|
||||
entries = _parse_guardrail_info_from_payload(payload)
|
||||
ids_by_name = MappingProxyType(
|
||||
{
|
||||
e["guardrail_name"]: e["guardrail_id"]
|
||||
for e in entries
|
||||
if e.get("guardrail_id") and isinstance(e.get("guardrail_name"), str) and e["guardrail_name"]
|
||||
}
|
||||
)
|
||||
for entry in entries:
|
||||
raw_name = entry.get("guardrail_name")
|
||||
guardrail_name = raw_name if isinstance(raw_name, str) else ""
|
||||
guardrail_id = entry.get("guardrail_id") or ids_by_name.get(guardrail_name) or guardrail_name
|
||||
if not isinstance(guardrail_id, str) or not guardrail_id:
|
||||
continue
|
||||
key = _MetricsKey(guardrail_id, date_key)
|
||||
daily_guardrail[key]["requests_evaluated"] += 1
|
||||
action = guardrail_status_to_action(entry.get("guardrail_status"))
|
||||
if action == "passed":
|
||||
daily_guardrail[key]["passed_count"] += 1
|
||||
elif action == "blocked":
|
||||
daily_guardrail[key]["blocked_count"] += 1
|
||||
else:
|
||||
daily_guardrail[key]["flagged_count"] += 1
|
||||
if action != "not_run":
|
||||
key = _MetricsKey(guardrail_id, date_key)
|
||||
daily_guardrail[key]["requests_evaluated"] += 1
|
||||
if action == "passed":
|
||||
daily_guardrail[key]["passed_count"] += 1
|
||||
elif action == "blocked":
|
||||
daily_guardrail[key]["blocked_count"] += 1
|
||||
else:
|
||||
daily_guardrail[key]["flagged_count"] += 1
|
||||
policy_id = entry.get("policy_id")
|
||||
index_rows.append(
|
||||
{
|
||||
prior = index_rows_by_key.get((request_id, guardrail_id))
|
||||
if prior is None or (prior["policy_id"] is None and policy_id is not None):
|
||||
index_rows_by_key[(request_id, guardrail_id)] = {
|
||||
"request_id": request_id,
|
||||
"guardrail_id": guardrail_id,
|
||||
"policy_id": policy_id,
|
||||
"start_time": start_time,
|
||||
}
|
||||
)
|
||||
index_rows: Final = tuple(index_rows_by_key.values())
|
||||
|
||||
async with pending.lock:
|
||||
pending_metrics: Final = pending.metrics
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Contract machinery shared by every LiteLLM-defined list route, on any surface."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
from typing import Final
|
||||
from urllib.parse import urlencode
|
||||
|
||||
|
|
@ -7,6 +8,7 @@ from fastapi import Request
|
|||
from fastapi.dependencies.utils import get_flat_params
|
||||
from fastapi.params import ParamTypes
|
||||
from fastapi.responses import JSONResponse
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import (
|
||||
ListLinks,
|
||||
|
|
@ -56,6 +58,31 @@ def escape_like(value: str) -> str:
|
|||
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||||
|
||||
|
||||
class ValidationErrorDetail(TypedDict):
|
||||
"""The two keys of a pydantic/FastAPI validation error a problem document needs."""
|
||||
|
||||
loc: ReadOnly[tuple[int | str, ...]]
|
||||
msg: ReadOnly[str]
|
||||
|
||||
|
||||
def request_validation_problem(errors: Sequence[ValidationErrorDetail]) -> ProblemDetail:
|
||||
"""A body that fails validation (an unknown field included) is 422; a bad query parameter is 400."""
|
||||
detail: Final = "; ".join(f"{'.'.join(str(part) for part in error['loc'][1:])}: {error['msg']}" for error in errors)
|
||||
if any(error["loc"] and error["loc"][0] == "body" for error in errors):
|
||||
return ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}invalid-request-body",
|
||||
title="Invalid request body",
|
||||
status=422,
|
||||
detail=detail or "The request body is invalid.",
|
||||
)
|
||||
return ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}invalid-query-parameter",
|
||||
title="Invalid query parameter",
|
||||
status=400,
|
||||
detail=detail or "The request query parameters are invalid.",
|
||||
)
|
||||
|
||||
|
||||
def unknown_query_param_problem(unknown: tuple[str, ...], allowed: tuple[str, ...]) -> ProblemDetail:
|
||||
return ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}unknown-query-parameter",
|
||||
|
|
|
|||
|
|
@ -221,6 +221,8 @@ LITELLM_TRACE_CONTROL_METADATA_FIELDS: Final = frozenset(
|
|||
)
|
||||
|
||||
_UNTRUSTED_ROOT_CONTROL_FIELDS: Final = (
|
||||
"weights",
|
||||
"_router_weights",
|
||||
"proxy_server_request",
|
||||
"standard_logging_object",
|
||||
"secret_fields",
|
||||
|
|
@ -334,7 +336,7 @@ _CLIENT_PRICING_METADATA_FIELDS: Final = frozenset({"model_info", "standard_logg
|
|||
# and read by spend logs as fact; a client value has no legitimate meaning and no
|
||||
# key or team setting keeps it, so the strip is never gated.
|
||||
_ROUTER_RESERVED_METADATA_FIELDS: Final = frozenset(
|
||||
{"attempted_fallbacks", "original_model_group", CLIENT_OUTPUT_CEILING_METADATA_KEY}
|
||||
{"attempted_fallbacks", "original_model_group", "request_retry_count", CLIENT_OUTPUT_CEILING_METADATA_KEY}
|
||||
)
|
||||
_ALLOW_CLIENT_PRICING_OVERRIDE_METADATA_KEY: Final = "allow_client_pricing_override"
|
||||
|
||||
|
|
|
|||
|
|
@ -567,7 +567,7 @@ async def new_user(
|
|||
teams = check_if_default_team_set()
|
||||
organization_ids: Final = cast(list[str] | None, data_json.pop("organizations", None))
|
||||
|
||||
response: Final = await generate_key_helper_fn(request_type="user", **data_json)
|
||||
response: Final = await generate_key_helper_fn(request_type="user", **data_json, llm_router=None)
|
||||
# Admin UI Logic
|
||||
# Add User to Team and Organization
|
||||
# if team_id passed add this user to the team
|
||||
|
|
|
|||
|
|
@ -96,6 +96,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
_add_model_to_db,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights
|
||||
from litellm.proxy.management_helpers.access_group_key_sync import (
|
||||
sync_key_access_group_membership,
|
||||
sync_key_regeneration_access_group_membership,
|
||||
|
|
@ -148,6 +149,7 @@ from litellm.types.proxy.management_endpoints.key_management_endpoints import (
|
|||
BulkUpdateKeyRequest,
|
||||
BulkUpdateKeyResponse,
|
||||
BulkUpdateTeamKeysRequest,
|
||||
CustomKeyPolicyRequest,
|
||||
FailedKeyUpdate,
|
||||
KeySearchWhere,
|
||||
SuccessfulKeyUpdate,
|
||||
|
|
@ -201,6 +203,10 @@ class _KeyUpdateResult(TypedDict):
|
|||
data: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class _StoredKeyRouterSettings(BaseModel):
|
||||
router_settings: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
class _KeyRowWhere(TypedDict):
|
||||
token: ReadOnly[str]
|
||||
|
||||
|
|
@ -280,6 +286,7 @@ def _config_table(prisma_client: PrismaClient) -> _ConfigTableActions:
|
|||
class _CustomKeyHooksModule(Protocol):
|
||||
user_custom_key_generate: Callable[..., Awaitable[Mapping[str, object]]] | None
|
||||
user_custom_key_update: Callable[..., Awaitable[Mapping[str, object]]] | None
|
||||
user_custom_key_policy: Callable[..., Awaitable[Mapping[str, object]]] | None
|
||||
|
||||
|
||||
def _custom_key_generate_hook(
|
||||
|
|
@ -294,6 +301,161 @@ def _custom_key_update_hook(
|
|||
return hooks.user_custom_key_update
|
||||
|
||||
|
||||
def _custom_key_policy_hook(
|
||||
hooks: _CustomKeyHooksModule,
|
||||
) -> Callable[..., Awaitable[Mapping[str, object]]] | None:
|
||||
return hooks.user_custom_key_policy
|
||||
|
||||
|
||||
async def _enforce_custom_key_update_policy(
|
||||
hook: Callable[..., Awaitable[Mapping[str, object]]] | None,
|
||||
data: UpdateKeyRequest,
|
||||
) -> None:
|
||||
if hook is None:
|
||||
return
|
||||
if not inspect.iscoroutinefunction(hook):
|
||||
raise ValueError("user_custom_key_update must be a coroutine")
|
||||
result: Final = await hook(data)
|
||||
if not result.get("decision", True):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=result.get("message", "Authentication Failed - Custom Auth Rule"),
|
||||
)
|
||||
|
||||
|
||||
async def _enforce_custom_key_policy(
|
||||
hook: Callable[..., Awaitable[Mapping[str, object]]] | None,
|
||||
build_policy_request: Callable[[], CustomKeyPolicyRequest],
|
||||
) -> None:
|
||||
if hook is None:
|
||||
return
|
||||
if not inspect.iscoroutinefunction(hook):
|
||||
raise ValueError("user_custom_key_policy must be a coroutine")
|
||||
result: Final = await hook(build_policy_request())
|
||||
if not result.get("decision", True):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=result.get("message", "Authentication Failed - Custom Auth Rule"),
|
||||
)
|
||||
|
||||
|
||||
_KEY_UPDATE_JSON_STRING_COLUMNS: Final = frozenset({"router_settings", "budget_limits"})
|
||||
|
||||
_KEY_METADATA_REQUEST_FIELDS: Final = frozenset(
|
||||
(*LiteLLM_ManagementEndpoint_MetadataFields_Premium, *LiteLLM_ManagementEndpoint_MetadataFields)
|
||||
)
|
||||
|
||||
|
||||
def _decode_json_string_column(column: str, value: object) -> object:
|
||||
if column in _KEY_UPDATE_JSON_STRING_COLUMNS and isinstance(value, str):
|
||||
return json.loads(value)
|
||||
return value
|
||||
|
||||
|
||||
def _verification_token_from_row(row: Mapping[str, object]) -> LiteLLM_VerificationToken:
|
||||
org_id: Final = row["organization_id"] if "organization_id" in row else row.get("org_id")
|
||||
return LiteLLM_VerificationToken.model_validate(MappingProxyType({**row, "org_id": org_id}))
|
||||
|
||||
|
||||
def _effective_key_after_update(
|
||||
existing_key_row: LiteLLM_VerificationToken,
|
||||
non_default_values: Mapping[str, object],
|
||||
) -> LiteLLM_VerificationToken:
|
||||
overlay: Final = MappingProxyType(
|
||||
{column: _decode_json_string_column(column, value) for column, value in non_default_values.items()}
|
||||
)
|
||||
return _verification_token_from_row(
|
||||
MappingProxyType({**existing_key_row.model_dump(), **overlay, "object_permission": None})
|
||||
)
|
||||
|
||||
|
||||
def _update_policy_request(
|
||||
operation: Literal["update", "regenerate"],
|
||||
existing_key_row: LiteLLM_VerificationToken,
|
||||
non_default_values: Mapping[str, object],
|
||||
request: UpdateKeyRequest | RegenerateKeyRequest,
|
||||
) -> CustomKeyPolicyRequest:
|
||||
return CustomKeyPolicyRequest(
|
||||
operation=operation,
|
||||
existing_key=_verification_token_from_row(existing_key_row.model_dump()),
|
||||
effective_key=_effective_key_after_update(
|
||||
existing_key_row=existing_key_row, non_default_values=non_default_values
|
||||
),
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
def _generate_budget_windows(
|
||||
budget_limits: Sequence[BudgetLimitEntry] | None,
|
||||
) -> tuple[Mapping[str, object], ...] | None:
|
||||
if not budget_limits:
|
||||
return None
|
||||
return tuple(
|
||||
MappingProxyType(
|
||||
{
|
||||
**window.model_dump(),
|
||||
"reset_at": get_budget_reset_time(budget_duration=window.budget_duration).isoformat(),
|
||||
}
|
||||
)
|
||||
for window in budget_limits
|
||||
)
|
||||
|
||||
|
||||
def _effective_key_for_generate(data: GenerateKeyRequest, now: datetime) -> LiteLLM_VerificationToken:
|
||||
requested: Final = data.model_dump(exclude_unset=True, exclude_none=True)
|
||||
metadata_fields: Final = MappingProxyType(
|
||||
{field: value for field, value in requested.items() if field in _KEY_METADATA_REQUEST_FIELDS}
|
||||
)
|
||||
column_fields: Final = MappingProxyType(
|
||||
{field: value for field, value in requested.items() if field not in _KEY_METADATA_REQUEST_FIELDS}
|
||||
)
|
||||
metadata: Final = data.metadata or MappingProxyType({})
|
||||
folded_metadata: Final = {**metadata, **metadata_fields} # mutable-ok: encrypt_callback_vars needs a dict
|
||||
columns: Final = handle_key_type(data, {**column_fields}) # mutable-ok: handle_key_type mutates in place
|
||||
expires: Final = (
|
||||
now + timedelta(seconds=duration_in_seconds(duration=data.duration)) if data.duration is not None else None
|
||||
)
|
||||
budget_reset_at: Final = (
|
||||
get_budget_reset_time(budget_duration=data.budget_duration) if data.budget_duration is not None else None
|
||||
)
|
||||
key_rotation_at: Final = (
|
||||
now + timedelta(seconds=duration_in_seconds(duration=data.rotation_interval))
|
||||
if data.auto_rotate and data.rotation_interval
|
||||
else None
|
||||
)
|
||||
return _verification_token_from_row(
|
||||
MappingProxyType(
|
||||
{
|
||||
**columns,
|
||||
"metadata": encrypt_callback_vars(folded_metadata),
|
||||
"expires": expires,
|
||||
"budget_reset_at": budget_reset_at,
|
||||
"key_rotation_at": key_rotation_at,
|
||||
"budget_limits": _generate_budget_windows(data.budget_limits),
|
||||
"object_permission": None,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
_EMPTY_DURATION_MEANS_UNCHANGED: Final = frozenset({"duration", "budget_duration"})
|
||||
|
||||
|
||||
def _regenerate_request_as_update_request(key: str, data: RegenerateKeyRequest) -> UpdateKeyRequest | None:
|
||||
changed_fields: Final = MappingProxyType(
|
||||
{
|
||||
field: value
|
||||
for field, value in data.model_dump(exclude_unset=True).items()
|
||||
if field in UpdateKeyRequest.model_fields
|
||||
and field != "key"
|
||||
and not (field in _EMPTY_DURATION_MEANS_UNCHANGED and value == "")
|
||||
}
|
||||
)
|
||||
if not changed_fields:
|
||||
return None
|
||||
return UpdateKeyRequest(key=key, **changed_fields)
|
||||
|
||||
|
||||
class _LegacyDumpable(Protocol):
|
||||
def dict(self) -> Mapping[str, object]: ...
|
||||
|
||||
|
|
@ -987,6 +1149,7 @@ async def _common_key_generation_helper(
|
|||
litellm_changed_by: str | None,
|
||||
team_table: LiteLLM_TeamTableCachedObj | None,
|
||||
) -> GenerateKeyResponse:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.proxy_server import (
|
||||
litellm_proxy_admin_name,
|
||||
llm_router,
|
||||
|
|
@ -1135,6 +1298,16 @@ async def _common_key_generation_helper(
|
|||
"litellm.proxy.proxy_server.generate_key_fn(): Enterprise key management params not applied - %s", e
|
||||
)
|
||||
|
||||
await _enforce_custom_key_policy(
|
||||
hook=_custom_key_policy_hook(proxy_server),
|
||||
build_policy_request=lambda: CustomKeyPolicyRequest(
|
||||
operation="generate",
|
||||
existing_key=None,
|
||||
effective_key=_effective_key_for_generate(data=data, now=datetime.now(timezone.utc)),
|
||||
request=data,
|
||||
),
|
||||
)
|
||||
|
||||
# TODO: @ishaan-jaff: Migrate all budget tracking to use LiteLLM_BudgetTable
|
||||
_budget_id = data.budget_id
|
||||
if prisma_client is not None and data.soft_budget is not None:
|
||||
|
|
@ -1330,7 +1503,7 @@ async def _common_key_generation_helper(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key")
|
||||
response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router)
|
||||
|
||||
response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response
|
||||
|
||||
|
|
@ -2234,7 +2407,26 @@ async def _update_key_row_with_soft_budget(
|
|||
async def prepare_key_update_data(
|
||||
data: UpdateKeyRequest | RegenerateKeyRequest,
|
||||
existing_key_row: LiteLLM_VerificationToken,
|
||||
*,
|
||||
prisma_client: PrismaClient | None = None,
|
||||
llm_router: Router | None = None,
|
||||
):
|
||||
if data.router_settings is not None or (
|
||||
"router_settings" not in data.model_fields_set
|
||||
and "team_id" in data.model_fields_set
|
||||
and data.team_id != existing_key_row.team_id
|
||||
):
|
||||
effective_settings: Final = (
|
||||
data.router_settings
|
||||
if data.router_settings is not None
|
||||
else _StoredKeyRouterSettings.model_validate(existing_key_row, from_attributes=True).router_settings
|
||||
)
|
||||
await validate_router_settings_weights(
|
||||
effective_settings,
|
||||
team_id=data.team_id if "team_id" in data.model_fields_set else existing_key_row.team_id,
|
||||
prisma_client=prisma_client,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
data_json: Final[dict] = data.model_dump(exclude_unset=True)
|
||||
data_json.pop("key", None)
|
||||
data_json.pop("new_key", None)
|
||||
|
|
@ -2301,12 +2493,6 @@ async def prepare_key_update_data(
|
|||
# sentinel for Json? columns, so store the JSON literal null
|
||||
non_default_values["budget_limits"] = json.dumps(None)
|
||||
|
||||
if "object_permission" in non_default_values:
|
||||
non_default_values = await _handle_update_object_permission(
|
||||
data_json=non_default_values,
|
||||
existing_key_row=existing_key_row,
|
||||
)
|
||||
|
||||
_metadata: Final = existing_key_row.metadata or {}
|
||||
|
||||
# validate model_max_budget
|
||||
|
|
@ -2327,13 +2513,12 @@ async def prepare_key_update_data(
|
|||
async def _handle_update_object_permission(
|
||||
data_json: dict,
|
||||
existing_key_row: LiteLLM_VerificationToken,
|
||||
prisma_client: PrismaClient,
|
||||
) -> dict:
|
||||
"""
|
||||
Handle the update of object permission.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
"""Persist the requested object permission row and swap it for its id, only after the key policy allowed the write."""
|
||||
if "object_permission" not in data_json:
|
||||
return data_json
|
||||
|
||||
# Use the common helper to handle the object permission update
|
||||
object_permission_id: Final = await handle_update_object_permission_common(
|
||||
data_json=data_json,
|
||||
existing_object_permission_id=existing_key_row.object_permission_id,
|
||||
|
|
@ -2467,6 +2652,7 @@ async def _process_single_key_update(
|
|||
llm_router: Router | None,
|
||||
user_custom_key_update: Callable | None = None,
|
||||
existing_key_row: LiteLLM_VerificationToken | None = None,
|
||||
user_custom_key_policy: Callable[..., Awaitable[Mapping[str, object]]] | None = None,
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Process a single key update with all validations and checks.
|
||||
|
|
@ -2575,7 +2761,19 @@ async def _process_single_key_update(
|
|||
)
|
||||
|
||||
# Prepare update data
|
||||
non_default_values = await prepare_key_update_data(data=update_key_request, existing_key_row=existing_key_row)
|
||||
non_default_values = await prepare_key_update_data(
|
||||
data=update_key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router
|
||||
)
|
||||
|
||||
await _enforce_custom_key_policy(
|
||||
hook=user_custom_key_policy,
|
||||
build_policy_request=lambda: _update_policy_request(
|
||||
operation="update",
|
||||
existing_key_row=existing_key_row,
|
||||
non_default_values=non_default_values,
|
||||
request=update_key_request,
|
||||
),
|
||||
)
|
||||
|
||||
# Update key in database
|
||||
if prisma_client is None:
|
||||
|
|
@ -2584,7 +2782,12 @@ async def _process_single_key_update(
|
|||
detail={"error": "Database not connected"},
|
||||
)
|
||||
|
||||
_data: Final = {**non_default_values, "token": update_key_request.key}
|
||||
update_values: Final = await _handle_update_object_permission(
|
||||
data_json=non_default_values,
|
||||
existing_key_row=existing_key_row,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
_data: Final = {**update_values, "token": update_key_request.key}
|
||||
response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict
|
||||
"Mapping[str, object] | None",
|
||||
await prisma_client.update_data(token=update_key_request.key, data=_data),
|
||||
|
|
@ -3077,23 +3280,13 @@ async def update_key_fn(
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
# Custom key update hook
|
||||
custom_key_update_hook: Final[Callable[..., Awaitable[Mapping[str, object]]] | None] = _custom_key_update_hook(
|
||||
proxy_server
|
||||
)
|
||||
if custom_key_update_hook is not None:
|
||||
if inspect.iscoroutinefunction(custom_key_update_hook):
|
||||
result: Final = await custom_key_update_hook(data)
|
||||
else:
|
||||
raise ValueError("user_custom_key_update must be a coroutine")
|
||||
decision: Final = result.get("decision", True)
|
||||
message: Final = result.get("message", "Authentication Failed - Custom Auth Rule")
|
||||
if not decision:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=message)
|
||||
await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=data)
|
||||
|
||||
# Enforce upperbound key params on update (don't fill defaults)
|
||||
_enforce_upperbound_key_params(data, fill_defaults=False)
|
||||
non_default_values: Final = await prepare_key_update_data(data=data, existing_key_row=existing_key_row)
|
||||
non_default_values: Final = await prepare_key_update_data(
|
||||
data=data, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router
|
||||
)
|
||||
|
||||
# Only validate key_alias format if it's actually being changed
|
||||
new_key_alias: Final = non_default_values.get("key_alias", None)
|
||||
|
|
@ -3114,21 +3307,36 @@ async def update_key_fn(
|
|||
existing_key_alias=existing_key_row.key_alias,
|
||||
)
|
||||
|
||||
await _enforce_custom_key_policy(
|
||||
hook=_custom_key_policy_hook(proxy_server),
|
||||
build_policy_request=lambda: _update_policy_request(
|
||||
operation="update",
|
||||
existing_key_row=existing_key_row,
|
||||
non_default_values=non_default_values,
|
||||
request=data,
|
||||
),
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise Exception("Not connected to DB!")
|
||||
|
||||
update_values: Final = await _handle_update_object_permission(
|
||||
data_json=non_default_values,
|
||||
existing_key_row=existing_key_row,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name
|
||||
response: Final = (
|
||||
await _update_key_row_with_soft_budget(
|
||||
prisma_client=prisma_client,
|
||||
key=key,
|
||||
data=data,
|
||||
non_default_values=non_default_values,
|
||||
non_default_values=update_values,
|
||||
existing_key_row=existing_key_row,
|
||||
changed_by=changed_by,
|
||||
)
|
||||
if "soft_budget" in data.model_fields_set
|
||||
else await prisma_client.update_data(token=key, data=MappingProxyType({**non_default_values, "token": key}))
|
||||
else await prisma_client.update_data(token=key, data=MappingProxyType({**update_values, "token": key}))
|
||||
)
|
||||
|
||||
# Delete - key from cache, since it's been updated!
|
||||
|
|
@ -3263,6 +3471,7 @@ async def bulk_update_keys(
|
|||
)
|
||||
|
||||
custom_key_update_hook: Final = _custom_key_update_hook(proxy_server)
|
||||
custom_key_policy_hook: Final = _custom_key_policy_hook(proxy_server)
|
||||
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
|
||||
raise HTTPException(
|
||||
|
|
@ -3310,6 +3519,7 @@ async def bulk_update_keys(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
user_custom_key_update=custom_key_update_hook,
|
||||
user_custom_key_policy=custom_key_policy_hook,
|
||||
)
|
||||
|
||||
successful_updates.append(
|
||||
|
|
@ -3427,6 +3637,7 @@ async def bulk_update_team_keys(
|
|||
)
|
||||
|
||||
custom_key_update_hook: Final = _custom_key_update_hook(proxy_server)
|
||||
custom_key_policy_hook: Final = _custom_key_policy_hook(proxy_server)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
|
|
@ -3557,6 +3768,7 @@ async def bulk_update_team_keys(
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
user_custom_key_update=custom_key_update_hook,
|
||||
user_custom_key_policy=custom_key_policy_hook,
|
||||
existing_key_row=existing_by_token[db_token],
|
||||
)
|
||||
|
||||
|
|
@ -4082,6 +4294,40 @@ def _check_model_access_group(models: list[str] | None, llm_router: Router | Non
|
|||
return True
|
||||
|
||||
|
||||
_NO_METADATA: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
def metadata_json_with_limits(
|
||||
metadata: Mapping[str, object] | None,
|
||||
*,
|
||||
model_rpm_limit: Mapping[str, object] | None,
|
||||
model_tpm_limit: Mapping[str, object] | None,
|
||||
mcp_rpm_limit: Mapping[str, int] | None,
|
||||
tag_rpm_limit: Mapping[str, int] | None,
|
||||
guardrails: Sequence[str] | None,
|
||||
policies: Sequence[str] | None,
|
||||
prompts: Sequence[str] | None,
|
||||
) -> str:
|
||||
"""Serialize the stored metadata blob with the per-model, MCP, tag, guardrail, policy and prompt settings folded in."""
|
||||
limits: Final = tuple(
|
||||
(name, value)
|
||||
for name, value in (
|
||||
("model_rpm_limit", model_rpm_limit),
|
||||
("model_tpm_limit", model_tpm_limit),
|
||||
("mcp_rpm_limit", mcp_rpm_limit),
|
||||
("tag_rpm_limit", tag_rpm_limit),
|
||||
("guardrails", guardrails),
|
||||
("policies", policies),
|
||||
("prompts", prompts),
|
||||
)
|
||||
if value is not None
|
||||
)
|
||||
if metadata is None and not limits:
|
||||
return json.dumps(None)
|
||||
merged: Final = {**(metadata or _NO_METADATA), **dict(limits)} # mutable-ok: encrypt_callback_vars takes a dict
|
||||
return json.dumps(encrypt_callback_vars(merged))
|
||||
|
||||
|
||||
async def generate_key_helper_fn(
|
||||
request_type: Literal["user", "key"], # identifies if this request is from /user/new or /key/generate
|
||||
duration: str | None = None,
|
||||
|
|
@ -4137,15 +4383,24 @@ async def generate_key_helper_fn(
|
|||
object_permission: LiteLLM_ObjectPermissionBase | None = None,
|
||||
auto_rotate: bool | None = None,
|
||||
rotation_interval: str | None = None,
|
||||
router_settings: dict | None = None,
|
||||
router_settings: dict[str, object] | None = None,
|
||||
access_group_ids: list[str] | None = None,
|
||||
budget_limits: list | None = None, # multiple concurrent budget windows
|
||||
*,
|
||||
llm_router: Router | None = None,
|
||||
):
|
||||
from litellm.proxy.proxy_server import premium_user, prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise Exception("Connect Proxy to database to generate keys - https://docs.litellm.ai/docs/proxy/virtual_keys ")
|
||||
|
||||
await validate_router_settings_weights(
|
||||
router_settings,
|
||||
team_id=team_id,
|
||||
prisma_client=prisma_client,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
if token is None:
|
||||
if key is not None:
|
||||
token = key
|
||||
|
|
@ -4184,31 +4439,16 @@ async def generate_key_helper_fn(
|
|||
permissions_json: Final = json.dumps(permissions)
|
||||
router_settings_json: Final = safe_dumps(router_settings) if router_settings is not None else safe_dumps({})
|
||||
|
||||
# Add model_rpm_limit and model_tpm_limit to metadata
|
||||
if model_rpm_limit is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["model_rpm_limit"] = model_rpm_limit
|
||||
if model_tpm_limit is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["model_tpm_limit"] = model_tpm_limit
|
||||
if mcp_rpm_limit is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["mcp_rpm_limit"] = mcp_rpm_limit
|
||||
if tag_rpm_limit is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["tag_rpm_limit"] = tag_rpm_limit
|
||||
if guardrails is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["guardrails"] = guardrails
|
||||
if policies is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["policies"] = policies
|
||||
if prompts is not None:
|
||||
metadata = metadata or {}
|
||||
metadata["prompts"] = prompts
|
||||
|
||||
metadata = encrypt_callback_vars(metadata)
|
||||
metadata_json: Final = json.dumps(metadata)
|
||||
metadata_json: Final = metadata_json_with_limits(
|
||||
metadata,
|
||||
model_rpm_limit=model_rpm_limit,
|
||||
model_tpm_limit=model_tpm_limit,
|
||||
mcp_rpm_limit=mcp_rpm_limit,
|
||||
tag_rpm_limit=tag_rpm_limit,
|
||||
guardrails=guardrails,
|
||||
policies=policies,
|
||||
prompts=prompts,
|
||||
)
|
||||
validate_model_max_budget(model_max_budget)
|
||||
model_max_budget_json: Final = json.dumps(model_max_budget)
|
||||
budget_fallbacks_json: Final = json.dumps(budget_fallbacks or {})
|
||||
|
|
@ -5070,6 +5310,7 @@ async def _insert_deprecated_key(
|
|||
async def _execute_virtual_key_regeneration(
|
||||
*,
|
||||
prisma_client: PrismaClient,
|
||||
llm_router: Router | None = None,
|
||||
key_in_db: LiteLLM_VerificationToken,
|
||||
hashed_api_key: str,
|
||||
key: str,
|
||||
|
|
@ -5080,6 +5321,7 @@ async def _execute_virtual_key_regeneration(
|
|||
proxy_logging_obj: ProxyLogging,
|
||||
) -> GenerateKeyResponse:
|
||||
"""Generate new token, update DB, invalidate cache, and return response."""
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.proxy_server import hash_token
|
||||
|
||||
# Mirror the /key/update ownership rebind guard. See helper docstring.
|
||||
|
|
@ -5127,15 +5369,34 @@ async def _execute_virtual_key_regeneration(
|
|||
|
||||
non_default_values = {}
|
||||
if data is not None:
|
||||
update_request: Final = _regenerate_request_as_update_request(key=hashed_api_key, data=data)
|
||||
if update_request is not None:
|
||||
await _enforce_custom_key_update_policy(hook=_custom_key_update_hook(proxy_server), data=update_request)
|
||||
# Enforce upperbound key params on regenerate (don't fill defaults)
|
||||
_enforce_upperbound_key_params(data, fill_defaults=False)
|
||||
non_default_values = await prepare_key_update_data(data=data, existing_key_row=key_in_db)
|
||||
non_default_values = await prepare_key_update_data(
|
||||
data=data, existing_key_row=key_in_db, prisma_client=prisma_client, llm_router=llm_router
|
||||
)
|
||||
# Only validate key_alias format if it's actually being changed
|
||||
new_key_alias: Final = non_default_values.get("key_alias")
|
||||
if new_key_alias != key_in_db.key_alias:
|
||||
_validate_key_alias_format(key_alias=new_key_alias)
|
||||
verbose_proxy_logger.debug("non_default_values: %s", non_default_values)
|
||||
update_data.update(non_default_values)
|
||||
await _enforce_custom_key_policy(
|
||||
hook=_custom_key_policy_hook(proxy_server),
|
||||
build_policy_request=lambda: _update_policy_request(
|
||||
operation="regenerate",
|
||||
existing_key_row=key_in_db,
|
||||
non_default_values=non_default_values,
|
||||
request=data if data is not None else RegenerateKeyRequest(),
|
||||
),
|
||||
)
|
||||
update_values: Final = await _handle_update_object_permission(
|
||||
data_json=non_default_values,
|
||||
existing_key_row=key_in_db,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
update_data.update(update_values)
|
||||
jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data)
|
||||
|
||||
# Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash,
|
||||
|
|
@ -5145,6 +5406,13 @@ async def _execute_virtual_key_regeneration(
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
await _persist_deleted_verification_tokens(
|
||||
keys=[key_in_db],
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
)
|
||||
|
||||
# If grace period set, insert deprecated key so old key remains valid
|
||||
await _insert_deprecated_key(
|
||||
prisma_client=prisma_client,
|
||||
|
|
@ -5268,6 +5536,7 @@ async def regenerate_key_fn(
|
|||
try:
|
||||
from litellm.proxy.proxy_server import (
|
||||
hash_token,
|
||||
llm_router,
|
||||
master_key,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
|
|
@ -5443,19 +5712,9 @@ async def regenerate_key_fn(
|
|||
if litellm_changed_by is not None and not isinstance(litellm_changed_by, str):
|
||||
litellm_changed_by = None
|
||||
|
||||
# Save the old key record to deleted table before regeneration.
|
||||
# This preserves key_alias and team_id metadata for historical spend records.
|
||||
# If this fails, abort the regeneration to avoid permanently losing the
|
||||
# old hash→metadata mapping.
|
||||
await _persist_deleted_verification_tokens(
|
||||
keys=[_key_in_db],
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=litellm_changed_by,
|
||||
)
|
||||
|
||||
return await _execute_virtual_key_regeneration(
|
||||
prisma_client=prisma_client,
|
||||
llm_router=llm_router,
|
||||
key_in_db=_key_in_db,
|
||||
hashed_api_key=hashed_api_key,
|
||||
key=key,
|
||||
|
|
|
|||
|
|
@ -10,9 +10,13 @@ from litellm.proxy.management_endpoints.management_v1.budgets import (
|
|||
from litellm.proxy.management_endpoints.management_v1.spend_logs import (
|
||||
router as spend_logs_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.management_v1.users import (
|
||||
router as users_router,
|
||||
)
|
||||
|
||||
router: Final = APIRouter()
|
||||
router.include_router(budgets_router)
|
||||
router.include_router(spend_logs_router)
|
||||
router.include_router(users_router)
|
||||
|
||||
__all__ = ["router"]
|
||||
|
|
|
|||
105
litellm/proxy/management_endpoints/management_v1/users.py
Normal file
105
litellm/proxy/management_endpoints/management_v1/users.py
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
"""`POST /management/v1/users/bulk`."""
|
||||
|
||||
from typing import Annotated, Final
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.list_api.common import PROBLEM_TYPE_BASE, ManagementProblem
|
||||
from litellm.proxy.management_endpoints.management_v1.common import MANAGEMENT_V1_PREFIX
|
||||
from litellm.proxy.management_helpers.bulk_user_creation import bulk_create_users
|
||||
from litellm.proxy.management_helpers.utils import (
|
||||
management_endpoint_wrapper, # pyright: ignore[reportUnknownVariableType] # legacy untyped decorator
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
|
||||
BulkNewUserRequest,
|
||||
BulkNewUserResponse,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail
|
||||
|
||||
router: Final = APIRouter(prefix=MANAGEMENT_V1_PREFIX)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/users/bulk",
|
||||
tags=["Internal User management"], # mutable-ok: fastapi types tags as list[str | Enum]
|
||||
dependencies=(Depends(user_api_key_auth),),
|
||||
response_model=BulkNewUserResponse,
|
||||
)
|
||||
@management_endpoint_wrapper
|
||||
async def bulk_create_users_route(
|
||||
data: BulkNewUserRequest,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
) -> BulkNewUserResponse:
|
||||
"""
|
||||
Create up to 500 internal users in one request, optionally adding each one to teams.
|
||||
|
||||
Every entry in `users` takes the same fields as `/user/new`, with two differences: `auto_create_key`
|
||||
defaults to `false` (opt in per user to also get a virtual key back) and `send_invite_email` is not
|
||||
supported. Unknown fields are rejected with 422. Rows are validated together (duplicate ids or emails,
|
||||
unknown teams, roles the caller may not grant), inserted in one statement, and each referenced team is
|
||||
written once for all of its new members.
|
||||
|
||||
Rows fail independently: a bad row is reported in `data` with `success: false` and an `error`, and the
|
||||
other rows still get created. A user that was created but could not be added to one of its teams is
|
||||
reported with `success: true`, `teams` listing where they did land, and `error` naming the failed team.
|
||||
The whole request is refused with a 403 problem document only if creating the valid rows would exceed
|
||||
the license seat limit.
|
||||
|
||||
Example curl:
|
||||
```
|
||||
curl -X POST "http://localhost:4000/management/v1/users/bulk" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-H "Authorization: Bearer sk-1234" \\
|
||||
-d '{
|
||||
"users": [
|
||||
{"user_email": "a@example.com", "user_role": "internal_user", "teams": ["team-1"]},
|
||||
{"user_email": "b@example.com", "user_role": "internal_user", "auto_create_key": true}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
Returns `data` (one entry per input row, in order, with `user_id`, `user_email`, `success`, `teams`,
|
||||
`key`, `error`) and `meta` with `total_requested`, `created` and `failed`.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import (
|
||||
_license_check, # pyright: ignore[reportPrivateUsage] # same proxy license singleton /user/new reads
|
||||
litellm_proxy_admin_name,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise ManagementProblem(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}database-not-connected",
|
||||
title="Database not connected",
|
||||
status=503,
|
||||
detail=CommonProxyErrors.db_not_connected_error.value,
|
||||
)
|
||||
)
|
||||
|
||||
return await bulk_create_users(
|
||||
users=data.users,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
prisma_client=prisma_client,
|
||||
license_check=_license_check,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
except ManagementProblem:
|
||||
raise
|
||||
except Exception: # noqa: BLE001 # a driver error answers as a problem document, not the OpenAI error shape
|
||||
verbose_proxy_logger.exception("/management/v1/users/bulk: Exception occurred")
|
||||
raise ManagementProblem(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}internal-server-error",
|
||||
title="Internal server error",
|
||||
status=500,
|
||||
detail="Failed to create users.",
|
||||
)
|
||||
)
|
||||
129
litellm/proxy/management_endpoints/router_weights.py
Normal file
129
litellm/proxy/management_endpoints/router_weights.py
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
from abc import abstractmethod
|
||||
from collections.abc import Mapping
|
||||
from typing import Annotated, Final, Protocol
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import BaseModel, BeforeValidator, ValidationError
|
||||
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.types.router_weights import RouterWeights
|
||||
|
||||
|
||||
class _StoredModel(Protocol):
|
||||
@property
|
||||
@abstractmethod
|
||||
def model_id(self) -> str:
|
||||
pass
|
||||
|
||||
|
||||
class _ModelDb(Protocol):
|
||||
@property
|
||||
@abstractmethod
|
||||
def litellm_proxymodeltable(self) -> TableActions[_StoredModel]:
|
||||
pass
|
||||
|
||||
|
||||
class _PrismaClient(Protocol):
|
||||
@property
|
||||
@abstractmethod
|
||||
def db(self) -> _ModelDb:
|
||||
pass
|
||||
|
||||
|
||||
class _Router(Protocol):
|
||||
@abstractmethod
|
||||
def get_deployment(self, model_id: str) -> object | None:
|
||||
pass
|
||||
|
||||
|
||||
class _RouterWeightSettings(BaseModel):
|
||||
weights: RouterWeights | None = None
|
||||
|
||||
|
||||
class _RouterWeightModelInfo(BaseModel):
|
||||
team_id: str | None = None
|
||||
db_model: bool | None = None
|
||||
team_public_model_name: str | None = None
|
||||
|
||||
|
||||
def _router_weight_model_info(value: object) -> _RouterWeightModelInfo:
|
||||
if isinstance(value, str):
|
||||
return _RouterWeightModelInfo.model_validate_json(value)
|
||||
return _RouterWeightModelInfo.model_validate(value or {}, from_attributes=True)
|
||||
|
||||
|
||||
class _RouterWeightDeployment(BaseModel):
|
||||
model_name: str
|
||||
model_info: Annotated[_RouterWeightModelInfo, BeforeValidator(_router_weight_model_info)]
|
||||
|
||||
|
||||
def _validate_router_weight_reference(
|
||||
model_group: str,
|
||||
deployment_id: str,
|
||||
team_id: str | None,
|
||||
stored: _RouterWeightDeployment | None,
|
||||
configured: object | None,
|
||||
) -> None:
|
||||
reference: Final = (
|
||||
stored
|
||||
if stored is not None
|
||||
else (
|
||||
_RouterWeightDeployment.model_validate(configured, from_attributes=True) if configured is not None else None
|
||||
)
|
||||
)
|
||||
if (
|
||||
reference is None
|
||||
or (stored is None and reference.model_info.db_model)
|
||||
or (reference.model_info.team_id is not None and reference.model_info.team_id != team_id)
|
||||
):
|
||||
raise HTTPException(status_code=400, detail=f"Unknown deployment ID in router weights: {deployment_id}")
|
||||
canonical_group: Final = (
|
||||
reference.model_info.team_public_model_name if reference.model_info.team_id is not None else None
|
||||
) or reference.model_name
|
||||
if model_group != canonical_group:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Deployment {deployment_id} does not belong to model group {model_group}",
|
||||
)
|
||||
|
||||
|
||||
async def validate_router_settings_weights(
|
||||
router_settings: BaseModel | Mapping[str, object] | None,
|
||||
*,
|
||||
team_id: str | None,
|
||||
prisma_client: _PrismaClient | None,
|
||||
llm_router: _Router | None,
|
||||
) -> None:
|
||||
try:
|
||||
weights: Final = (
|
||||
_RouterWeightSettings.model_validate(router_settings, from_attributes=True).weights
|
||||
if router_settings is not None
|
||||
else None
|
||||
)
|
||||
except ValidationError:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Invalid router weights. Replace or clear router_settings.weights.",
|
||||
) from None
|
||||
if not weights:
|
||||
return
|
||||
deployment_ids: Final = frozenset(deployment_id for group in weights.values() for deployment_id in group)
|
||||
if not deployment_ids:
|
||||
return
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=503, detail="Database unavailable while validating router weights")
|
||||
stored_models: Final = await prisma_client.db.litellm_proxymodeltable.find_many(
|
||||
where={"model_id": {"in": list(deployment_ids)}}
|
||||
)
|
||||
stored_by_id: Final = {
|
||||
row.model_id: _RouterWeightDeployment.model_validate(row, from_attributes=True) for row in stored_models
|
||||
}
|
||||
for model_group, group_weights in weights.items():
|
||||
for deployment_id in group_weights:
|
||||
_validate_router_weight_reference(
|
||||
model_group,
|
||||
deployment_id,
|
||||
team_id,
|
||||
stored_by_id.get(deployment_id),
|
||||
llm_router.get_deployment(model_id=deployment_id) if llm_router is not None else None,
|
||||
)
|
||||
|
|
@ -112,6 +112,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
from litellm.proxy.management_endpoints.organization_endpoints import (
|
||||
add_member_to_organization,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.router_weights import validate_router_settings_weights
|
||||
from litellm.proxy.management_endpoints.tag_management_endpoints import (
|
||||
get_daily_activity,
|
||||
)
|
||||
|
|
@ -1288,6 +1289,7 @@ async def new_team(
|
|||
create_audit_log_for_update,
|
||||
general_settings,
|
||||
litellm_proxy_admin_name,
|
||||
llm_router,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
|
@ -1462,6 +1464,13 @@ async def new_team(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
await validate_router_settings_weights(
|
||||
data.router_settings,
|
||||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
## ADD TO MODEL TABLE
|
||||
_model_id = None
|
||||
if data.model_aliases is not None and isinstance(data.model_aliases, dict):
|
||||
|
|
@ -2075,6 +2084,13 @@ async def update_team(
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
await validate_router_settings_weights(
|
||||
data.router_settings,
|
||||
team_id=data.team_id,
|
||||
prisma_client=prisma_client,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
_existing_team_metadata: Final[object] = getattr(existing_team_row, "metadata", None)
|
||||
enforce_output_token_estimates_are_admin_only(
|
||||
data=data,
|
||||
|
|
|
|||
|
|
@ -3592,6 +3592,7 @@ class SSOAuthenticationHandler:
|
|||
verbose_proxy_logger.info("user_defined_values for creating ui key: %s", user_defined_values)
|
||||
|
||||
response: Final = await generate_key_helper_fn(
|
||||
llm_router=None,
|
||||
request_type="key",
|
||||
duration=LITELLM_UI_SESSION_DURATION,
|
||||
key_max_budget=litellm.max_ui_session_budget,
|
||||
|
|
|
|||
871
litellm/proxy/management_helpers/bulk_user_creation.py
Normal file
871
litellm/proxy/management_helpers/bulk_user_creation.py
Normal file
|
|
@ -0,0 +1,871 @@
|
|||
"""Batched internal user creation behind `POST /management/v1/users/bulk`.
|
||||
|
||||
The batch is validated with set queries, user rows land in one `create_many`, and every
|
||||
referenced team is written once under its advisory lock instead of once per user.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal, TypeAlias, TypeVar
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_TeamTable,
|
||||
LitellmUserRoles,
|
||||
Member,
|
||||
NewUserRequestTeam,
|
||||
OrganizationMemberAddRequest,
|
||||
OrgMember,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.auth_checks import invalidate_team_member_spend_state
|
||||
from litellm.proxy.auth.litellm_license import LicenseCheck
|
||||
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventHooks
|
||||
from litellm.proxy.list_api.common import PROBLEM_TYPE_BASE, ManagementProblem
|
||||
from litellm.proxy.management_endpoints.common_utils import (
|
||||
_is_user_org_admin_for_team, # pyright: ignore[reportPrivateUsage] # same team-admin check /user/new uses
|
||||
_is_user_team_admin, # pyright: ignore[reportPrivateUsage] # same team-admin check /user/new uses
|
||||
validate_budget_duration,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
||||
_update_internal_new_user_params, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # /user/new defaults; result validated below
|
||||
check_if_default_team_set,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
_check_permissions_caller_permission, # pyright: ignore[reportPrivateUsage] # same permission check /user/new uses
|
||||
generate_key_helper_fn, # pyright: ignore[reportUnknownVariableType] # legacy untyped helper; result validated by _KEY_RESPONSE
|
||||
metadata_json_with_limits,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.organization_endpoints import organization_member_add
|
||||
from litellm.proxy.management_helpers.access_group_team_sync import TEAM_ADVISORY_LOCK_SQL
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
_set_object_permission, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared with /user/new; result validated below
|
||||
)
|
||||
from litellm.proxy.management_helpers.utils import (
|
||||
_resolve_member_budget_id, # pyright: ignore[reportPrivateUsage] # shared with /team/member_add
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.repositories.team_repository import TeamRepository
|
||||
from litellm.repositories.user_repository import UserRepository
|
||||
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
|
||||
BulkNewUserItem,
|
||||
BulkNewUserMeta,
|
||||
BulkNewUserResponse,
|
||||
UserCreateResult,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import Prisma
|
||||
from prisma import models as prisma_models
|
||||
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
BULK_NEW_USER_CONCURRENCY: Final = 10
|
||||
|
||||
TeamRole: TypeAlias = Literal["user", "admin"]
|
||||
KeyGenerator: TypeAlias = Callable[..., Awaitable[object]]
|
||||
_T: Final = TypeVar("_T")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _RowFailure:
|
||||
index: int
|
||||
user_id: str | None
|
||||
user_email: str | None
|
||||
error: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PendingUser:
|
||||
index: int
|
||||
request: BulkNewUserItem
|
||||
user_id: str
|
||||
teams: tuple[NewUserRequestTeam, ...]
|
||||
|
||||
|
||||
class _UserRow(BaseModel):
|
||||
"""The `/user/new` body after defaults and object permission were applied."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
user_id: str
|
||||
user_email: str | None = None
|
||||
user_alias: str | None = None
|
||||
user_role: str | None = None
|
||||
team_id: str | None = None
|
||||
max_budget: float | None = None
|
||||
spend: float | None = 0.0
|
||||
models: tuple[str, ...] | None = None
|
||||
metadata: Mapping[str, object] | None = None
|
||||
max_parallel_requests: int | None = None
|
||||
tpm_limit: int | None = None
|
||||
rpm_limit: int | None = None
|
||||
budget_duration: str | None = None
|
||||
allowed_cache_controls: tuple[str, ...] | None = None
|
||||
sso_user_id: str | None = None
|
||||
object_permission_id: str | None = None
|
||||
model_max_budget: Mapping[str, object] | None = None
|
||||
model_rpm_limit: Mapping[str, object] | None = None
|
||||
model_tpm_limit: Mapping[str, object] | None = None
|
||||
mcp_rpm_limit: Mapping[str, int] | None = None
|
||||
tag_rpm_limit: Mapping[str, int] | None = None
|
||||
guardrails: tuple[str, ...] | None = None
|
||||
policies: tuple[str, ...] | None = None
|
||||
prompts: tuple[str, ...] | None = None
|
||||
duration: str | None = None
|
||||
key_alias: str | None = None
|
||||
aliases: Mapping[str, object] | None = None
|
||||
config: Mapping[str, object] | None = None
|
||||
permissions: Mapping[str, object] | None = None
|
||||
blocked: bool | None = None
|
||||
agent_id: str | None = None
|
||||
budget_fallbacks: Mapping[str, tuple[str, ...]] | None = None
|
||||
budget_limits: tuple[Mapping[str, object], ...] | None = None
|
||||
organizations: tuple[str, ...] | None = None
|
||||
|
||||
|
||||
_USER_ROW: Final = TypeAdapter(_UserRow)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PreparedUser:
|
||||
pending: _PendingUser
|
||||
row: _UserRow
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _TeamAssignment:
|
||||
user_id: str
|
||||
user_email: str | None
|
||||
role: TeamRole
|
||||
max_budget_in_team: float | None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _TeamWrite:
|
||||
"""Outcome of one locked roster write. `failed` maps user ids to the reason they were not added."""
|
||||
|
||||
team_id: str
|
||||
after: tuple[Member, ...]
|
||||
added: frozenset[str]
|
||||
failed: Mapping[str, str]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _CreatedUser:
|
||||
prepared: _PreparedUser
|
||||
teams: tuple[str, ...]
|
||||
key: str | None
|
||||
errors: tuple[str, ...]
|
||||
|
||||
|
||||
_ERROR_DETAIL: Final = TypeAdapter(Mapping[str, object])
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
class _KeyResponse(BaseModel):
|
||||
token: str
|
||||
|
||||
|
||||
_KEY_RESPONSE: Final = TypeAdapter(_KeyResponse)
|
||||
|
||||
|
||||
def _error_message(exc: BaseException) -> str:
|
||||
if not isinstance(exc, HTTPException):
|
||||
return str(exc)
|
||||
try:
|
||||
detail: Final = _ERROR_DETAIL.validate_python(exc.detail)
|
||||
except ValidationError:
|
||||
return str(exc.detail)
|
||||
return str(detail.get("error", detail))
|
||||
|
||||
|
||||
def _requested_teams(item: BulkNewUserItem) -> tuple[NewUserRequestTeam, ...]:
|
||||
if item.team_id is not None:
|
||||
return (NewUserRequestTeam(team_id=item.team_id),)
|
||||
teams: Final = item.teams if item.teams is not None else check_if_default_team_set()
|
||||
if teams is None:
|
||||
return ()
|
||||
return tuple(team if isinstance(team, NewUserRequestTeam) else NewUserRequestTeam(team_id=team) for team in teams)
|
||||
|
||||
|
||||
def _row_error(item: BulkNewUserItem, user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
||||
if (
|
||||
item.user_role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)
|
||||
and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN
|
||||
):
|
||||
return (
|
||||
"Only proxy admins can create administrative users (proxy_admin, proxy_admin_viewer). "
|
||||
f"Attempted to create user with role: {item.user_role}. Your role: {user_api_key_dict.user_role}"
|
||||
)
|
||||
try:
|
||||
validate_budget_duration(item.budget_duration)
|
||||
_check_permissions_caller_permission(data=item, user_api_key_dict=user_api_key_dict)
|
||||
except Exception as exc: # noqa: BLE001 # any validation failure is reported on this row only
|
||||
return _error_message(exc)
|
||||
return None
|
||||
|
||||
|
||||
def _normalized_email(email: str | None) -> str | None:
|
||||
return email.strip().lower() if email else None
|
||||
|
||||
|
||||
def _partition_rows(
|
||||
users: Sequence[BulkNewUserItem], user_api_key_dict: UserAPIKeyAuth
|
||||
) -> tuple[tuple[_PendingUser, ...], tuple[_RowFailure, ...]]:
|
||||
"""Assign ids, run the per-row checks and fail later rows that repeat an earlier row's id or email."""
|
||||
user_ids: Final = tuple(item.user_id or str(uuid.uuid4()) for item in users)
|
||||
first_index_by_id: Final = MappingProxyType(
|
||||
{user_id: index for index, user_id in reversed(tuple(enumerate(user_ids)))}
|
||||
)
|
||||
first_index_by_email: Final = MappingProxyType(
|
||||
{
|
||||
email: index
|
||||
for index, email in reversed(tuple(enumerate(_normalized_email(item.user_email) for item in users)))
|
||||
if email is not None
|
||||
}
|
||||
)
|
||||
|
||||
def classify(index: int, item: BulkNewUserItem) -> _PendingUser | _RowFailure:
|
||||
user_id: Final = user_ids[index]
|
||||
email: Final = _normalized_email(item.user_email)
|
||||
if first_index_by_id[user_id] != index:
|
||||
return _RowFailure(index, user_id, item.user_email, f"Duplicate user_id in request: {user_id}")
|
||||
if email is not None and first_index_by_email[email] != index:
|
||||
return _RowFailure(index, user_id, item.user_email, f"Duplicate user_email in request: {item.user_email}")
|
||||
error: Final = _row_error(item, user_api_key_dict)
|
||||
if error is not None:
|
||||
return _RowFailure(index, user_id, item.user_email, error)
|
||||
return _PendingUser(index, item, user_id, _requested_teams(item))
|
||||
|
||||
outcomes: Final = tuple(classify(index, item) for index, item in enumerate(users))
|
||||
return (
|
||||
tuple(outcome for outcome in outcomes if isinstance(outcome, _PendingUser)),
|
||||
tuple(outcome for outcome in outcomes if isinstance(outcome, _RowFailure)),
|
||||
)
|
||||
|
||||
|
||||
def _user_table(prisma_client: PrismaClient) -> "TableActions[prisma_models.LiteLLM_UserTable]":
|
||||
return UserRepository(prisma_client).table
|
||||
|
||||
|
||||
async def _existing_user_conflicts(
|
||||
prisma_client: PrismaClient, pending: Sequence[_PendingUser]
|
||||
) -> tuple[frozenset[str], frozenset[str]]:
|
||||
"""Return the requested user ids and (lowercased) emails that already exist, using one query each."""
|
||||
user_ids: Final = sorted(user.user_id for user in pending)
|
||||
emails: Final = sorted(frozenset(user.request.user_email for user in pending if user.request.user_email))
|
||||
if not user_ids:
|
||||
return frozenset(), frozenset()
|
||||
table: Final = _user_table(prisma_client)
|
||||
id_filter: Final = {"user_id": {"in": user_ids}} # mutable-ok: Prisma query filters are dict-shaped
|
||||
email_filter: Final = {"user_email": {"in": emails, "mode": "insensitive"}} # mutable-ok: Prisma filter
|
||||
id_rows: Final = await table.find_many(where=id_filter)
|
||||
email_rows: Final = await table.find_many(where=email_filter) if emails else ()
|
||||
return (
|
||||
frozenset(row.user_id for row in id_rows),
|
||||
frozenset(lowered for row in email_rows if (lowered := _normalized_email(row.user_email)) is not None),
|
||||
)
|
||||
|
||||
|
||||
async def _load_teams(prisma_client: PrismaClient, team_ids: frozenset[str]) -> Mapping[str, LiteLLM_TeamTable]:
|
||||
if not team_ids:
|
||||
return MappingProxyType({})
|
||||
rows: Final = await TeamRepository(prisma_client).table.find_many(
|
||||
where={"team_id": {"in": sorted(team_ids)}} # mutable-ok: Prisma query filters are dict-shaped
|
||||
)
|
||||
return MappingProxyType({row.team_id: LiteLLM_TeamTable.model_validate(row.model_dump()) for row in rows})
|
||||
|
||||
|
||||
async def _team_permission_error(team: LiteLLM_TeamTable, user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
|
||||
return None
|
||||
if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team):
|
||||
return None
|
||||
if await _is_user_org_admin_for_team(user_api_key_dict=user_api_key_dict, team_obj=team):
|
||||
return None
|
||||
return f"Call not allowed. User not proxy admin OR team admin. team_id={team.team_id}"
|
||||
|
||||
|
||||
async def _unusable_teams(
|
||||
prisma_client: PrismaClient,
|
||||
pending: Sequence[_PendingUser],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> tuple[Mapping[str, LiteLLM_TeamTable], Mapping[str, str]]:
|
||||
"""Load every referenced team once and explain, per team id, why rows naming it cannot proceed."""
|
||||
team_ids: Final = frozenset(team.team_id for user in pending for team in user.teams)
|
||||
teams: Final = await _load_teams(prisma_client, team_ids)
|
||||
permission_errors: Final = await asyncio.gather(
|
||||
*(_team_permission_error(team, user_api_key_dict) for team in teams.values())
|
||||
)
|
||||
missing: Final = tuple(
|
||||
(team_id, f"Team id={team_id} does not exist") for team_id in team_ids if team_id not in teams
|
||||
)
|
||||
denied: Final = tuple(
|
||||
(team.team_id, error)
|
||||
for team, error in zip(teams.values(), permission_errors, strict=True)
|
||||
if error is not None
|
||||
)
|
||||
return teams, MappingProxyType({team_id: error for team_id, error in (*missing, *denied)})
|
||||
|
||||
|
||||
def _db_failure(
|
||||
user: _PendingUser,
|
||||
existing_ids: frozenset[str],
|
||||
existing_emails: frozenset[str],
|
||||
team_errors: Mapping[str, str],
|
||||
) -> _RowFailure | None:
|
||||
email: Final = _normalized_email(user.request.user_email)
|
||||
if user.user_id in existing_ids:
|
||||
return _RowFailure(user.index, user.user_id, user.request.user_email, f"User id={user.user_id} already exists")
|
||||
if email is not None and email in existing_emails:
|
||||
return _RowFailure(
|
||||
user.index, user.user_id, user.request.user_email, f"User email={user.request.user_email} already exists"
|
||||
)
|
||||
errors: Final = tuple(team_errors[team.team_id] for team in user.teams if team.team_id in team_errors)
|
||||
if errors:
|
||||
return _RowFailure(user.index, user.user_id, user.request.user_email, "; ".join(errors))
|
||||
return None
|
||||
|
||||
|
||||
async def _prepare_user(user: _PendingUser, prisma_client: PrismaClient) -> _PreparedUser | _RowFailure:
|
||||
try:
|
||||
dumped: Final = user.request.model_dump(exclude={"user_id"}) # mutable-ok: pydantic IncEx takes a set
|
||||
data: Final = {**dumped, "user_id": user.user_id} # mutable-ok: /user/new defaults helper mutates in place
|
||||
data_json: Final = _JSON_OBJECT.validate_python(_update_internal_new_user_params(data, user.request))
|
||||
with_permission: Final = _JSON_OBJECT.validate_python(
|
||||
await _set_object_permission(data_json=data_json, prisma_client=prisma_client) # pyright: ignore[reportUnknownArgumentType] # validated by the adapter
|
||||
)
|
||||
return _PreparedUser(user, _USER_ROW.validate_python(with_permission))
|
||||
except Exception as exc: # noqa: BLE001 # any preparation failure is reported on this row only
|
||||
verbose_proxy_logger.warning("/user/bulk_new: could not prepare row %d - %s", user.index, type(exc).__name__)
|
||||
return _RowFailure(user.index, user.user_id, user.request.user_email, _error_message(exc))
|
||||
|
||||
|
||||
class _UserCreateData(TypedDict):
|
||||
"""One `LiteLLM_UserTable` row as `create_many` takes it; JSON columns are pre-serialized."""
|
||||
|
||||
user_id: ReadOnly[str]
|
||||
user_email: ReadOnly[str | None]
|
||||
user_alias: ReadOnly[str | None]
|
||||
user_role: ReadOnly[str | None]
|
||||
team_id: ReadOnly[str | None]
|
||||
max_budget: ReadOnly[float | None]
|
||||
spend: ReadOnly[float]
|
||||
models: ReadOnly[tuple[str, ...]]
|
||||
metadata: ReadOnly[str]
|
||||
max_parallel_requests: ReadOnly[int | None]
|
||||
tpm_limit: ReadOnly[int | None]
|
||||
rpm_limit: ReadOnly[int | None]
|
||||
budget_duration: ReadOnly[str | None]
|
||||
budget_reset_at: ReadOnly[datetime | None]
|
||||
allowed_cache_controls: ReadOnly[tuple[str, ...]]
|
||||
sso_user_id: ReadOnly[str | None]
|
||||
object_permission_id: ReadOnly[str | None]
|
||||
teams: ReadOnly[tuple[str, ...]]
|
||||
model_max_budget: ReadOnly[str]
|
||||
|
||||
|
||||
def _user_create_payload(prepared: _PreparedUser) -> _UserCreateData:
|
||||
row: Final = prepared.row
|
||||
metadata_json: Final = metadata_json_with_limits(
|
||||
row.metadata,
|
||||
model_rpm_limit=row.model_rpm_limit,
|
||||
model_tpm_limit=row.model_tpm_limit,
|
||||
mcp_rpm_limit=row.mcp_rpm_limit,
|
||||
tag_rpm_limit=row.tag_rpm_limit,
|
||||
guardrails=row.guardrails,
|
||||
policies=row.policies,
|
||||
prompts=row.prompts,
|
||||
)
|
||||
payload: Final[_UserCreateData] = {
|
||||
"user_id": row.user_id,
|
||||
"user_email": row.user_email,
|
||||
"user_alias": row.user_alias,
|
||||
"user_role": row.user_role,
|
||||
"team_id": row.team_id,
|
||||
"max_budget": row.max_budget,
|
||||
"spend": row.spend or 0.0,
|
||||
"models": row.models or (),
|
||||
"metadata": metadata_json,
|
||||
"max_parallel_requests": row.max_parallel_requests,
|
||||
"tpm_limit": row.tpm_limit,
|
||||
"rpm_limit": row.rpm_limit,
|
||||
"budget_duration": row.budget_duration,
|
||||
"budget_reset_at": get_budget_reset_time(row.budget_duration) if row.budget_duration else None,
|
||||
"allowed_cache_controls": row.allowed_cache_controls or (),
|
||||
"sso_user_id": row.sso_user_id,
|
||||
"object_permission_id": row.object_permission_id,
|
||||
"teams": tuple(team.team_id for team in prepared.pending.teams),
|
||||
"model_max_budget": json.dumps(row.model_max_budget) if row.model_max_budget else "{}",
|
||||
}
|
||||
return payload
|
||||
|
||||
|
||||
async def _bounded(limit: int, awaitables: Sequence[Awaitable[_T]]) -> tuple[_T | BaseException, ...]:
|
||||
semaphore: Final = asyncio.Semaphore(limit)
|
||||
|
||||
async def run(awaitable: Awaitable[_T]) -> _T:
|
||||
async with semaphore:
|
||||
return await awaitable
|
||||
|
||||
return tuple(await asyncio.gather(*(run(awaitable) for awaitable in awaitables), return_exceptions=True))
|
||||
|
||||
|
||||
async def _insert_users(
|
||||
prisma_client: PrismaClient, prepared: Sequence[_PreparedUser]
|
||||
) -> tuple[tuple[_PreparedUser, ...], tuple[_RowFailure, ...]]:
|
||||
"""Insert every row in one statement. If that fails, retry rows one at a time so the error lands on its row."""
|
||||
if not prepared:
|
||||
return (), ()
|
||||
table: Final = _user_table(prisma_client)
|
||||
payloads: Final = tuple(_user_create_payload(user) for user in prepared)
|
||||
try:
|
||||
await table.create_many(data=payloads)
|
||||
return tuple(prepared), ()
|
||||
except Exception as exc: # noqa: BLE001 # fall back to per-row inserts so the failing row can be identified
|
||||
verbose_proxy_logger.warning("/user/bulk_new: create_many failed, retrying rows individually", exc_info=True)
|
||||
outcome_unknown: Final = PrismaDBExceptionHandler.is_database_infrastructure_error(exc)
|
||||
requested: Final = frozenset(payload["user_id"] for payload in payloads)
|
||||
landed_rows: Final = await table.find_many(where={"user_id": {"in": list(requested)}}) # mutable-ok: Prisma filter
|
||||
landed: Final = frozenset(row.user_id for row in landed_rows)
|
||||
# create_many is one INSERT: after a lost response the full set is ours, any partial set belongs to another request
|
||||
if outcome_unknown and landed == requested:
|
||||
return tuple(prepared), ()
|
||||
taken: Final = tuple(user for user in prepared if user.row.user_id in landed)
|
||||
retried: Final = tuple(user for user in prepared if user.row.user_id not in landed)
|
||||
outcomes: Final = await _bounded(
|
||||
BULK_NEW_USER_CONCURRENCY, tuple(table.create(data=_user_create_payload(user)) for user in retried)
|
||||
)
|
||||
failed: Final = MappingProxyType(
|
||||
{
|
||||
**{
|
||||
user.row.user_id: _RowFailure(
|
||||
user.pending.index,
|
||||
user.pending.user_id,
|
||||
user.row.user_email,
|
||||
f"User id={user.row.user_id} already exists",
|
||||
)
|
||||
for user in taken
|
||||
},
|
||||
**{
|
||||
user.row.user_id: _RowFailure(
|
||||
user.pending.index, user.pending.user_id, user.row.user_email, _error_message(outcome)
|
||||
)
|
||||
for user, outcome in zip(retried, outcomes, strict=True)
|
||||
if isinstance(outcome, BaseException)
|
||||
},
|
||||
}
|
||||
)
|
||||
return (
|
||||
tuple(user for user in prepared if user.row.user_id not in failed),
|
||||
tuple(failed.values()),
|
||||
)
|
||||
|
||||
|
||||
def _assignments_by_team(created: Sequence[_PreparedUser]) -> Mapping[str, tuple[_TeamAssignment, ...]]:
|
||||
team_ids: Final = tuple(dict.fromkeys(team.team_id for user in created for team in user.pending.teams))
|
||||
return MappingProxyType(
|
||||
{
|
||||
team_id: tuple(
|
||||
_TeamAssignment(user.pending.user_id, user.row.user_email, team.user_role, team.max_budget_in_team)
|
||||
for user in created
|
||||
for team in user.pending.teams
|
||||
if team.team_id == team_id
|
||||
)
|
||||
for team_id in team_ids
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class _MembershipData(TypedDict):
|
||||
team_id: ReadOnly[str]
|
||||
user_id: ReadOnly[str]
|
||||
budget_id: ReadOnly[str | None]
|
||||
|
||||
|
||||
class _RosterData(TypedDict):
|
||||
members_with_roles: ReadOnly[str]
|
||||
|
||||
|
||||
class _TeamsData(TypedDict):
|
||||
teams: ReadOnly[tuple[str, ...]]
|
||||
|
||||
|
||||
def _default_member_budget_id(team: LiteLLM_TeamTable) -> str | None:
|
||||
metadata: Final = (
|
||||
_JSON_OBJECT.validate_python(
|
||||
team.metadata # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_TeamTable.metadata is a bare dict; validated by the adapter
|
||||
)
|
||||
if team.metadata # pyright: ignore[reportUnknownMemberType] # same bare dict
|
||||
else None
|
||||
)
|
||||
budget_id: Final = metadata.get("team_member_budget_id") if metadata is not None else None
|
||||
return budget_id if isinstance(budget_id, str) else None
|
||||
|
||||
|
||||
def _team_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_TeamTable]":
|
||||
return tx.litellm_teamtable # pyright: ignore[reportReturnType] # TableActions widens the generated inputs to Mapping, as the repositories do
|
||||
|
||||
|
||||
def _membership_tx_db(tx: "Prisma") -> "TableActions[prisma_models.LiteLLM_TeamMembership]":
|
||||
return tx.litellm_teammembership # pyright: ignore[reportReturnType] # TableActions widens the generated inputs to Mapping, as the repositories do
|
||||
|
||||
|
||||
async def _write_team_roster(
|
||||
prisma_client: PrismaClient,
|
||||
team: LiteLLM_TeamTable,
|
||||
members: Sequence[_TeamAssignment],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_proxy_admin_name: str,
|
||||
) -> _TeamWrite:
|
||||
"""Add every new member to one team under its advisory lock: one roster rewrite and one membership insert."""
|
||||
try:
|
||||
async with prisma_client.tx() as tx:
|
||||
await tx.query_raw(TEAM_ADVISORY_LOCK_SQL, team.team_id)
|
||||
roster: Final = await TeamRepository(prisma_client).get_members_with_roles_locked(tx, team.team_id)
|
||||
if roster is None:
|
||||
raise ValueError(f"Team id={team.team_id} does not exist")
|
||||
already_present: Final = frozenset(member.user_id for member in roster if member.user_id)
|
||||
new_members: Final = tuple(member for member in members if member.user_id not in already_present)
|
||||
budget_ids: Final = tuple(
|
||||
[ # mutable-ok: budgets are created one at a time on the transaction's single connection
|
||||
await _resolve_member_budget_id(
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
max_budget_in_team=member.max_budget_in_team,
|
||||
allowed_models=team.default_team_member_models or None,
|
||||
budget_duration=None,
|
||||
default_team_budget_id=_default_member_budget_id(team),
|
||||
tx=tx, # pyright: ignore[reportArgumentType] # MemberWriteTx lags the generated Prisma signatures, same as /team/member_add
|
||||
)
|
||||
for member in new_members
|
||||
]
|
||||
)
|
||||
await _membership_tx_db(tx).create_many(
|
||||
data=tuple(
|
||||
_MembershipData(team_id=team.team_id, user_id=member.user_id, budget_id=budget_id)
|
||||
for member, budget_id in zip(new_members, budget_ids, strict=True)
|
||||
),
|
||||
skip_duplicates=True,
|
||||
)
|
||||
after: Final = (
|
||||
*roster,
|
||||
*(Member(user_id=m.user_id, user_email=m.user_email, role=m.role) for m in new_members),
|
||||
)
|
||||
await _team_tx_db(tx).update(
|
||||
where={"team_id": team.team_id}, # mutable-ok: Prisma query filters are dict-shaped
|
||||
data=_RosterData(members_with_roles=json.dumps(tuple(member.model_dump() for member in after))),
|
||||
)
|
||||
return _TeamWrite(
|
||||
team_id=team.team_id,
|
||||
after=after,
|
||||
added=frozenset(member.user_id for member in members),
|
||||
failed=MappingProxyType({}),
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 # the team write failure is reported on each affected row
|
||||
verbose_proxy_logger.exception("/user/bulk_new: failed to add %d members to a team", len(members))
|
||||
message: Final = f"Failed to add user to team {team.team_id}: {_error_message(exc)}"
|
||||
return _TeamWrite(
|
||||
team_id=team.team_id,
|
||||
after=(),
|
||||
added=frozenset(),
|
||||
failed=MappingProxyType({member.user_id: message for member in members}),
|
||||
)
|
||||
|
||||
|
||||
async def _detach_failed_teams(
|
||||
prisma_client: PrismaClient, created: Sequence[_PreparedUser], writes: Mapping[str, _TeamWrite]
|
||||
) -> None:
|
||||
"""Users are inserted with `teams` already set; drop the teams whose roster write did not take them."""
|
||||
table: Final = _user_table(prisma_client)
|
||||
updates: Final = tuple(
|
||||
table.update(
|
||||
where={"user_id": user.row.user_id}, # mutable-ok: Prisma query filters are dict-shaped
|
||||
data=_TeamsData(teams=landed),
|
||||
)
|
||||
for user in created
|
||||
if (landed := _row_teams(user, writes)[0]) != tuple(team.team_id for team in user.pending.teams)
|
||||
)
|
||||
for outcome in await _bounded(BULK_NEW_USER_CONCURRENCY, updates):
|
||||
if isinstance(outcome, BaseException):
|
||||
verbose_proxy_logger.warning(
|
||||
"/user/bulk_new: could not detach failed teams from user - %s", type(outcome).__name__
|
||||
)
|
||||
|
||||
|
||||
async def _publish_team_writes(writes: Sequence[_TeamWrite], user_api_key_cache: "UserApiKeyCache") -> None:
|
||||
prometheus_logger: Final = PrometheusLogger.get_instance()
|
||||
for write in writes:
|
||||
if prometheus_logger is None or not write.added:
|
||||
continue
|
||||
try:
|
||||
prometheus_logger.set_team_members_metric(
|
||||
LiteLLM_TeamTable(
|
||||
team_id=write.team_id,
|
||||
members_with_roles=write.after, # pyright: ignore[reportArgumentType] # pydantic coerces the tuple into the declared list
|
||||
)
|
||||
)
|
||||
except Exception: # noqa: BLE001 # metrics are best-effort and must not fail the request
|
||||
verbose_proxy_logger.debug("Prometheus: failed to emit team members metric", exc_info=True)
|
||||
evictions: Final = await _bounded(
|
||||
BULK_NEW_USER_CONCURRENCY,
|
||||
tuple(
|
||||
invalidate_team_member_spend_state(
|
||||
user_id=user_id, team_id=write.team_id, user_api_key_cache=user_api_key_cache
|
||||
)
|
||||
for write in writes
|
||||
for user_id in write.added
|
||||
),
|
||||
)
|
||||
for eviction in evictions:
|
||||
if isinstance(eviction, BaseException):
|
||||
verbose_proxy_logger.warning("/user/bulk_new: cache eviction failed - %s", type(eviction).__name__)
|
||||
|
||||
|
||||
_KEY_FIELDS: Final = MappingProxyType(
|
||||
{
|
||||
name: True
|
||||
for name in (
|
||||
"user_id",
|
||||
"team_id",
|
||||
"agent_id",
|
||||
"duration",
|
||||
"key_alias",
|
||||
"models",
|
||||
"aliases",
|
||||
"config",
|
||||
"permissions",
|
||||
"blocked",
|
||||
"spend",
|
||||
"budget_fallbacks",
|
||||
"budget_limits",
|
||||
"metadata",
|
||||
"max_parallel_requests",
|
||||
"tpm_limit",
|
||||
"rpm_limit",
|
||||
"allowed_cache_controls",
|
||||
"model_max_budget",
|
||||
"model_rpm_limit",
|
||||
"model_tpm_limit",
|
||||
"mcp_rpm_limit",
|
||||
"tag_rpm_limit",
|
||||
"guardrails",
|
||||
"policies",
|
||||
"prompts",
|
||||
"object_permission_id",
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def _generate_key(prepared: _PreparedUser, generate_key: KeyGenerator) -> str:
|
||||
response: Final = _KEY_RESPONSE.validate_python(
|
||||
await generate_key(
|
||||
request_type="key", table_name="key", **prepared.row.model_dump(include=_KEY_FIELDS, exclude_none=True)
|
||||
)
|
||||
)
|
||||
return response.token
|
||||
|
||||
|
||||
async def _add_to_organizations(
|
||||
prepared: _PreparedUser, organizations: Sequence[str], user_api_key_dict: UserAPIKeyAuth
|
||||
) -> None:
|
||||
for organization_id in organizations:
|
||||
await organization_member_add(
|
||||
data=OrganizationMemberAddRequest(
|
||||
organization_id=organization_id,
|
||||
member=OrgMember(user_id=prepared.row.user_id, role=LitellmUserRoles.INTERNAL_USER),
|
||||
),
|
||||
http_request=Request(scope={"type": "http", "path": "/user/bulk_new"}), # mutable-ok: ASGI scopes are dicts
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
|
||||
async def _run_per_user(
|
||||
created: Sequence[_PreparedUser],
|
||||
select: Callable[[_PreparedUser], bool],
|
||||
action: Callable[[_PreparedUser], Awaitable[_T]],
|
||||
) -> Mapping[str, _T | BaseException]:
|
||||
chosen: Final = tuple(user for user in created if select(user))
|
||||
outcomes: Final = await _bounded(BULK_NEW_USER_CONCURRENCY, tuple(action(user) for user in chosen))
|
||||
return MappingProxyType({user.row.user_id: outcome for user, outcome in zip(chosen, outcomes, strict=True)})
|
||||
|
||||
|
||||
async def _write_audit_logs(
|
||||
prisma_client: PrismaClient,
|
||||
created: Sequence[_PreparedUser],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_proxy_admin_name: str,
|
||||
) -> None:
|
||||
if not created:
|
||||
return
|
||||
created_ids: Final = sorted(user.row.user_id for user in created)
|
||||
created_filter: Final = {"user_id": {"in": created_ids}} # mutable-ok: Prisma query filters are dict-shaped
|
||||
rows: Final = await _user_table(prisma_client).find_many(where=created_filter)
|
||||
outcomes: Final = await _bounded(
|
||||
BULK_NEW_USER_CONCURRENCY,
|
||||
tuple(
|
||||
UserManagementEventHooks.create_internal_user_audit_log(
|
||||
user_id=row.user_id,
|
||||
action="created",
|
||||
litellm_changed_by=user_api_key_dict.user_id,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
before_value=None,
|
||||
after_value=row.model_dump_json(exclude_none=True),
|
||||
)
|
||||
for row in rows
|
||||
),
|
||||
)
|
||||
for outcome in outcomes:
|
||||
if isinstance(outcome, BaseException):
|
||||
verbose_proxy_logger.warning(
|
||||
"Unable to create audit log for user on `/user/bulk_new` - %s", type(outcome).__name__
|
||||
)
|
||||
|
||||
|
||||
def _row_teams(prepared: _PreparedUser, writes: Mapping[str, _TeamWrite]) -> tuple[tuple[str, ...], tuple[str, ...]]:
|
||||
"""Split a user's requested teams into the ones they landed in and the errors for the ones they did not."""
|
||||
requested: Final = tuple(team.team_id for team in prepared.pending.teams)
|
||||
return (
|
||||
tuple(team_id for team_id in requested if prepared.row.user_id in writes[team_id].added),
|
||||
tuple(
|
||||
writes[team_id].failed[prepared.row.user_id]
|
||||
for team_id in requested
|
||||
if prepared.row.user_id in writes[team_id].failed
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _to_result(created: _CreatedUser) -> UserCreateResult:
|
||||
return UserCreateResult(
|
||||
user_id=created.prepared.row.user_id,
|
||||
user_email=created.prepared.row.user_email,
|
||||
success=True,
|
||||
teams=created.teams,
|
||||
key=created.key,
|
||||
error="; ".join(created.errors) if created.errors else None,
|
||||
)
|
||||
|
||||
|
||||
def _failure_result(failure: _RowFailure) -> UserCreateResult:
|
||||
return UserCreateResult(user_id=failure.user_id, user_email=failure.user_email, success=False, error=failure.error)
|
||||
|
||||
|
||||
async def bulk_create_users(
|
||||
users: Sequence[BulkNewUserItem],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
prisma_client: PrismaClient,
|
||||
license_check: LicenseCheck,
|
||||
litellm_proxy_admin_name: str,
|
||||
user_api_key_cache: "UserApiKeyCache",
|
||||
generate_key: KeyGenerator = generate_key_helper_fn,
|
||||
) -> BulkNewUserResponse:
|
||||
"""Create every valid row in `users`; rows that fail validation or a write are reported, not raised.
|
||||
|
||||
Raises a 403 `ManagementProblem` only when the whole batch would push the deployment over its license seat
|
||||
limit.
|
||||
"""
|
||||
pending, request_failures = _partition_rows(users, user_api_key_dict)
|
||||
existing_ids, existing_emails = await _existing_user_conflicts(prisma_client, pending)
|
||||
teams, team_errors = await _unusable_teams(prisma_client, pending, user_api_key_dict)
|
||||
db_failures: Final = tuple(
|
||||
failure
|
||||
for user in pending
|
||||
if (failure := _db_failure(user, existing_ids, existing_emails, team_errors)) is not None
|
||||
)
|
||||
failed_indexes: Final = frozenset(failure.index for failure in db_failures)
|
||||
creatable: Final = tuple(user for user in pending if user.index not in failed_indexes)
|
||||
|
||||
billable_users: Final = await UserRepository(prisma_client).count_billable_users()
|
||||
if creatable and license_check.is_over_limit(total_users=billable_users + len(creatable)):
|
||||
raise ManagementProblem(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}license-limit-exceeded",
|
||||
title="License limit exceeded",
|
||||
status=403,
|
||||
detail="License is over limit. Please contact support@berri.ai to upgrade your license.",
|
||||
)
|
||||
)
|
||||
|
||||
prepared_outcomes: Final = tuple([await _prepare_user(user, prisma_client) for user in creatable])
|
||||
prepare_failures: Final = tuple(o for o in prepared_outcomes if isinstance(o, _RowFailure))
|
||||
created, insert_failures = await _insert_users(
|
||||
prisma_client, tuple(o for o in prepared_outcomes if isinstance(o, _PreparedUser))
|
||||
)
|
||||
|
||||
team_writes: Final = MappingProxyType(
|
||||
{
|
||||
team_id: await _write_team_roster(
|
||||
prisma_client, teams[team_id], members, user_api_key_dict, litellm_proxy_admin_name
|
||||
)
|
||||
for team_id, members in _assignments_by_team(created).items()
|
||||
}
|
||||
)
|
||||
await _detach_failed_teams(prisma_client, created, team_writes)
|
||||
await _publish_team_writes(tuple(team_writes.values()), user_api_key_cache)
|
||||
|
||||
keys: Final = await _run_per_user(
|
||||
created, lambda user: user.pending.request.auto_create_key, lambda user: _generate_key(user, generate_key)
|
||||
)
|
||||
org_outcomes: Final = await _run_per_user(
|
||||
created,
|
||||
lambda user: bool(user.row.organizations),
|
||||
lambda user: _add_to_organizations(user, user.row.organizations or (), user_api_key_dict),
|
||||
)
|
||||
await _write_audit_logs(prisma_client, created, user_api_key_dict, litellm_proxy_admin_name)
|
||||
|
||||
def finish(prepared: _PreparedUser) -> _CreatedUser:
|
||||
landed, team_failures = _row_teams(prepared, team_writes)
|
||||
key_outcome: Final = keys.get(prepared.row.user_id)
|
||||
org_outcome: Final = org_outcomes.get(prepared.row.user_id)
|
||||
return _CreatedUser(
|
||||
prepared=prepared,
|
||||
teams=landed,
|
||||
key=key_outcome if isinstance(key_outcome, str) else None,
|
||||
errors=(
|
||||
*team_failures,
|
||||
*(
|
||||
(f"Failed to create key: {_error_message(key_outcome)}",)
|
||||
if isinstance(key_outcome, BaseException)
|
||||
else ()
|
||||
),
|
||||
*(
|
||||
(f"Failed to add user to organizations: {_error_message(org_outcome)}",)
|
||||
if isinstance(org_outcome, BaseException)
|
||||
else ()
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
failures: Final = MappingProxyType(
|
||||
{
|
||||
failure.index: _failure_result(failure)
|
||||
for failure in (*request_failures, *db_failures, *prepare_failures, *insert_failures)
|
||||
}
|
||||
)
|
||||
successes_by_index: Final = MappingProxyType({user.pending.index: _to_result(finish(user)) for user in created})
|
||||
results: Final = tuple(
|
||||
failures[index] if index in failures else successes_by_index[index] for index in range(len(users))
|
||||
)
|
||||
successes: Final = sum(1 for result in results if result.success)
|
||||
return BulkNewUserResponse(
|
||||
data=results,
|
||||
meta=BulkNewUserMeta(total_requested=len(users), created=successes, failed=len(users) - successes),
|
||||
)
|
||||
|
|
@ -5,18 +5,105 @@ Handles cost tracking and logging for Vertex AI Live API WebSocket passthrough e
|
|||
Supports different modalities: text, audio, video, and web search.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Any, Final
|
||||
from itertools import chain, pairwise
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.vertex_ai.gemini.grounding_requests import GroundingRequests, calculate_grounding_requests
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough_logging_handler import (
|
||||
BasePassthroughLoggingHandler,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler import (
|
||||
PassThroughEndpointLoggingTypedDict,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders, ModelResponse, Usage
|
||||
from litellm.utils import get_model_info
|
||||
from litellm.types.utils import (
|
||||
CompletionTokensDetailsWrapper,
|
||||
CostBreakdown,
|
||||
LlmProviders,
|
||||
ModelResponse,
|
||||
PromptTokensDetailsWrapper,
|
||||
Usage,
|
||||
)
|
||||
|
||||
_NO_GROUNDING: Final = GroundingRequests(web_search_requests=None, google_maps_grounding_requests=None)
|
||||
|
||||
_AGGREGATED_FIELDS: Final = frozenset(
|
||||
{
|
||||
"promptTokenCount",
|
||||
"candidatesTokenCount",
|
||||
"totalTokenCount",
|
||||
"toolUsePromptTokenCount",
|
||||
"promptTokensDetails",
|
||||
"candidatesTokensDetails",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _detail_entries(raw: object) -> tuple[Mapping[str, object], ...]:
|
||||
"""Narrow one turn's ``*TokensDetails`` value to the entries that are actually shaped like one."""
|
||||
return tuple(entry for entry in raw if isinstance(entry, Mapping)) if isinstance(raw, Sequence) else ()
|
||||
|
||||
|
||||
def _grounding_metadata(websocket_messages: Sequence[object]) -> tuple[Mapping[str, object], ...]:
|
||||
"""Collect every ``serverContent.groundingMetadata`` a session emitted.
|
||||
|
||||
Live reports grounding in the server frames, never in ``usageMetadata``, so the per-query
|
||||
charge has to be counted here rather than derived from the token totals.
|
||||
"""
|
||||
return tuple(
|
||||
metadata
|
||||
for message in websocket_messages
|
||||
if isinstance(message, Mapping)
|
||||
for server_content in (message.get("serverContent"),)
|
||||
if isinstance(server_content, Mapping)
|
||||
for metadata in (server_content.get("groundingMetadata"),)
|
||||
if isinstance(metadata, Mapping)
|
||||
)
|
||||
|
||||
|
||||
def _turns(websocket_messages: Sequence[object]) -> tuple[tuple[object, ...], ...]:
|
||||
"""Split a session at every ``usageMetadata`` frame; frames after the last one never got their usage."""
|
||||
closes: Final = tuple(
|
||||
index + 1
|
||||
for index, message in enumerate(websocket_messages)
|
||||
if isinstance(message, Mapping) and isinstance(message.get("usageMetadata"), dict)
|
||||
)
|
||||
return tuple(tuple(websocket_messages[start:end]) for start, end in pairwise((0, *closes)))
|
||||
|
||||
|
||||
def _session_grounding_requests(websocket_messages: Sequence[object]) -> GroundingRequests:
|
||||
per_turn: Final = tuple(
|
||||
calculate_grounding_requests(_grounding_metadata(turn)) for turn in _turns(websocket_messages)
|
||||
)
|
||||
web_search_requests: Final = sum(requests.web_search_requests or 0 for requests in per_turn)
|
||||
google_maps_grounding_requests: Final = sum(requests.google_maps_grounding_requests or 0 for requests in per_turn)
|
||||
return GroundingRequests(
|
||||
web_search_requests=web_search_requests or None,
|
||||
google_maps_grounding_requests=google_maps_grounding_requests or None,
|
||||
)
|
||||
|
||||
|
||||
_SummedField: TypeAlias = Literal[
|
||||
"input_cost",
|
||||
"output_cost",
|
||||
"tool_usage_cost",
|
||||
"cache_read_cost",
|
||||
"cache_creation_cost",
|
||||
"reasoning_cost",
|
||||
"original_cost",
|
||||
"discount_amount",
|
||||
"margin_fixed_amount",
|
||||
"margin_total_amount",
|
||||
]
|
||||
|
||||
|
||||
def _summed(breakdowns: Sequence[CostBreakdown], field: _SummedField) -> float | None:
|
||||
values: Final = tuple(value for breakdown in breakdowns if (value := breakdown.get(field)) is not None)
|
||||
return sum(values) if values else None
|
||||
|
||||
|
||||
class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
||||
|
|
@ -48,186 +135,110 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
|||
"""Return the LLM provider name."""
|
||||
return LlmProviders.VERTEX_AI
|
||||
|
||||
@staticmethod
|
||||
def _resolve_detail_counts(
|
||||
details: Sequence[Mapping[str, object]],
|
||||
declared_total: object,
|
||||
) -> tuple[tuple[str, int], ...]:
|
||||
"""
|
||||
Pair each of one turn's ``*TokensDetails`` entries with its token count.
|
||||
|
||||
Live sometimes names the modality that carries the rest of a turn without a
|
||||
``tokenCount``, and reading the absent key as zero drops those tokens from the
|
||||
breakdown, so real audio ends up priced as text. A lone unpriced entry therefore takes
|
||||
whatever the turn's declared count leaves over. Two or more cannot be told apart, so
|
||||
they are left out and the cost calculator charges the remainder as text.
|
||||
"""
|
||||
priced: Final = tuple(
|
||||
(str(detail.get("modality", "TEXT")), count)
|
||||
for detail in details
|
||||
if isinstance(count := detail.get("tokenCount"), int)
|
||||
)
|
||||
unpriced: Final = tuple(
|
||||
str(detail.get("modality", "TEXT")) for detail in details if not isinstance(detail.get("tokenCount"), int)
|
||||
)
|
||||
if len(unpriced) != 1 or not isinstance(declared_total, int):
|
||||
return priced
|
||||
residual: Final = declared_total - sum(count for _, count in priced)
|
||||
return priced if residual <= 0 else (*priced, (unpriced[0], residual))
|
||||
|
||||
@staticmethod
|
||||
def _sum_by_modality(counts: Sequence[tuple[str, int]]) -> Mapping[str, int]:
|
||||
"""Total the (modality, tokenCount) pairs of one or more turns per modality."""
|
||||
return MappingProxyType({modality: sum(c for m, c in counts if m == modality) for modality, _ in counts})
|
||||
|
||||
@staticmethod
|
||||
def _merged_modality_totals(
|
||||
snapshots: Sequence[Mapping[str, object]],
|
||||
count_key: str,
|
||||
details_key: str,
|
||||
) -> Mapping[str, int]:
|
||||
"""Total every turn's per-modality counts, so the breakdown adds up the way the totals do."""
|
||||
return VertexAILivePassthroughLoggingHandler._sum_by_modality(
|
||||
tuple(
|
||||
chain.from_iterable(
|
||||
VertexAILivePassthroughLoggingHandler._resolve_detail_counts(
|
||||
_detail_entries(snapshot.get(details_key)), snapshot.get(count_key)
|
||||
)
|
||||
for snapshot in snapshots
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _extract_usage_metadata_from_websocket_messages(
|
||||
websocket_messages: list[dict],
|
||||
websocket_messages: Sequence[object],
|
||||
) -> dict | None:
|
||||
"""
|
||||
Extract and aggregate usage metadata from a list of WebSocket messages.
|
||||
|
||||
Live emits one ``usageMetadata`` per turn and Google charges per turn for every token in
|
||||
the session context window, which is the current turn's tokens plus all accumulated
|
||||
tokens from previous turns, so the turns add up rather than restating each other. See
|
||||
the Live API note under https://cloud.google.com/vertex-ai/generative-ai/pricing.
|
||||
|
||||
Args:
|
||||
websocket_messages: List of WebSocket messages from the Live API
|
||||
|
||||
Returns:
|
||||
Dictionary containing aggregated usage metadata, or None if not found
|
||||
"""
|
||||
all_usage_metadata: Final = []
|
||||
snapshots: Final = tuple(
|
||||
metadata
|
||||
for message in websocket_messages
|
||||
if isinstance(message, Mapping)
|
||||
for metadata in (message.get("usageMetadata"),)
|
||||
if isinstance(metadata, dict)
|
||||
)
|
||||
|
||||
# Collect all usage metadata messages
|
||||
for message in websocket_messages:
|
||||
if isinstance(message, dict) and "usageMetadata" in message:
|
||||
all_usage_metadata.append(message["usageMetadata"])
|
||||
|
||||
if not all_usage_metadata:
|
||||
if not snapshots:
|
||||
return None
|
||||
|
||||
# If only one usage metadata, return it as-is
|
||||
if len(all_usage_metadata) == 1:
|
||||
return all_usage_metadata[0]
|
||||
|
||||
# Aggregate multiple usage metadata messages
|
||||
aggregated: Final[dict[str, Any]] = {
|
||||
"promptTokenCount": 0,
|
||||
"candidatesTokenCount": 0,
|
||||
"totalTokenCount": 0,
|
||||
"promptTokensDetails": [],
|
||||
"candidatesTokensDetails": [],
|
||||
prompt_totals: Final = VertexAILivePassthroughLoggingHandler._merged_modality_totals(
|
||||
snapshots, "promptTokenCount", "promptTokensDetails"
|
||||
)
|
||||
candidate_totals: Final = VertexAILivePassthroughLoggingHandler._merged_modality_totals(
|
||||
snapshots, "candidatesTokenCount", "candidatesTokensDetails"
|
||||
)
|
||||
return {
|
||||
**{key: value for key, value in snapshots[0].items() if key not in _AGGREGATED_FIELDS},
|
||||
"promptTokenCount": sum(snapshot.get("promptTokenCount", 0) for snapshot in snapshots),
|
||||
"candidatesTokenCount": sum(snapshot.get("candidatesTokenCount", 0) for snapshot in snapshots),
|
||||
"totalTokenCount": sum(snapshot.get("totalTokenCount", 0) for snapshot in snapshots),
|
||||
"toolUsePromptTokenCount": sum(snapshot.get("toolUsePromptTokenCount", 0) for snapshot in snapshots),
|
||||
"promptTokensDetails": [
|
||||
{"modality": modality, "tokenCount": count} for modality, count in prompt_totals.items() if count > 0
|
||||
],
|
||||
"candidatesTokensDetails": [
|
||||
{"modality": modality, "tokenCount": count} for modality, count in candidate_totals.items() if count > 0
|
||||
],
|
||||
}
|
||||
|
||||
# Aggregate token counts
|
||||
for usage in all_usage_metadata:
|
||||
aggregated["promptTokenCount"] += usage.get("promptTokenCount", 0)
|
||||
aggregated["candidatesTokenCount"] += usage.get("candidatesTokenCount", 0)
|
||||
aggregated["totalTokenCount"] += usage.get("totalTokenCount", 0)
|
||||
|
||||
# Aggregate token details by modality
|
||||
modality_totals: Final = {}
|
||||
|
||||
for usage in all_usage_metadata:
|
||||
# Process prompt tokens details
|
||||
for detail in usage.get("promptTokensDetails", []):
|
||||
modality = detail.get("modality", "TEXT")
|
||||
token_count = detail.get("tokenCount", 0)
|
||||
|
||||
if modality not in modality_totals:
|
||||
modality_totals[modality] = {"prompt": 0, "candidate": 0}
|
||||
modality_totals[modality]["prompt"] += token_count
|
||||
|
||||
# Process candidate tokens details
|
||||
for detail in usage.get("candidatesTokensDetails", []):
|
||||
modality = detail.get("modality", "TEXT")
|
||||
token_count = detail.get("tokenCount", 0)
|
||||
|
||||
if modality not in modality_totals:
|
||||
modality_totals[modality] = {"prompt": 0, "candidate": 0}
|
||||
modality_totals[modality]["candidate"] += token_count
|
||||
|
||||
# Convert aggregated modality totals back to details format
|
||||
for modality, totals in modality_totals.items():
|
||||
if totals["prompt"] > 0:
|
||||
aggregated["promptTokensDetails"].append({"modality": modality, "tokenCount": totals["prompt"]})
|
||||
if totals["candidate"] > 0:
|
||||
aggregated["candidatesTokensDetails"].append({"modality": modality, "tokenCount": totals["candidate"]})
|
||||
|
||||
# Add any additional fields from the first usage metadata
|
||||
first_usage: Final = all_usage_metadata[0]
|
||||
for key, value in first_usage.items():
|
||||
if key not in aggregated:
|
||||
aggregated[key] = value
|
||||
|
||||
return aggregated
|
||||
|
||||
@staticmethod
|
||||
def _calculate_live_api_cost(
|
||||
model: str,
|
||||
usage_metadata: dict,
|
||||
custom_llm_provider: str = "vertex_ai",
|
||||
) -> float:
|
||||
"""
|
||||
Calculate cost for Vertex AI Live API based on usage metadata.
|
||||
|
||||
Args:
|
||||
model: The model name (e.g., "gemini-2.0-flash-live-preview-04-09")
|
||||
usage_metadata: Usage metadata from the Live API response
|
||||
custom_llm_provider: The LLM provider (default: "vertex_ai")
|
||||
|
||||
Returns:
|
||||
Total cost in USD
|
||||
"""
|
||||
try:
|
||||
# Get model pricing information
|
||||
model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
verbose_proxy_logger.debug("Vertex AI Live API model info for '%s': %s", model, model_info)
|
||||
|
||||
# Check if pricing info is available
|
||||
if not model_info or not model_info.get("input_cost_per_token"):
|
||||
verbose_proxy_logger.error("No pricing info found for %s in local model pricing database", model)
|
||||
return 0.0
|
||||
|
||||
total_cost = 0.0
|
||||
|
||||
# Extract token counts from usage metadata
|
||||
prompt_token_count: Final = usage_metadata.get("promptTokenCount", 0)
|
||||
candidates_token_count: Final = usage_metadata.get("candidatesTokenCount", 0)
|
||||
|
||||
# Calculate base text token costs
|
||||
input_cost_per_token: Final = model_info.get("input_cost_per_token", 0.0)
|
||||
output_cost_per_token: Final = model_info.get("output_cost_per_token", 0.0)
|
||||
|
||||
total_cost += prompt_token_count * input_cost_per_token
|
||||
total_cost += candidates_token_count * output_cost_per_token
|
||||
|
||||
# Handle modality-specific costs if present
|
||||
prompt_tokens_details: Final = usage_metadata.get("promptTokensDetails", [])
|
||||
candidates_tokens_details: Final = usage_metadata.get("candidatesTokensDetails", [])
|
||||
|
||||
# Process prompt tokens by modality
|
||||
for detail in prompt_tokens_details:
|
||||
modality = detail.get("modality", "TEXT")
|
||||
token_count = detail.get("tokenCount", 0)
|
||||
|
||||
if modality == "AUDIO":
|
||||
audio_cost_per_token = model_info.get("input_cost_per_audio_token", 0.0)
|
||||
total_cost += token_count * audio_cost_per_token
|
||||
elif modality == "VIDEO":
|
||||
# Video tokens are typically per second, but we'll treat as per token for now
|
||||
video_cost_per_token = model_info.get("input_cost_per_video_per_second", 0.0)
|
||||
total_cost += token_count * video_cost_per_token
|
||||
# TEXT tokens are already handled above
|
||||
|
||||
# Process candidate tokens by modality
|
||||
for detail in candidates_tokens_details:
|
||||
modality = detail.get("modality", "TEXT")
|
||||
token_count = detail.get("tokenCount", 0)
|
||||
|
||||
if modality == "AUDIO":
|
||||
audio_cost_per_token = model_info.get("output_cost_per_audio_token", 0.0)
|
||||
total_cost += token_count * audio_cost_per_token
|
||||
elif modality == "VIDEO":
|
||||
# Video tokens are typically per second, but we'll treat as per token for now
|
||||
video_cost_per_token = model_info.get("output_cost_per_video_per_second", 0.0)
|
||||
total_cost += token_count * video_cost_per_token
|
||||
# TEXT tokens are already handled above
|
||||
|
||||
# Handle web search costs if present
|
||||
tool_use_prompt_token_count: Final = usage_metadata.get("toolUsePromptTokenCount", 0)
|
||||
if tool_use_prompt_token_count > 0:
|
||||
# Web search typically has a fixed cost per request
|
||||
web_search_cost: Final = model_info.get("web_search_cost_per_request", 0.0)
|
||||
if isinstance(web_search_cost, (int, float)) and web_search_cost > 0:
|
||||
total_cost += web_search_cost
|
||||
else:
|
||||
# Fallback to token-based pricing for tool use
|
||||
total_cost += tool_use_prompt_token_count * input_cost_per_token
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Vertex AI Live API cost calculation - Model: {model}, "
|
||||
f"Prompt tokens: {prompt_token_count}, "
|
||||
f"Candidate tokens: {candidates_token_count}, "
|
||||
f"Total cost: ${total_cost:.6f}"
|
||||
)
|
||||
|
||||
return total_cost
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Error calculating Vertex AI Live API cost: %s", e)
|
||||
return 0.0
|
||||
|
||||
@staticmethod
|
||||
def _create_usage_object_from_metadata(
|
||||
usage_metadata: dict,
|
||||
model: str,
|
||||
grounding_requests: GroundingRequests = _NO_GROUNDING,
|
||||
) -> Usage:
|
||||
"""
|
||||
Create a LiteLLM Usage object from Live API usage metadata.
|
||||
|
|
@ -235,48 +246,124 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
|||
Args:
|
||||
usage_metadata: Usage metadata from the Live API response
|
||||
model: The model name
|
||||
grounding_requests: The Search and Maps grounding requests summed over the session's
|
||||
turns, matching the per-turn charge
|
||||
|
||||
Returns:
|
||||
LiteLLM Usage object
|
||||
"""
|
||||
prompt_tokens: Final = usage_metadata.get("promptTokenCount", 0)
|
||||
completion_tokens: Final = usage_metadata.get("candidatesTokenCount", 0)
|
||||
total_tokens: Final = usage_metadata.get("totalTokenCount", 0)
|
||||
prompt_by_modality: Final = VertexAILivePassthroughLoggingHandler._sum_by_modality(
|
||||
VertexAILivePassthroughLoggingHandler._resolve_detail_counts(
|
||||
_detail_entries(usage_metadata.get("promptTokensDetails")), usage_metadata.get("promptTokenCount")
|
||||
)
|
||||
)
|
||||
candidates_by_modality: Final = VertexAILivePassthroughLoggingHandler._sum_by_modality(
|
||||
VertexAILivePassthroughLoggingHandler._resolve_detail_counts(
|
||||
_detail_entries(usage_metadata.get("candidatesTokensDetails")),
|
||||
usage_metadata.get("candidatesTokenCount"),
|
||||
)
|
||||
)
|
||||
|
||||
# Create modality-specific token details if available
|
||||
prompt_tokens_details: Final = usage_metadata.get("promptTokensDetails", [])
|
||||
candidates_tokens_details: Final = usage_metadata.get("candidatesTokensDetails", [])
|
||||
|
||||
# Extract text tokens from details
|
||||
text_prompt_tokens = 0
|
||||
text_completion_tokens = 0
|
||||
|
||||
for detail in prompt_tokens_details:
|
||||
if detail.get("modality") == "TEXT":
|
||||
text_prompt_tokens = detail.get("tokenCount", 0)
|
||||
break
|
||||
|
||||
for detail in candidates_tokens_details:
|
||||
if detail.get("modality") == "TEXT":
|
||||
text_completion_tokens = detail.get("tokenCount", 0)
|
||||
break
|
||||
|
||||
# If no text tokens found in details, use total counts
|
||||
if text_prompt_tokens == 0:
|
||||
text_prompt_tokens = prompt_tokens
|
||||
if text_completion_tokens == 0:
|
||||
text_completion_tokens = completion_tokens
|
||||
prompt_tokens: Final = usage_metadata.get("promptTokenCount", 0) or sum(prompt_by_modality.values())
|
||||
completion_tokens: Final = usage_metadata.get("candidatesTokenCount", 0) or sum(candidates_by_modality.values())
|
||||
|
||||
return Usage(
|
||||
prompt_tokens=text_prompt_tokens,
|
||||
completion_tokens=text_completion_tokens,
|
||||
total_tokens=total_tokens,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=usage_metadata.get("totalTokenCount", 0) or (prompt_tokens + completion_tokens),
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=prompt_by_modality.get("TEXT"),
|
||||
audio_tokens=prompt_by_modality.get("AUDIO"),
|
||||
image_tokens=prompt_by_modality.get("IMAGE"),
|
||||
video_tokens=prompt_by_modality.get("VIDEO"),
|
||||
tool_use_tokens=usage_metadata.get("toolUsePromptTokenCount") or None,
|
||||
web_search_requests=grounding_requests.web_search_requests,
|
||||
google_maps_grounding_requests=grounding_requests.google_maps_grounding_requests,
|
||||
),
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
text_tokens=candidates_by_modality.get("TEXT"),
|
||||
audio_tokens=candidates_by_modality.get("AUDIO"),
|
||||
image_tokens=candidates_by_modality.get("IMAGE"),
|
||||
video_tokens=candidates_by_modality.get("VIDEO"),
|
||||
),
|
||||
)
|
||||
|
||||
def _session_usage(self, websocket_messages: Sequence[object], model: str) -> Usage | None:
|
||||
usage_metadata: Final = self._extract_usage_metadata_from_websocket_messages(websocket_messages)
|
||||
if usage_metadata is None:
|
||||
return None
|
||||
return self._create_usage_object_from_metadata(
|
||||
usage_metadata=usage_metadata,
|
||||
grounding_requests=_session_grounding_requests(websocket_messages),
|
||||
model=model,
|
||||
)
|
||||
|
||||
def _turn_cost(
|
||||
self,
|
||||
turn: Sequence[object],
|
||||
model: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> tuple[float, CostBreakdown] | None:
|
||||
usage: Final = self._session_usage(turn, model)
|
||||
if usage is None:
|
||||
return None
|
||||
cost: Final = logging_obj._response_cost_calculator( # pyright: ignore[reportPrivateUsage] # the call's own calculator keeps custom pricing and the deployment's region in step with the spend row
|
||||
result=ModelResponse(model=model, usage=usage),
|
||||
litellm_model_name=model,
|
||||
)
|
||||
if cost is None:
|
||||
return None
|
||||
breakdown: Final = logging_obj.cost_breakdown
|
||||
return None if breakdown is None else (cost, breakdown)
|
||||
|
||||
def _session_cost(
|
||||
self,
|
||||
websocket_messages: Sequence[object],
|
||||
model: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> float | None:
|
||||
"""Price each turn on its own tokens and grounding, so two grounded turns pay the query fee twice.
|
||||
|
||||
The fixed cost margin is a flat per-request fee, so the session's single spend row carries it once
|
||||
rather than once per turn.
|
||||
"""
|
||||
turn_costs: Final = tuple(self._turn_cost(turn, model, logging_obj) for turn in _turns(websocket_messages))
|
||||
priced: Final = tuple(turn_cost for turn_cost in turn_costs if turn_cost is not None)
|
||||
if not priced or len(priced) != len(turn_costs):
|
||||
return None
|
||||
breakdowns: Final = tuple(breakdown for _, breakdown in priced)
|
||||
first: Final = breakdowns[0]
|
||||
fixed_margin: Final = first.get("margin_fixed_amount") or 0.0
|
||||
duplicated_fixed_margin: Final = fixed_margin * (len(priced) - 1)
|
||||
total_cost: Final = sum(cost for cost, _ in priced) - duplicated_fixed_margin
|
||||
summed_margin_total: Final = _summed(breakdowns, "margin_total_amount")
|
||||
margin_total_amount: Final = (
|
||||
None if summed_margin_total is None else summed_margin_total - duplicated_fixed_margin
|
||||
)
|
||||
logging_obj.set_cost_breakdown(
|
||||
input_cost=_summed(breakdowns, "input_cost") or 0.0,
|
||||
output_cost=_summed(breakdowns, "output_cost") or 0.0,
|
||||
total_cost=total_cost,
|
||||
cost_for_built_in_tools_cost_usd_dollar=_summed(breakdowns, "tool_usage_cost") or 0.0,
|
||||
original_cost=_summed(breakdowns, "original_cost"),
|
||||
discount_percent=first.get("discount_percent"),
|
||||
discount_amount=_summed(breakdowns, "discount_amount"),
|
||||
margin_percent=first.get("margin_percent"),
|
||||
margin_fixed_amount=first.get("margin_fixed_amount"),
|
||||
margin_total_amount=margin_total_amount,
|
||||
cache_read_cost=_summed(breakdowns, "cache_read_cost"),
|
||||
cache_creation_cost=_summed(breakdowns, "cache_creation_cost"),
|
||||
reasoning_cost=_summed(breakdowns, "reasoning_cost"),
|
||||
service_tier=first.get("service_tier"),
|
||||
data_residency=first.get("data_residency"),
|
||||
vertex_location=first.get("vertex_location"),
|
||||
)
|
||||
return total_cost
|
||||
|
||||
def vertex_ai_live_passthrough_handler(
|
||||
self,
|
||||
websocket_messages: list[dict],
|
||||
logging_obj,
|
||||
websocket_messages: Sequence[object],
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
url_route: str,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
|
|
@ -300,34 +387,25 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
|||
"""
|
||||
try:
|
||||
# Extract model from request body or kwargs
|
||||
model: Final = kwargs.get("model", "gemini-2.0-flash-live-preview-04-09")
|
||||
requested_model: Final = kwargs.get("model")
|
||||
model: Final = (
|
||||
requested_model if isinstance(requested_model, str) else "gemini-2.0-flash-live-preview-04-09"
|
||||
)
|
||||
custom_llm_provider: Final = kwargs.get("custom_llm_provider", "vertex_ai")
|
||||
verbose_proxy_logger.debug(
|
||||
"Vertex AI Live API model: %s, custom_llm_provider: %s", model, custom_llm_provider
|
||||
)
|
||||
|
||||
# Extract usage metadata from WebSocket messages
|
||||
usage_metadata: Final = self._extract_usage_metadata_from_websocket_messages(websocket_messages)
|
||||
usage: Final = self._session_usage(websocket_messages, model)
|
||||
|
||||
if not usage_metadata:
|
||||
if usage is None:
|
||||
verbose_proxy_logger.warning("No usage metadata found in Vertex AI Live API WebSocket messages")
|
||||
return {
|
||||
"result": None,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
|
||||
# Calculate cost using Live API specific pricing
|
||||
response_cost: Final = self._calculate_live_api_cost(
|
||||
model=model,
|
||||
usage_metadata=usage_metadata,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
# Create Usage object for standard LiteLLM logging
|
||||
usage: Final = self._create_usage_object_from_metadata(
|
||||
usage_metadata=usage_metadata,
|
||||
model=model,
|
||||
)
|
||||
response_cost: Final = self._session_cost(websocket_messages, model, logging_obj)
|
||||
|
||||
# Create a mock ModelResponse for standard logging
|
||||
litellm_model_response: Final = ModelResponse(
|
||||
|
|
@ -338,9 +416,9 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
|||
usage=usage,
|
||||
choices=[],
|
||||
)
|
||||
if response_cost is not None:
|
||||
litellm_model_response._hidden_params["response_cost"] = response_cost # pyright: ignore[reportPrivateUsage] # the logger reads the cost off the response's hidden params; the constructor's hidden_params kwarg is reset by pydantic
|
||||
|
||||
# Update kwargs with cost information
|
||||
kwargs["response_cost"] = response_cost
|
||||
kwargs["model"] = model
|
||||
kwargs["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
|
|
@ -348,12 +426,15 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
|
|||
import re
|
||||
|
||||
allowed_pattern: Final = re.compile(r"^[A-Za-z0-9._\-:]+$")
|
||||
safe_model: Final = model if isinstance(model, str) and allowed_pattern.match(model) else "[REDACTED]"
|
||||
safe_model: Final = model if allowed_pattern.match(model) else "[REDACTED]"
|
||||
verbose_proxy_logger.debug(
|
||||
f"Vertex AI Live API passthrough cost tracking - "
|
||||
f"Model: {safe_model}, Cost: ${response_cost:.6f}, "
|
||||
f"Prompt tokens: {usage.prompt_tokens}, "
|
||||
f"Completion tokens: {usage.completion_tokens}"
|
||||
"Vertex AI Live API passthrough cost tracking - Model: %s, "
|
||||
"Prompt tokens: %s %s, Completion tokens: %s %s",
|
||||
safe_model,
|
||||
usage.prompt_tokens,
|
||||
usage.prompt_tokens_details,
|
||||
usage.completion_tokens,
|
||||
usage.completion_tokens_details,
|
||||
)
|
||||
|
||||
return {
|
||||
|
|
|
|||
|
|
@ -2090,6 +2090,22 @@ def _rewrite_vertex_live_setup_model(text_data: str, setup_model_rewriter: Calla
|
|||
return json.dumps({**message, "setup": {**setup, "model": rewritten_model}}) # mutable-ok: one-shot json payload
|
||||
|
||||
|
||||
def _resolved_vertex_live_setup(
|
||||
setup_data: Mapping[str, object], setup_model_rewriter: Callable[[str], str] | None
|
||||
) -> Mapping[str, object]:
|
||||
"""
|
||||
Give the model extractor the same fully qualified path the upstream will receive.
|
||||
|
||||
Clients may name a bare gateway alias, which the rewriter turns into a ``projects/...`` path before
|
||||
it reaches Vertex. The extractor only reads a path containing ``/models/``, so running it on the raw
|
||||
frame logs the session as ``unknown`` at no cost, which is precisely the supported client form
|
||||
"""
|
||||
setup_model: Final = setup_data.get("model")
|
||||
if setup_model_rewriter is None or not isinstance(setup_model, str):
|
||||
return setup_data
|
||||
return {**setup_data, "model": setup_model_rewriter(setup_model)}
|
||||
|
||||
|
||||
def _truncated_close_reason(reason: str) -> str:
|
||||
"""
|
||||
Fit a close reason inside the byte budget a WebSocket close frame allows, without splitting a character
|
||||
|
|
@ -2314,7 +2330,9 @@ async def websocket_passthrough_request(
|
|||
setup_data,
|
||||
)
|
||||
if isinstance(setup_data, dict) and "model" in setup_data:
|
||||
extracted_model = _extract_model_from_vertex_ai_setup(setup_data)
|
||||
extracted_model = _extract_model_from_vertex_ai_setup(
|
||||
_resolved_vertex_live_setup(setup_data, setup_model_rewriter)
|
||||
)
|
||||
if extracted_model:
|
||||
kwargs["model"] = extracted_model
|
||||
kwargs["custom_llm_provider"] = "vertex_ai-language-models"
|
||||
|
|
|
|||
|
|
@ -58,15 +58,6 @@ class UndeliverableStreamRewrite(Exception):
|
|||
self.guardrail_name: Final = guardrail_name
|
||||
|
||||
|
||||
class UnappliableRequestRewrite(Exception):
|
||||
def __init__(self, guardrail_name: str) -> None:
|
||||
super().__init__(
|
||||
f"Guardrail '{guardrail_name}' rewrote the request in a way this endpoint cannot apply, "
|
||||
"so the request was rejected rather than sent unrewritten"
|
||||
)
|
||||
self.guardrail_name: Final = guardrail_name
|
||||
|
||||
|
||||
def _tool_call_shape(tool_call: object) -> tuple[object, object]:
|
||||
plain: Final = tool_call.model_dump() if isinstance(tool_call, BaseModel) else tool_call
|
||||
function: Final = plain.get("function") if isinstance(plain, Mapping) else None
|
||||
|
|
|
|||
|
|
@ -476,9 +476,10 @@ from litellm.proxy.hooks.prompt_injection_detection import (
|
|||
from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_spend_event
|
||||
from litellm.proxy.image_endpoints.endpoints import router as image_router
|
||||
from litellm.proxy.list_api.common import (
|
||||
PROBLEM_TYPE_BASE,
|
||||
ManagementProblem,
|
||||
ValidationErrorDetail,
|
||||
problem_response,
|
||||
request_validation_problem,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
from litellm.proxy.logging_endpoints.callback_logs_endpoints import (
|
||||
|
|
@ -601,7 +602,6 @@ from litellm.proxy.spend_tracking.spend_event_producer import (
|
|||
SpendEventProducer,
|
||||
build_spend_event_producer,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail
|
||||
|
||||
try:
|
||||
from litellm.proxy.enterprise_billing.billing_metrics import (
|
||||
|
|
@ -928,6 +928,7 @@ def cleanup_router_config_variables():
|
|||
user_custom_auth_path, \
|
||||
user_custom_key_generate, \
|
||||
user_custom_key_update, \
|
||||
user_custom_key_policy, \
|
||||
user_custom_sso, \
|
||||
user_custom_ui_sso_sign_in_handler, \
|
||||
use_background_health_checks, \
|
||||
|
|
@ -945,6 +946,7 @@ def cleanup_router_config_variables():
|
|||
user_custom_auth_path = None
|
||||
user_custom_key_generate = None
|
||||
user_custom_key_update = None
|
||||
user_custom_key_policy = None
|
||||
TEAM_METADATA_VALIDATOR_REGISTRY.set(None)
|
||||
TEAM_METADATA_SCHEMA_REGISTRY.set(())
|
||||
user_custom_sso = None
|
||||
|
|
@ -1787,27 +1789,13 @@ class _ExceptionRow(TypedDict, total=False):
|
|||
exception_counts: Mapping[str, int]
|
||||
|
||||
|
||||
class _ValidationErrorDetail(TypedDict):
|
||||
loc: tuple[int | str, ...]
|
||||
msg: str
|
||||
|
||||
|
||||
@app.exception_handler(RequestValidationError)
|
||||
async def otel_request_validation_exception_handler(request: Request, exc: RequestValidationError):
|
||||
if request.url.path.startswith(MANAGEMENT_V1_PREFIX):
|
||||
_close_dangling_otel_server_span(request, 400, exc=exc)
|
||||
validation_errors: Final[Sequence[_ValidationErrorDetail]] = exc.errors()
|
||||
return problem_response(
|
||||
ProblemDetail(
|
||||
type=f"{PROBLEM_TYPE_BASE}invalid-query-parameter",
|
||||
title="Invalid query parameter",
|
||||
status=400,
|
||||
detail="; ".join(
|
||||
f"{'.'.join(str(part) for part in error['loc'][1:])}: {error['msg']}" for error in validation_errors
|
||||
)
|
||||
or "The request query parameters are invalid.",
|
||||
)
|
||||
)
|
||||
validation_errors: Final[Sequence[ValidationErrorDetail]] = exc.errors()
|
||||
problem: Final = request_validation_problem(validation_errors)
|
||||
_close_dangling_otel_server_span(request, problem.status, exc=exc)
|
||||
return problem_response(problem)
|
||||
_close_dangling_otel_server_span(request, 422, exc=exc)
|
||||
return JSONResponse(
|
||||
status_code=422,
|
||||
|
|
@ -2369,6 +2357,7 @@ user_custom_key_generate = None
|
|||
_pkce_no_redis_warning_emitted: bool = False
|
||||
_cp_no_redis_warning_emitted: bool = False
|
||||
user_custom_key_update = None
|
||||
user_custom_key_policy = None
|
||||
user_custom_sso = None
|
||||
user_custom_ui_sso_sign_in_handler = None
|
||||
use_background_health_checks = None
|
||||
|
|
@ -4256,6 +4245,7 @@ _DB_OVERLAY_REMOTE_MODULE_STR_FIELDS: Final[dict[str, tuple[str, ...]]] = {
|
|||
"custom_auth",
|
||||
"custom_key_generate",
|
||||
"custom_key_update",
|
||||
"custom_key_policy",
|
||||
"custom_team_metadata_validate",
|
||||
"custom_sso",
|
||||
"custom_ui_sso_sign_in_handler",
|
||||
|
|
@ -5405,6 +5395,7 @@ class ProxyConfig:
|
|||
user_custom_auth_path, \
|
||||
user_custom_key_generate, \
|
||||
user_custom_key_update, \
|
||||
user_custom_key_policy, \
|
||||
user_custom_sso, \
|
||||
user_custom_ui_sso_sign_in_handler, \
|
||||
use_background_health_checks, \
|
||||
|
|
@ -5942,6 +5933,10 @@ class ProxyConfig:
|
|||
if custom_key_update is not None:
|
||||
user_custom_key_update = get_instance_fn(value=custom_key_update, config_file_path=config_file_path)
|
||||
|
||||
custom_key_policy: Final = general_settings.get("custom_key_policy", None)
|
||||
if custom_key_policy is not None:
|
||||
user_custom_key_policy = get_instance_fn(value=custom_key_policy, config_file_path=config_file_path)
|
||||
|
||||
custom_team_metadata_validate: Final = general_settings.get("custom_team_metadata_validate", None)
|
||||
TEAM_METADATA_VALIDATOR_REGISTRY.set(
|
||||
get_instance_fn(value=custom_team_metadata_validate, config_file_path=config_file_path)
|
||||
|
|
@ -9546,6 +9541,7 @@ class ProxyStartupEvent:
|
|||
gate the first duration window.
|
||||
"""
|
||||
await generate_key_helper_fn(
|
||||
llm_router=llm_router,
|
||||
request_type="user",
|
||||
table_name="user",
|
||||
user_id=LITELLM_PROXY_BUDGET_NAME,
|
||||
|
|
@ -16290,6 +16286,7 @@ async def _generate_onboarding_ui_session_token(user_obj: _UserTableRow) -> str:
|
|||
global master_key, general_settings
|
||||
|
||||
response: Final = await generate_key_helper_fn(
|
||||
llm_router=llm_router,
|
||||
request_type="key",
|
||||
**{
|
||||
"user_role": user_obj.user_role,
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
from litellm.types.utils import CallTypes, LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
from ..litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from ..litellm_core_utils.get_litellm_params import get_litellm_params
|
||||
from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from ..llms.azure.common_utils import get_azure_ad_token
|
||||
|
|
@ -54,6 +55,17 @@ xai_realtime: Final = XAIRealtime()
|
|||
vertex_llm_base: Final = VertexBase()
|
||||
base_llm_http_handler = BaseLLMHTTPHandler()
|
||||
_EMPTY_MODEL_PARAMS: Final[Mapping[str, Any]] = MappingProxyType({})
|
||||
_EMPTY_AUTH_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _model_params_with_stored_credentials(model_params: Mapping[str, Any]) -> Mapping[str, Any]:
|
||||
credential_name: Final = model_params.get("litellm_credential_name")
|
||||
credential_values: Final = (
|
||||
CredentialAccessor.get_credential_values(credential_name)
|
||||
if isinstance(credential_name, str)
|
||||
else _EMPTY_MODEL_PARAMS
|
||||
)
|
||||
return MappingProxyType({**credential_values, **model_params})
|
||||
|
||||
|
||||
def _with_resolved_session_model(session: dict[str, object], model_name: str) -> dict[str, object]:
|
||||
|
|
@ -591,13 +603,15 @@ def _azure_realtime_health_protocol(
|
|||
|
||||
def _realtime_health_check_auth_headers(
|
||||
custom_llm_provider: str, api_key: str | None, model_params: Mapping[str, Any]
|
||||
) -> Mapping[str, str | None]:
|
||||
if custom_llm_provider != "azure":
|
||||
return MappingProxyType({"api-key": api_key})
|
||||
return azure_realtime.get_auth_headers(
|
||||
api_key=api_key,
|
||||
azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))),
|
||||
)
|
||||
) -> Mapping[str, str]:
|
||||
if custom_llm_provider == "azure":
|
||||
return azure_realtime.get_auth_headers(
|
||||
api_key=api_key,
|
||||
azure_ad_token=(None if api_key else get_azure_ad_token(GenericLiteLLMParams(**model_params))),
|
||||
)
|
||||
if api_key is None:
|
||||
return _EMPTY_AUTH_HEADERS
|
||||
return MappingProxyType({"Authorization": f"Bearer {api_key}"})
|
||||
|
||||
|
||||
async def _realtime_health_check(
|
||||
|
|
@ -629,34 +643,46 @@ async def _realtime_health_check(
|
|||
"""
|
||||
import websockets
|
||||
|
||||
resolved_params: Final = _model_params_with_stored_credentials(model_params or _EMPTY_MODEL_PARAMS)
|
||||
resolved_api_key: Final = cast( # cast-ok: provider parameters expose optional string credentials
|
||||
str | None, api_key or resolved_params.get("api_key")
|
||||
)
|
||||
resolved_api_base: Final = cast( # cast-ok: provider parameters expose optional string endpoints
|
||||
str | None, api_base or resolved_params.get("api_base")
|
||||
)
|
||||
resolved_api_version: Final = cast( # cast-ok: provider parameters expose optional string versions
|
||||
str | None, api_version or resolved_params.get("api_version")
|
||||
)
|
||||
url: str | None = None
|
||||
auth_headers: Final = _realtime_health_check_auth_headers(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
api_key=api_key,
|
||||
model_params=model_params or _EMPTY_MODEL_PARAMS,
|
||||
api_key=resolved_api_key,
|
||||
model_params=resolved_params,
|
||||
)
|
||||
if custom_llm_provider == "azure":
|
||||
resolved_protocol, azure_query_params = _azure_realtime_health_protocol(
|
||||
model=model,
|
||||
realtime_protocol=realtime_protocol,
|
||||
model_params=model_params or _EMPTY_MODEL_PARAMS,
|
||||
model_params=resolved_params,
|
||||
)
|
||||
url = azure_realtime._construct_url(
|
||||
api_base=api_base or "",
|
||||
api_base=resolved_api_base or "",
|
||||
model=model,
|
||||
api_version=api_version or "2024-10-01-preview",
|
||||
api_version=resolved_api_version or "2024-10-01-preview",
|
||||
realtime_protocol=resolved_protocol,
|
||||
query_params=azure_query_params,
|
||||
)
|
||||
elif custom_llm_provider == "openai":
|
||||
url = openai_realtime._construct_url(
|
||||
api_base=api_base or "https://api.openai.com/",
|
||||
api_base=resolved_api_base or "https://api.openai.com/",
|
||||
query_params={"model": model},
|
||||
)
|
||||
elif custom_llm_provider == "xai":
|
||||
url = xai_realtime._construct_url(api_base=api_base or "https://api.x.ai/v1", query_params={"model": model})
|
||||
url = xai_realtime._construct_url(
|
||||
api_base=resolved_api_base or "https://api.x.ai/v1", query_params={"model": model}
|
||||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
vertex_model_params: Final = model_params or {}
|
||||
vertex_model_params: Final = dict(resolved_params)
|
||||
resolved_location: Final = vertex_llm_base.get_vertex_region(
|
||||
vertex_region=VertexBase.safe_get_vertex_ai_location(vertex_model_params),
|
||||
model=model,
|
||||
|
|
@ -675,19 +701,19 @@ async def _realtime_health_check(
|
|||
project=resolved_project,
|
||||
location=resolved_location,
|
||||
)
|
||||
url = vertex_realtime_config.get_complete_url(api_base=api_base, model=model)
|
||||
ssl_context = get_shared_realtime_ssl_context()
|
||||
url = vertex_realtime_config.get_complete_url(api_base=resolved_api_base, model=model)
|
||||
vertex_ssl_context: Final = get_shared_realtime_ssl_context()
|
||||
headers: Final = vertex_realtime_config.validate_environment(headers={}, model=model, api_key=None)
|
||||
async with websockets.connect(
|
||||
url,
|
||||
additional_headers=headers,
|
||||
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
||||
ssl=ssl_context,
|
||||
ssl=vertex_ssl_context,
|
||||
):
|
||||
return True
|
||||
else:
|
||||
raise ValueError(f"Unsupported model: {model}")
|
||||
ssl_context = get_shared_realtime_ssl_context()
|
||||
ssl_context: Final = get_shared_realtime_ssl_context()
|
||||
async with websockets.connect(
|
||||
url,
|
||||
additional_headers=auth_headers,
|
||||
|
|
|
|||
|
|
@ -117,3 +117,13 @@ class BaseRepository(ABC, Generic[T]):
|
|||
"""Check if a record exists."""
|
||||
record: Final = await self.table.find_unique(where={id_field: id_value})
|
||||
return record is not None
|
||||
|
||||
|
||||
def is_unique_violation(exc: BaseException) -> bool:
|
||||
try:
|
||||
from prisma.errors import UniqueViolationError
|
||||
except ImportError:
|
||||
return "P2002" in str(exc) or "unique constraint" in str(exc).lower()
|
||||
if isinstance(exc, UniqueViolationError):
|
||||
return True
|
||||
return getattr(exc, "code", None) == "P2002"
|
||||
|
|
|
|||
65
litellm/responses/additional_tools.py
Normal file
65
litellm/responses/additional_tools.py
Normal file
|
|
@ -0,0 +1,65 @@
|
|||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, cast # noqa: TID251 # validating the openai tool union strips vendor keys from raw tools
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.llms.openai import ALL_RESPONSES_API_TOOL_PARAMS, ResponseInputParam
|
||||
|
||||
ADDITIONAL_TOOLS_INPUT_ITEM_TYPE: Final = "additional_tools"
|
||||
|
||||
|
||||
class _InputItemType(BaseModel):
|
||||
type: str = ""
|
||||
|
||||
|
||||
class _AdditionalToolsItem(BaseModel):
|
||||
tools: tuple[dict[str, object], ...] = ()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class HoistedAdditionalTools:
|
||||
input: str | ResponseInputParam
|
||||
tools: tuple[ALL_RESPONSES_API_TOOL_PARAMS, ...]
|
||||
hoisted: tuple[ALL_RESPONSES_API_TOOL_PARAMS, ...]
|
||||
|
||||
|
||||
def _is_additional_tools_item(item: object) -> bool:
|
||||
try:
|
||||
return _InputItemType.model_validate(item).type == ADDITIONAL_TOOLS_INPUT_ITEM_TYPE
|
||||
except ValidationError:
|
||||
return False
|
||||
|
||||
|
||||
def _tools_of_item(item: object) -> tuple[ALL_RESPONSES_API_TOOL_PARAMS, ...]:
|
||||
try:
|
||||
parsed: Final = _AdditionalToolsItem.model_validate(item)
|
||||
except ValidationError:
|
||||
return ()
|
||||
return tuple(
|
||||
cast(
|
||||
"ALL_RESPONSES_API_TOOL_PARAMS", tool
|
||||
) # cast-ok: nested tools carry the same raw tool JSON as top-level tools
|
||||
for tool in parsed.tools
|
||||
)
|
||||
|
||||
|
||||
def hoist_additional_tools(
|
||||
input: str | ResponseInputParam,
|
||||
tools: Sequence[ALL_RESPONSES_API_TOOL_PARAMS] | None,
|
||||
) -> HoistedAdditionalTools:
|
||||
existing: Final = tuple(tools or ())
|
||||
if isinstance(input, str):
|
||||
return HoistedAdditionalTools(input=input, tools=existing, hoisted=())
|
||||
items: Final = tuple(item for item in input if _is_additional_tools_item(item))
|
||||
if not items:
|
||||
return HoistedAdditionalTools(input=input, tools=existing, hoisted=())
|
||||
hoisted: Final = tuple(tool for item in items for tool in _tools_of_item(item))
|
||||
verbose_logger.debug(
|
||||
"Responses API: hoisting %d tool(s) out of %d 'additional_tools' input item(s) into the top-level tools param.",
|
||||
len(hoisted),
|
||||
len(items),
|
||||
)
|
||||
remaining_input: Final = [item for item in input if not _is_additional_tools_item(item)]
|
||||
return HoistedAdditionalTools(input=remaining_input, tools=(*existing, *hoisted), hoisted=hoisted)
|
||||
|
|
@ -39,15 +39,38 @@ def openai_shaped_tool_call_item_id(item_type: str, tool_id: str) -> str:
|
|||
return f"{prefix}_{tool_id}"
|
||||
|
||||
|
||||
class _ToolNameFields(BaseModel):
|
||||
type: str = ""
|
||||
name: str = ""
|
||||
tools: tuple[object, ...] = ()
|
||||
|
||||
|
||||
def _tool_name_fields_of(tool: object) -> _ToolNameFields | None:
|
||||
try:
|
||||
return _ToolNameFields.model_validate(tool)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _custom_tool_name_of(tool: object) -> str | None:
|
||||
parsed: Final = _tool_name_fields_of(tool)
|
||||
if parsed is None or parsed.type != "custom" or not parsed.name:
|
||||
return None
|
||||
return parsed.name
|
||||
|
||||
|
||||
def _nested_tools_of(tool: object) -> tuple[object, ...]:
|
||||
parsed: Final = _tool_name_fields_of(tool)
|
||||
if parsed is None or parsed.type != "namespace":
|
||||
return ()
|
||||
return parsed.tools
|
||||
|
||||
|
||||
def extract_custom_tool_names(tools: Sequence[object] | None) -> set[str]:
|
||||
"""Extract names of tools originally defined as ``type: "custom"``."""
|
||||
if not tools:
|
||||
return set()
|
||||
names: Final[set[str]] = set()
|
||||
for tool in tools:
|
||||
if isinstance(tool, dict) and tool.get("type") == "custom" and "name" in tool:
|
||||
names.add(tool["name"])
|
||||
return names
|
||||
"""Extract names of ``type: "custom"`` tools, at the top level or one level inside a ``namespace`` tool."""
|
||||
top_level: Final = tuple(tools or ())
|
||||
nested: Final = tuple(nested_tool for tool in top_level for nested_tool in _nested_tools_of(tool))
|
||||
return {name for tool in (*top_level, *nested) if (name := _custom_tool_name_of(tool)) is not None}
|
||||
|
||||
|
||||
def is_custom_tool_call(tool_name: str, custom_tool_names: set[str]) -> bool:
|
||||
|
|
@ -143,7 +166,7 @@ def validated_allowed_callers(value: object) -> list[str] | None:
|
|||
raise ValueError("allowed_callers must be a list of strings") from exc
|
||||
|
||||
|
||||
def _grammar_suffix(fmt: object) -> str:
|
||||
def custom_tool_grammar_suffix(fmt: object) -> str:
|
||||
try:
|
||||
parsed: Final = _CustomToolFormat.model_validate(fmt)
|
||||
except ValidationError:
|
||||
|
|
@ -167,7 +190,9 @@ def convert_custom_tool_to_function_tool(tool: Mapping[str, object]) -> ChatComp
|
|||
raw_name: Final = tool.get("name")
|
||||
name: Final = raw_name if isinstance(raw_name, str) else ""
|
||||
raw_description: Final = tool.get("description")
|
||||
description = (raw_description if isinstance(raw_description, str) else "") + _grammar_suffix(tool.get("format"))
|
||||
description: Final = (raw_description if isinstance(raw_description, str) else "") + custom_tool_grammar_suffix(
|
||||
tool.get("format")
|
||||
)
|
||||
allowed_callers: Final = validated_allowed_callers(tool.get("allowed_callers"))
|
||||
function_chunk: Final = ChatCompletionToolParamFunctionChunk(
|
||||
name=name,
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from collections.abc import Coroutine, Mapping
|
|||
from typing import Final
|
||||
|
||||
import litellm
|
||||
from litellm.responses.additional_tools import hoist_additional_tools
|
||||
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
|
||||
LiteLLMCompletionStreamingIterator,
|
||||
)
|
||||
|
|
@ -37,11 +38,16 @@ class LiteLLMCompletionTransformationHandler:
|
|||
| BaseResponsesAPIStreamingIterator
|
||||
| Coroutine[object, object, ResponsesAPIResponse | BaseResponsesAPIStreamingIterator]
|
||||
):
|
||||
hoisted: Final = hoist_additional_tools(input, responses_api_request.get("tools"))
|
||||
bridged_input: Final = hoisted.input
|
||||
bridged_request: Final[ResponsesAPIOptionalRequestParams] = (
|
||||
{**responses_api_request, "tools": list(hoisted.tools)} if hoisted.hoisted else responses_api_request
|
||||
)
|
||||
litellm_completion_request: Final[dict] = (
|
||||
LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request(
|
||||
model=model,
|
||||
input=input,
|
||||
responses_api_request=responses_api_request,
|
||||
input=bridged_input,
|
||||
responses_api_request=bridged_request,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
stream=stream,
|
||||
extra_headers=extra_headers,
|
||||
|
|
@ -52,8 +58,8 @@ class LiteLLMCompletionTransformationHandler:
|
|||
if _is_async:
|
||||
return self.async_response_api_handler(
|
||||
litellm_completion_request=litellm_completion_request,
|
||||
request_input=input,
|
||||
responses_api_request=responses_api_request,
|
||||
request_input=bridged_input,
|
||||
responses_api_request=bridged_request,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -70,8 +76,8 @@ class LiteLLMCompletionTransformationHandler:
|
|||
responses_api_response: Final[ResponsesAPIResponse] = (
|
||||
LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
|
||||
chat_completion_response=litellm_completion_response,
|
||||
request_input=input,
|
||||
responses_api_request=responses_api_request,
|
||||
request_input=bridged_input,
|
||||
responses_api_request=bridged_request,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -81,8 +87,8 @@ class LiteLLMCompletionTransformationHandler:
|
|||
return LiteLLMCompletionStreamingIterator(
|
||||
model=model,
|
||||
litellm_custom_stream_wrapper=litellm_completion_response,
|
||||
request_input=input,
|
||||
responses_api_request=responses_api_request,
|
||||
request_input=bridged_input,
|
||||
responses_api_request=bridged_request,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_metadata=kwargs.get("litellm_metadata", {}),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from litellm.main import stream_chunk_builder
|
|||
from litellm.responses.litellm_completion_transformation.custom_tools import (
|
||||
build_tool_call_item_kwargs,
|
||||
extract_custom_tool_names,
|
||||
is_custom_tool_call,
|
||||
serialize_tool_call_arguments,
|
||||
)
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
|
|
@ -166,6 +167,14 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
return tool_name, namespace
|
||||
return fn_name, None
|
||||
|
||||
def _tool_call_item_kwargs(self, call_id: str, fn_name: str, arguments: str, status: str) -> dict[str, str]:
|
||||
item_kwargs: Final = build_tool_call_item_kwargs(call_id, fn_name, arguments, status, self._custom_tool_names)
|
||||
if is_custom_tool_call(fn_name, self._custom_tool_names):
|
||||
return item_kwargs
|
||||
tool_name, tool_namespace = self._responses_namespace_tool_call_fields(fn_name)
|
||||
namespace_kwargs: Final = {"namespace": tool_namespace} if tool_namespace else {}
|
||||
return {**item_kwargs, "name": tool_name, **namespace_kwargs}
|
||||
|
||||
def _is_reasoning_end(self, chunk):
|
||||
delta: Final = chunk.choices[0].delta
|
||||
|
||||
|
|
@ -244,17 +253,13 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
else:
|
||||
fn_name = str(getattr(fn, "name", "") or "")
|
||||
fn_args_delta = serialize_tool_call_arguments(getattr(fn, "arguments", ""))
|
||||
tool_name, tool_namespace = self._responses_namespace_tool_call_fields(fn_name)
|
||||
output_index = self._get_or_assign_tool_output_index(call_id)
|
||||
|
||||
if call_id not in self._tool_args_by_call_id:
|
||||
self._tool_args_by_call_id[call_id] = ""
|
||||
self._sequence_number += 1
|
||||
names = self._custom_tool_names
|
||||
item_kwargs = build_tool_call_item_kwargs(call_id, tool_name, "", "in_progress", names)
|
||||
item_kwargs = self._tool_call_item_kwargs(call_id, fn_name, "", "in_progress")
|
||||
self._tool_item_id_by_call_id[call_id] = item_kwargs["id"]
|
||||
if tool_namespace:
|
||||
item_kwargs["namespace"] = tool_namespace
|
||||
event = OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=output_index,
|
||||
|
|
@ -315,7 +320,6 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
else:
|
||||
fn_name = str(getattr(fn, "name", "") or "")
|
||||
fn_args = serialize_tool_call_arguments(getattr(fn, "arguments", ""))
|
||||
tool_name, tool_namespace = self._responses_namespace_tool_call_fields(fn_name)
|
||||
web_search_call = self._web_search_calls.get(call_id)
|
||||
if web_search_call is not None:
|
||||
if call_id not in self._queued_web_search_call_ids:
|
||||
|
|
@ -330,11 +334,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
if is_new_tool_call:
|
||||
self._tool_args_by_call_id[call_id] = ""
|
||||
self._sequence_number += 1
|
||||
names = self._custom_tool_names
|
||||
item_kwargs = build_tool_call_item_kwargs(call_id, tool_name, "", "in_progress", names)
|
||||
item_kwargs = self._tool_call_item_kwargs(call_id, fn_name, "", "in_progress")
|
||||
self._tool_item_id_by_call_id[call_id] = item_kwargs["id"]
|
||||
if tool_namespace:
|
||||
item_kwargs["namespace"] = tool_namespace
|
||||
event = OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=output_index,
|
||||
|
|
@ -376,11 +377,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
self._pending_tool_events.append(done_event)
|
||||
|
||||
self._sequence_number += 1
|
||||
names = self._custom_tool_names
|
||||
item_kwargs = build_tool_call_item_kwargs(call_id, tool_name, final_args, "completed", names)
|
||||
item_kwargs = self._tool_call_item_kwargs(call_id, fn_name, final_args, "completed")
|
||||
item_kwargs["id"] = self._tool_item_id_by_call_id.setdefault(call_id, item_kwargs["id"])
|
||||
if tool_namespace:
|
||||
item_kwargs["namespace"] = tool_namespace
|
||||
item_done_event = OutputItemDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
output_index=output_index,
|
||||
|
|
|
|||
|
|
@ -110,6 +110,7 @@ NamespaceTool: TypeAlias = Mapping[str, object]
|
|||
ResponseTools: TypeAlias = Sequence[Mapping[str, object]] | None
|
||||
ChatToolParam: TypeAlias = ChatCompletionToolParam | OpenAIMcpServerTool
|
||||
NAMESPACE_DESCRIPTION_SEPARATOR: Final = "\n\n"
|
||||
NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS: Final = frozenset({"function", "custom"})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -1891,9 +1892,21 @@ class LiteLLMCompletionResponsesConfig:
|
|||
namespace_tool: NamespaceTool,
|
||||
nested: bool,
|
||||
) -> ChatCompletionToolParam | None:
|
||||
if nested and namespace_tool.get("type") != "function":
|
||||
tool_type: Final = namespace_tool.get("type")
|
||||
if nested and tool_type not in NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS:
|
||||
return None
|
||||
|
||||
raw_description: Final = str(namespace_tool.get("description") or "")
|
||||
description: Final = (
|
||||
f"{namespace_description}{NAMESPACE_DESCRIPTION_SEPARATOR}{raw_description}"
|
||||
if nested and namespace_description and raw_description
|
||||
else namespace_description
|
||||
if nested and namespace_description
|
||||
else raw_description
|
||||
)
|
||||
if nested and tool_type == "custom":
|
||||
return convert_custom_tool_to_function_tool({**namespace_tool, "description": description})
|
||||
|
||||
raw_parameters: Final = namespace_tool.get("parameters")
|
||||
parameters: Final = (
|
||||
MappingProxyType(raw_parameters) if isinstance(raw_parameters, Mapping) else MappingProxyType({})
|
||||
|
|
@ -1902,14 +1915,6 @@ class LiteLLMCompletionResponsesConfig:
|
|||
parameters if parameters and "type" in parameters else MappingProxyType({**parameters, "type": "object"})
|
||||
)
|
||||
tool_name: Final = str(namespace_tool.get("name") or "")
|
||||
raw_description: Final = str(namespace_tool.get("description") or "")
|
||||
description: Final = (
|
||||
f"{namespace_description}{NAMESPACE_DESCRIPTION_SEPARATOR}{raw_description}"
|
||||
if nested and namespace_description and raw_description
|
||||
else namespace_description
|
||||
if nested and namespace_description
|
||||
else raw_description
|
||||
)
|
||||
chat_tool_name: Final = f"{namespace}__{tool_name}" if nested else tool_name
|
||||
function: Final = ChatCompletionToolParamFunctionChunk(
|
||||
name=chat_tool_name,
|
||||
|
|
@ -2826,6 +2831,22 @@ class LiteLLMCompletionResponsesConfig:
|
|||
if cache_write_tokens is not None
|
||||
else MappingProxyType({})
|
||||
)
|
||||
# The cost path reads the grounding counters off the input details, and a realtime
|
||||
# session's usage is rebuilt from its own response.done, so dropping them here bills
|
||||
# no per-query grounding fee at all.
|
||||
grounding_request_counts: Final[Mapping[str, int]] = MappingProxyType(
|
||||
{
|
||||
counter: count
|
||||
for counter, count in (
|
||||
("web_search_requests", getattr(prompt_details, "web_search_requests", None)),
|
||||
(
|
||||
"google_maps_grounding_requests",
|
||||
getattr(prompt_details, "google_maps_grounding_requests", None),
|
||||
),
|
||||
)
|
||||
if count is not None
|
||||
}
|
||||
)
|
||||
response_usage.input_tokens_details = InputTokensDetails(
|
||||
cached_tokens=prompt_details.cached_tokens if prompt_details.cached_tokens is not None else 0,
|
||||
text_tokens=prompt_details.text_tokens,
|
||||
|
|
@ -2834,6 +2855,7 @@ class LiteLLMCompletionResponsesConfig:
|
|||
cached_tokens_details if isinstance(cached_tokens_details, CachedTokensDetails) else None
|
||||
),
|
||||
**cache_write_extra,
|
||||
**grounding_request_counts,
|
||||
)
|
||||
|
||||
# Translate completion_tokens_details to output_tokens_details
|
||||
|
|
|
|||
|
|
@ -1183,6 +1183,10 @@ class ResponseAPILoggingUtils:
|
|||
response_api_usage.input_tokens_details, "cached_tokens_details", None
|
||||
),
|
||||
cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None),
|
||||
web_search_requests=getattr(response_api_usage.input_tokens_details, "web_search_requests", None),
|
||||
google_maps_grounding_requests=getattr(
|
||||
response_api_usage.input_tokens_details, "google_maps_grounding_requests", None
|
||||
),
|
||||
)
|
||||
completion_tokens_details: CompletionTokensDetailsWrapper | None = None
|
||||
output_tokens_details: Final[OutputTokensDetails | None] = getattr(
|
||||
|
|
|
|||
|
|
@ -4849,6 +4849,7 @@ class Router:
|
|||
model=model,
|
||||
messages=messages,
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
|
||||
data: Final = deployment["litellm_params"].copy()
|
||||
|
|
@ -5163,13 +5164,11 @@ class Router:
|
|||
return healthy_deployments[0]
|
||||
|
||||
# Use simple_shuffle for weighted selection
|
||||
return cast(
|
||||
GuardrailTypedDict,
|
||||
simple_shuffle(
|
||||
llm_router_instance=self,
|
||||
healthy_deployments=healthy_deployments,
|
||||
model=guardrail_name,
|
||||
),
|
||||
return simple_shuffle(
|
||||
resolve_model_alias=self._get_model_from_alias,
|
||||
healthy_deployments=healthy_deployments,
|
||||
model=guardrail_name,
|
||||
request_kwargs=None,
|
||||
)
|
||||
|
||||
async def _ageneric_api_call_with_fallbacks(self, model: str, original_function: Callable, **kwargs):
|
||||
|
|
@ -8378,7 +8377,8 @@ class Router:
|
|||
|
||||
def log_retry(self, kwargs: dict, e: Exception) -> dict:
|
||||
"""
|
||||
When a retry or fallback happens, record which model group, deployment and attempt just failed and why
|
||||
When a retry or fallback happens, record which model group, deployment and attempt just failed and why,
|
||||
and count it toward the request-wide num_retries_per_request cap
|
||||
"""
|
||||
from litellm.types.router import RetryAttemptRecord
|
||||
|
||||
|
|
@ -8402,7 +8402,10 @@ class Router:
|
|||
else ()
|
||||
)
|
||||
breadcrumbs: Final = (*kept_breadcrumbs, attempt_record)
|
||||
earlier: Final = request_metadata.get("request_retry_count")
|
||||
request_retry_count: Final = (earlier if type(earlier) is int and 0 <= earlier else 0) + 1
|
||||
kwargs[_metadata_var]["previous_models"] = breadcrumbs # rebind-ok: the logging object already holds this dict
|
||||
kwargs[_metadata_var]["request_retry_count"] = request_retry_count # rebind-ok: same dict, read by the cap
|
||||
return kwargs
|
||||
|
||||
def _update_usage(self, deployment_id: str, parent_otel_span: Span | None) -> int:
|
||||
|
|
@ -13045,9 +13048,10 @@ class Router:
|
|||
start_time: Final = time.time()
|
||||
if strategy == "simple-shuffle":
|
||||
return simple_shuffle(
|
||||
llm_router_instance=self,
|
||||
resolve_model_alias=self._get_model_from_alias,
|
||||
healthy_deployments=healthy_deployments,
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
deployment: Final = await self._select_deployment_async(
|
||||
strategy=strategy,
|
||||
|
|
@ -13190,9 +13194,10 @@ class Router:
|
|||
start_time: Final = time.perf_counter()
|
||||
if strategy == "simple-shuffle":
|
||||
return simple_shuffle(
|
||||
llm_router_instance=self,
|
||||
resolve_model_alias=self._get_model_from_alias,
|
||||
healthy_deployments=pass_through_deployments,
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
deployment: Final = await self._select_deployment_async(
|
||||
strategy=strategy,
|
||||
|
|
@ -13888,9 +13893,10 @@ class Router:
|
|||
# if users pass rpm or tpm, we do a random weighted pick - based on rpm/tpm
|
||||
############## Check 'weight' param set for weighted pick #################
|
||||
return simple_shuffle(
|
||||
llm_router_instance=self,
|
||||
resolve_model_alias=self._get_model_from_alias,
|
||||
healthy_deployments=healthy_deployments,
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
deployment: Final = self._select_deployment_sync(
|
||||
strategy=strategy,
|
||||
|
|
@ -13958,6 +13964,7 @@ class Router:
|
|||
messages=messages,
|
||||
input=input,
|
||||
specific_deployment=specific_deployment,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
|
||||
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
|
||||
|
|
@ -14040,9 +14047,10 @@ class Router:
|
|||
# 6. Apply load balancing strategy
|
||||
if strategy == "simple-shuffle":
|
||||
return simple_shuffle(
|
||||
llm_router_instance=self,
|
||||
resolve_model_alias=self._get_model_from_alias,
|
||||
healthy_deployments=pass_through_deployments,
|
||||
model=model,
|
||||
request_kwargs=request_kwargs,
|
||||
)
|
||||
deployment: Final = self._select_deployment_sync(
|
||||
strategy=strategy,
|
||||
|
|
|
|||
|
|
@ -16,6 +16,8 @@ Inspired by ClawRouter: https://github.com/BlockRunAI/ClawRouter
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
|
|
@ -28,6 +30,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, cast
|
|||
from pydantic import BaseModel, TypeAdapter, ValidationError, create_model
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.affinity_cache import claim_affinity_pin
|
||||
from litellm.constants import (
|
||||
EMPTY_MAPPING,
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY,
|
||||
|
|
@ -55,6 +58,7 @@ from litellm.router_strategy.complexity_router.tier_predictor import (
|
|||
TierSuccessPredictor,
|
||||
resolve_tier_artifact,
|
||||
)
|
||||
from litellm.router_utils.pre_call_checks.deployment_affinity_check import DeploymentAffinityCheck
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionImageObject,
|
||||
|
|
@ -1119,10 +1123,10 @@ class _ContextWindowPlacement(NamedTuple):
|
|||
|
||||
class _SessionAffinityPin(NamedTuple):
|
||||
model: str
|
||||
tier: ComplexityTier | None
|
||||
tier: ComplexityTier | str | None
|
||||
|
||||
|
||||
def _parse_session_affinity_pin(value: object) -> _SessionAffinityPin | None:
|
||||
def _parse_session_affinity_pin(value: object, active_tiers: tuple[str, ...]) -> _SessionAffinityPin | None:
|
||||
if isinstance(value, str):
|
||||
return _SessionAffinityPin(model=value, tier=None)
|
||||
parts: Final[tuple[object, object] | None] = (
|
||||
|
|
@ -1137,8 +1141,11 @@ def _parse_session_affinity_pin(value: object) -> _SessionAffinityPin | None:
|
|||
model, tier_value = parts
|
||||
if not isinstance(model, str):
|
||||
return None
|
||||
tier: Final = ComplexityTier(tier_value) if isinstance(tier_value, str) else None
|
||||
return _SessionAffinityPin(model=model, tier=tier)
|
||||
if tier_value is None:
|
||||
return _SessionAffinityPin(model=model, tier=None)
|
||||
if not isinstance(tier_value, str) or tier_value not in active_tiers:
|
||||
return None
|
||||
return _SessionAffinityPin(model=model, tier=_built_in_tier_or_none(tier_value) or tier_value)
|
||||
|
||||
|
||||
def _session_affinity_cache_value(model: str, tier: ComplexityTier | str | None) -> Mapping[str, str | None]:
|
||||
|
|
@ -1195,6 +1202,10 @@ class ComplexityRouter(CustomLogger):
|
|||
if default_model:
|
||||
self.config.default_model = default_model
|
||||
|
||||
self._tier_affinity_config = hashlib.sha256(
|
||||
self.config.model_dump_json(include=MappingProxyType({"tiers": True, "tier_model_configs": True})).encode()
|
||||
).hexdigest()
|
||||
|
||||
# Checked here rather than on the config model because the deployment's
|
||||
# complexity_router_default_model arrives outside complexity_router_config and is
|
||||
# applied just above, so a validator on the model would reject a deployment that
|
||||
|
|
@ -2259,6 +2270,51 @@ class ComplexityRouter(CustomLogger):
|
|||
def _tier_pools(self) -> dict[str, list[str]]:
|
||||
return {tier: (models if isinstance(models, list) else [models]) for tier, models in self.config.tiers.items()}
|
||||
|
||||
async def _pin_model_for_tier(
|
||||
self,
|
||||
tier: ComplexityTier | str,
|
||||
model: str,
|
||||
candidates: tuple[str, ...],
|
||||
request_kwargs: dict[str, object], # mutable-ok: adaptive feedback metadata must follow the selected model
|
||||
retained_pin: _SessionAffinityPin | None = None,
|
||||
) -> str:
|
||||
if not self._uses_deployment_pin or model not in candidates:
|
||||
return model
|
||||
retained_model: Final = (
|
||||
retained_pin.model
|
||||
if retained_pin is not None
|
||||
and retained_pin.tier is not None
|
||||
and _tier_name(retained_pin.tier) == _tier_name(tier)
|
||||
else None
|
||||
)
|
||||
if retained_model is not None and retained_model in candidates:
|
||||
self._restamp_adaptive_choice(request_kwargs, model, retained_model)
|
||||
return retained_model
|
||||
session_id: Final = self._get_session_id_from_request_kwargs(request_kwargs)
|
||||
if session_id is None:
|
||||
return model
|
||||
caller: Final = DeploymentAffinityCheck.get_user_key_from_request_kwargs(request_kwargs)
|
||||
identity: Final = (self.model_name, self._tier_affinity_config, caller, session_id, _tier_name(tier))
|
||||
cache_identity: Final = (
|
||||
(*identity, ("replay_fallback", retained_model)) if retained_model is not None else identity
|
||||
)
|
||||
cache_key: Final = (
|
||||
"complexity_router_tier_model_affinity:v1:"
|
||||
+ hashlib.sha256(json.dumps(cache_identity).encode()).hexdigest()
|
||||
)
|
||||
winner: Final = await claim_affinity_pin(
|
||||
self.litellm_router_instance.cache,
|
||||
cache_key,
|
||||
MappingProxyType({"model": model}),
|
||||
self.config.session_affinity_ttl_seconds,
|
||||
eligible_values=tuple(MappingProxyType({"model": candidate}) for candidate in candidates),
|
||||
)
|
||||
pinned: Final[object] = winner.get("model") if isinstance(winner, Mapping) else None
|
||||
if not isinstance(pinned, str) or pinned not in candidates:
|
||||
return model
|
||||
self._restamp_adaptive_choice(request_kwargs, model, pinned)
|
||||
return pinned
|
||||
|
||||
async def _pick_model_for_tier(
|
||||
self,
|
||||
tier: ComplexityTier | str,
|
||||
|
|
@ -2266,11 +2322,18 @@ class ComplexityRouter(CustomLogger):
|
|||
resolved_messages: list[dict[str, Any]] | None,
|
||||
request_kwargs: dict,
|
||||
allowed_models: tuple[str, ...] | None = None,
|
||||
retained_pin: _SessionAffinityPin | None = None,
|
||||
) -> str:
|
||||
if not self.config.plugins:
|
||||
if allowed_models is not None:
|
||||
return self._pick_from_tier_value(allowed_models, _tier_name(tier))
|
||||
return self.get_model_for_tier(tier)
|
||||
candidates: Final = (
|
||||
allowed_models if allowed_models is not None else tuple(self._tier_pools().get(_tier_name(tier), ()))
|
||||
)
|
||||
selected: Final = (
|
||||
self._pick_from_tier_value(allowed_models, _tier_name(tier))
|
||||
if allowed_models is not None
|
||||
else self.get_model_for_tier(tier)
|
||||
)
|
||||
return await self._pin_model_for_tier(tier, selected, candidates, request_kwargs, retained_pin)
|
||||
|
||||
from litellm.types.router import RoutingContext
|
||||
|
||||
|
|
@ -2369,6 +2432,40 @@ class ComplexityRouter(CustomLogger):
|
|||
self._adaptive_chosen_model_key = ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY
|
||||
return self.adaptive_router
|
||||
|
||||
def _adaptive_candidate_models(
|
||||
self,
|
||||
classified_tier: ComplexityTier | str,
|
||||
hard_floor: ComplexityTier | str | None = None,
|
||||
hard_ceiling: ComplexityTier | str | None = None,
|
||||
fit_filter: frozenset[str] | None = None,
|
||||
) -> tuple[str, ...]:
|
||||
pools: Final = self._tier_pools()
|
||||
candidates: Final = (
|
||||
tuple(pools.get(_tier_name(classified_tier), ()))
|
||||
if self.config.adaptive_eligible == "classified_tier"
|
||||
else tuple(dict.fromkeys(chain.from_iterable(pools.values())))
|
||||
)
|
||||
floor: Final = self._active_tier_severity(hard_floor) if hard_floor is not None else None
|
||||
ceiling: Final = self._active_tier_severity(hard_ceiling) if hard_ceiling is not None else None
|
||||
return tuple(
|
||||
model
|
||||
for model in _allowed(candidates, fit_filter)
|
||||
if (
|
||||
floor is None
|
||||
or any(
|
||||
self._active_tier_severity(tier) >= floor
|
||||
for tier in self._model_tiers.get(model, (classified_tier,))
|
||||
)
|
||||
)
|
||||
and (
|
||||
ceiling is None
|
||||
or any(
|
||||
self._active_tier_severity(tier) <= ceiling
|
||||
for tier in self._model_tiers.get(model, (classified_tier,))
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
def _soft_floor_pick(
|
||||
self,
|
||||
classified_tier: ComplexityTier | str,
|
||||
|
|
@ -2436,34 +2533,17 @@ class ComplexityRouter(CustomLogger):
|
|||
],
|
||||
}
|
||||
return chosen_model
|
||||
if self.config.adaptive_eligible == "classified_tier":
|
||||
candidates = list(classified_candidates)
|
||||
if not candidates:
|
||||
return self._fitting_tier_fallback(classified_tier, fit_filter)
|
||||
else:
|
||||
candidates = list(_allowed(tuple(adaptive.config.available_models), fit_filter))
|
||||
candidates: Final = self._adaptive_candidate_models(classified_tier, fit_filter=fit_filter)
|
||||
|
||||
all_costs: Final = [adaptive.model_to_cost.get(m, 0.0) for m in candidates]
|
||||
quality_weight: Final = self.config.adaptive_weights.quality
|
||||
cost_weight: Final = self.config.adaptive_weights.cost
|
||||
penalty_weight: Final = self.config.tier_distance_penalty
|
||||
|
||||
floor_severity: Final = self._active_tier_severity(hard_floor) if hard_floor is not None else None
|
||||
ceiling_severity: Final = self._active_tier_severity(hard_ceiling) if hard_ceiling is not None else None
|
||||
best_model: str | None = None
|
||||
best_score = float("-inf")
|
||||
candidate_scores: Final[list[dict[str, object]]] = []
|
||||
for model in candidates:
|
||||
if floor_severity is not None and all(
|
||||
self._active_tier_severity(model_tier) < floor_severity
|
||||
for model_tier in self._model_tiers.get(model, (classified_tier,))
|
||||
):
|
||||
continue
|
||||
if ceiling_severity is not None and all(
|
||||
self._active_tier_severity(model_tier) > ceiling_severity
|
||||
for model_tier in self._model_tiers.get(model, (classified_tier,))
|
||||
):
|
||||
continue
|
||||
for model in self._adaptive_candidate_models(classified_tier, hard_floor, hard_ceiling, fit_filter):
|
||||
cell = adaptive._cells[(request_type, model)]
|
||||
quality_sample = thompson_sample(cell)
|
||||
cost_score = normalized_cost(adaptive.model_to_cost.get(model, 0.0), all_costs)
|
||||
|
|
@ -2644,8 +2724,6 @@ class ComplexityRouter(CustomLogger):
|
|||
"""Prompt content the resolved message list never carries: the Responses API's
|
||||
`instructions`, the /v1/messages top-level `system` block, and tool definitions.
|
||||
A coding agent's context is dominated by these."""
|
||||
import json
|
||||
|
||||
instructions: Final = request_kwargs.get("instructions")
|
||||
proxy_request: Final = request_kwargs.get("proxy_server_request")
|
||||
body: Final = proxy_request.get("body") if isinstance(proxy_request, Mapping) else None
|
||||
|
|
@ -2831,19 +2909,21 @@ class ComplexityRouter(CustomLogger):
|
|||
)
|
||||
return higher_tiers[0] if higher_tiers else tier
|
||||
|
||||
def _escalated_pin(self, pinned_model: str) -> str | None:
|
||||
def _escalated_pin(self, pinned_model: str, tier: ComplexityTier | str | None = None) -> _SessionAffinityPin | None:
|
||||
"""Bump a session's pinned model to the next-higher configured tier.
|
||||
|
||||
Returns None when the pin no longer maps to any configured tier, signalling
|
||||
a full reclassification instead.
|
||||
"""
|
||||
pinned_tier: Final = self._tier_for_model(pinned_model)
|
||||
pinned_tier: Final = tier if tier is not None else self._tier_for_model(pinned_model)
|
||||
if pinned_tier is None:
|
||||
return None
|
||||
escalated_tier: Final = self._escalate_tier(pinned_tier)
|
||||
if escalated_tier == pinned_tier:
|
||||
return pinned_model
|
||||
return self.get_model_for_tier(escalated_tier)
|
||||
return _SessionAffinityPin(pinned_model, pinned_tier)
|
||||
return _SessionAffinityPin(
|
||||
self.get_model_for_tier(escalated_tier), _built_in_tier_or_none(_tier_name(escalated_tier))
|
||||
)
|
||||
|
||||
def _vision_verdicts(self, model_name: str) -> tuple[bool | None, ...]:
|
||||
"""Declared vision support per deployment serving the name: True, False, or None when
|
||||
|
|
@ -2907,6 +2987,7 @@ class ComplexityRouter(CustomLogger):
|
|||
resolved_messages: Sequence[Mapping[str, object]] | None,
|
||||
request_kwargs: dict, # mutable-ok: same shape the hook receives
|
||||
context_fit: _RequestContextFit | None = None,
|
||||
retained_pin: _SessionAffinityPin | None = None,
|
||||
) -> PreRoutingHookResponse:
|
||||
"""Replace a routed model that cannot accept this request's image input.
|
||||
|
||||
|
|
@ -2955,6 +3036,7 @@ class ComplexityRouter(CustomLogger):
|
|||
repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them
|
||||
request_kwargs,
|
||||
allowed_models=tuple(entry for entry in pools.get(capable, ()) if entry in eligible),
|
||||
retained_pin=retained_pin,
|
||||
)
|
||||
elif self._modality_default_model_usable(request_kwargs, resolved_messages, eligible):
|
||||
new_tier = None
|
||||
|
|
@ -3098,6 +3180,7 @@ class ComplexityRouter(CustomLogger):
|
|||
resolved_messages: Sequence[Mapping[str, object]] | None,
|
||||
request_kwargs: dict, # mutable-ok: same shape the hook receives
|
||||
context_fit: _RequestContextFit | None = None,
|
||||
retained_pin: _SessionAffinityPin | None = None,
|
||||
) -> PreRoutingHookResponse:
|
||||
"""Try compatible tier recovery before the default, preserving request policy and fit."""
|
||||
decision: Final = response.routing_decision
|
||||
|
|
@ -3155,6 +3238,7 @@ class ComplexityRouter(CustomLogger):
|
|||
repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them
|
||||
request_kwargs,
|
||||
allowed_models=live,
|
||||
retained_pin=retained_pin,
|
||||
)
|
||||
except ValueError as exc:
|
||||
verbose_router_logger.debug(
|
||||
|
|
@ -3247,8 +3331,13 @@ class ComplexityRouter(CustomLogger):
|
|||
"""The adaptive feedback loop reads its chosen-model marker from request metadata; a
|
||||
gate rewrite must move the marker with the model or rewards land on the displaced one."""
|
||||
metadata: Final = request_kwargs.get("metadata")
|
||||
if isinstance(metadata, dict) and metadata.get("adaptive_router_chosen_model") == old_model:
|
||||
if not isinstance(metadata, dict):
|
||||
return
|
||||
if metadata.get("adaptive_router_chosen_model") == old_model:
|
||||
metadata["adaptive_router_chosen_model"] = new_model
|
||||
decision: Final = metadata.get("adaptive_router_decision")
|
||||
if isinstance(decision, dict) and decision.get("chosen_model") == old_model:
|
||||
decision["chosen_model"] = new_model
|
||||
|
||||
def _lexical_tier_override(self, user_message: str) -> KeywordOverride | None:
|
||||
"""When keyword_tier_rules match literally, the most-severe matched tier wins.
|
||||
|
|
@ -3561,25 +3650,42 @@ class ComplexityRouter(CustomLogger):
|
|||
|
||||
if cache_key is not None and pin_replay_allowed:
|
||||
pinned_value: Final = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
|
||||
pinned_pin: Final = _parse_session_affinity_pin(pinned_value)
|
||||
pinned_pin: Final = _parse_session_affinity_pin(pinned_value, self.config.tier_names())
|
||||
if pinned_pin is not None:
|
||||
routed_model: str | None = pinned_pin.model
|
||||
pin_escalation_keyword: str | None = None
|
||||
if self.escalation_keywords:
|
||||
user_message: Final = (
|
||||
_newest_turn_ask(resolved_messages, marker_pairs) if resolved_messages else None
|
||||
user_message: Final = _newest_turn_ask(resolved_messages, marker_pairs) if resolved_messages else None
|
||||
pin_escalation_keyword: Final = (
|
||||
self._matched_escalation_keyword(user_message) if user_message is not None else None
|
||||
)
|
||||
selected_pin: Final = (
|
||||
self._escalated_pin(pinned_pin.model, pinned_pin.tier)
|
||||
if pin_escalation_keyword is not None
|
||||
else _SessionAffinityPin(
|
||||
pinned_pin.model,
|
||||
pinned_pin.tier if pinned_pin.tier is not None else self._tier_for_model(pinned_pin.model),
|
||||
)
|
||||
if user_message is not None:
|
||||
pin_escalation_keyword = self._matched_escalation_keyword(user_message)
|
||||
if pin_escalation_keyword is not None:
|
||||
routed_model = self._escalated_pin(pinned_pin.model)
|
||||
if routed_model is not None:
|
||||
escalated: Final = routed_model != pinned_pin.model
|
||||
resolved_pin_tier: Final = (
|
||||
pinned_pin.tier
|
||||
if not escalated and pinned_pin.tier is not None
|
||||
else self._tier_for_model(routed_model)
|
||||
)
|
||||
if selected_pin is not None:
|
||||
escalated: Final = selected_pin.model != pinned_pin.model or (
|
||||
pin_escalation_keyword is not None
|
||||
and pinned_pin.tier is not None
|
||||
and selected_pin.tier != pinned_pin.tier
|
||||
)
|
||||
resolved_pin_tier: Final = selected_pin.tier
|
||||
session_model: Final = (
|
||||
await self._pin_model_for_tier(
|
||||
resolved_pin_tier,
|
||||
selected_pin.model,
|
||||
tuple(self._tier_pools().get(_tier_name(resolved_pin_tier), ())),
|
||||
request_kwargs,
|
||||
)
|
||||
if escalated and resolved_pin_tier is not None
|
||||
else selected_pin.model
|
||||
)
|
||||
retained_pin: Final = _SessionAffinityPin(session_model, resolved_pin_tier)
|
||||
if resolved_pin_tier is not None:
|
||||
await self._pin_model_for_tier(
|
||||
resolved_pin_tier, session_model, (session_model,), request_kwargs
|
||||
)
|
||||
# The floor outranks the pin because plan mode is a transient state of the
|
||||
# session, not a request to move it: the turns carrying the sentinel route at
|
||||
# the floor, and the stored pin deliberately keeps the session's own model so
|
||||
|
|
@ -3590,16 +3696,28 @@ class ComplexityRouter(CustomLogger):
|
|||
plan_floored: Final = (
|
||||
pinned_tier is not None and self._apply_plan_mode_floor(pinned_tier) != pinned_tier
|
||||
)
|
||||
session_model: Final = routed_model
|
||||
if plan_floored and pinned_tier is not None:
|
||||
routed_model = self.get_model_for_tier(self._apply_plan_mode_floor(pinned_tier))
|
||||
pin_source_tier: Final = self._tier_for_model(routed_model)
|
||||
floor_model: Final = (
|
||||
await self._pick_model_for_tier(
|
||||
self._apply_plan_mode_floor(pinned_tier),
|
||||
messages,
|
||||
resolved_messages,
|
||||
request_kwargs,
|
||||
retained_pin=retained_pin,
|
||||
)
|
||||
if plan_floored and pinned_tier is not None
|
||||
else session_model
|
||||
)
|
||||
pin_source_tier: Final = (
|
||||
self._apply_plan_mode_floor(pinned_tier)
|
||||
if plan_floored and pinned_tier is not None
|
||||
else resolved_pin_tier
|
||||
)
|
||||
pin_placement: Final = (
|
||||
await self._context_window_placement(
|
||||
pin_source_tier,
|
||||
resolved_messages,
|
||||
request_kwargs,
|
||||
pool_override=(routed_model,),
|
||||
pool_override=(floor_model,),
|
||||
context_fit=context_fit,
|
||||
)
|
||||
if pin_source_tier is not None
|
||||
|
|
@ -3612,11 +3730,18 @@ class ComplexityRouter(CustomLogger):
|
|||
and _tier_name(pin_placement.tier) != _tier_name(pin_source_tier)
|
||||
else None
|
||||
)
|
||||
if pin_placement is not None and pin_context_original_tier is not None:
|
||||
# The stored pin below keeps the session's own model on purpose.
|
||||
routed_model = self._pick_from_tier_value(
|
||||
pin_placement.allowed_models, _tier_name(pin_placement.tier)
|
||||
routed_model: Final = (
|
||||
await self._pick_model_for_tier(
|
||||
pin_placement.tier,
|
||||
messages,
|
||||
resolved_messages,
|
||||
request_kwargs,
|
||||
allowed_models=pin_placement.allowed_models,
|
||||
retained_pin=retained_pin,
|
||||
)
|
||||
if pin_placement is not None and pin_context_original_tier is not None
|
||||
else floor_model
|
||||
)
|
||||
# Refresh the TTL on every hit so an active session doesn't lose its
|
||||
# pin mid-conversation just because it outlives the original write.
|
||||
await self.litellm_router_instance.cache.async_set_cache(
|
||||
|
|
@ -3644,7 +3769,7 @@ class ComplexityRouter(CustomLogger):
|
|||
routed_pin_tier: Final = (
|
||||
pin_placement.tier
|
||||
if pin_placement is not None and pin_context_original_tier is not None
|
||||
else (self._tier_for_model(routed_model) if plan_floored else resolved_pin_tier)
|
||||
else pin_source_tier
|
||||
)
|
||||
session_tier_litellm_params: Final = self._litellm_params_for_model(routed_pin_tier, routed_model)
|
||||
has_original_messages: Final = messages is not None and len(messages) > 0
|
||||
|
|
@ -3671,12 +3796,14 @@ class ComplexityRouter(CustomLogger):
|
|||
resolved_messages,
|
||||
request_kwargs,
|
||||
context_fit,
|
||||
retained_pin,
|
||||
),
|
||||
messages,
|
||||
input,
|
||||
resolved_messages,
|
||||
request_kwargs,
|
||||
context_fit,
|
||||
retained_pin,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -3961,13 +4088,21 @@ class ComplexityRouter(CustomLogger):
|
|||
housekeeping_ceiling: Final = tier if outcome.cause == "housekeeping" else None
|
||||
# A context-escalated tier becomes the hard floor: a floor the bandit can slide
|
||||
# under is not a floor.
|
||||
routed_model = self._soft_floor_pick(
|
||||
adaptive_floor: Final = tier if context_original_tier is not None else plan_floor
|
||||
adaptive_fit: Final = context_placement.holdable_models if context_placement is not None else None
|
||||
sampled_model: Final = self._soft_floor_pick(
|
||||
tier,
|
||||
ask,
|
||||
request_kwargs,
|
||||
hard_floor=tier if context_original_tier is not None else plan_floor,
|
||||
hard_floor=adaptive_floor,
|
||||
hard_ceiling=housekeeping_ceiling,
|
||||
fit_filter=context_placement.holdable_models if context_placement is not None else None,
|
||||
fit_filter=adaptive_fit,
|
||||
)
|
||||
routed_model = await self._pin_model_for_tier( # rebind-ok: reuse the eligible tier winner
|
||||
tier,
|
||||
sampled_model,
|
||||
self._adaptive_candidate_models(tier, adaptive_floor, housekeeping_ceiling, adaptive_fit),
|
||||
request_kwargs,
|
||||
)
|
||||
adaptive: Final = self._ensure_adaptive_router()
|
||||
if adaptive is not None:
|
||||
|
|
|
|||
|
|
@ -1256,20 +1256,16 @@ class ComplexityRouterConfig(BaseModel):
|
|||
deployment_affinity: bool = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"When True and a session_id is resolvable on the request, pin the deployment chosen "
|
||||
"inside each routed model group and reuse it whenever the session returns to that "
|
||||
"group, without pinning which group the session routes to. Independent of "
|
||||
"session_affinity, which pins the model group instead (and always carries this "
|
||||
"deployment pin with it): with session_affinity off, "
|
||||
"every turn is still classified on its own merits while a session that escalates to a "
|
||||
"stronger tier and comes back still lands on the deployment it used before, which is "
|
||||
"what keeps a provider prompt cache warm. Pins are held per model group, so switching "
|
||||
"tiers does not disturb the pin left behind in the previous group. On by default "
|
||||
"because re-shuffling a conversation across deployments of the same model discards "
|
||||
"that cache for no benefit; set False to keep every turn load-balanced across the "
|
||||
"group, which is what a deployment set with tight per-deployment rate limits wants. "
|
||||
"Inert when no session_id is resolvable, since there is nothing to key a pin on, and "
|
||||
"suppressed when plugins are configured, for the same reason session_affinity is."
|
||||
"When True and a client session_id is resolvable, reuse the session's chosen model "
|
||||
"for each classified tier and its deployment within each model group. With "
|
||||
"session_affinity off, every turn is still classified: moving to another tier leaves "
|
||||
"the previous tier's model pin intact for a later return. Pins yield to current "
|
||||
"candidate, context, modality, and availability constraints. Adaptive selection chooses "
|
||||
"the initial model from its eligible pool, then reuses that choice per tier. This "
|
||||
"reduces avoidable provider prompt-cache misses; it does not guarantee cache hits. "
|
||||
"Set False to select models and load-balance deployments on every turn, unless "
|
||||
"session_affinity or user_turn classification requires a pin. Inert without a client "
|
||||
"session_id and suppressed when plugins are configured."
|
||||
),
|
||||
)
|
||||
session_affinity_ttl_seconds: int = Field(
|
||||
|
|
@ -1277,7 +1273,7 @@ class ComplexityRouterConfig(BaseModel):
|
|||
gt=0,
|
||||
description=(
|
||||
"TTL for the session affinity pin; refreshed on every cache hit. Bounds both the "
|
||||
"session_affinity model pin and the deployment_affinity deployment pin, so it measures "
|
||||
"session_affinity model pin and the deployment_affinity per-tier model and deployment pins, so it measures "
|
||||
"idle time for the session's routing decisions rather than total session length"
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,71 +1,67 @@
|
|||
"""
|
||||
Returns a random deployment from the list of healthy deployments.
|
||||
"""Choose among eligible deployments using request weights, then global metrics."""
|
||||
|
||||
If weights are provided, it will return a deployment based on the weights.
|
||||
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from itertools import chain
|
||||
from typing import Final, TypeVar
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.types.router_weights import validate_router_weights
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router as _Router
|
||||
_DeploymentT = TypeVar("_DeploymentT", bound=Mapping[str, object])
|
||||
_ROUTER_LOGGER: Final = logging.getLogger("LiteLLM Router")
|
||||
|
||||
LitellmRouter = _Router
|
||||
else:
|
||||
LitellmRouter = Any
|
||||
|
||||
def _metric_weight(deployment: Mapping[str, object], metric: str) -> float:
|
||||
params: Final = deployment.get("litellm_params")
|
||||
value: Final = params.get(metric) if isinstance(params, Mapping) else None
|
||||
if value is None:
|
||||
return 0.0
|
||||
if isinstance(value, (int, float)):
|
||||
return float(value)
|
||||
raise TypeError(f"Deployment {metric} must be numeric")
|
||||
|
||||
|
||||
def _scoped_weights(
|
||||
deployments: Sequence[Mapping[str, object]],
|
||||
model: str,
|
||||
request_kwargs: Mapping[str, object] | None,
|
||||
) -> tuple[float, ...]:
|
||||
settings: Final = validate_router_weights((request_kwargs or {}).get("_router_weights"))
|
||||
model_weights: Final = settings.get(model) if settings is not None else None
|
||||
if not model_weights:
|
||||
return ()
|
||||
return tuple(
|
||||
model_weights.get(str(info.get("id")), 0.0) if isinstance(info, Mapping) else 0.0
|
||||
for deployment in deployments
|
||||
for info in (deployment.get("model_info"),)
|
||||
)
|
||||
|
||||
|
||||
def simple_shuffle(
|
||||
llm_router_instance: LitellmRouter,
|
||||
healthy_deployments: list[Any] | dict[Any, Any],
|
||||
resolve_model_alias: Callable[[str], str | None],
|
||||
healthy_deployments: Sequence[_DeploymentT],
|
||||
model: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Returns a random deployment from the list of healthy deployments.
|
||||
|
||||
If weights are provided, it will return a deployment based on the weights.
|
||||
|
||||
If users pass `rpm` or `tpm`, we do a random weighted pick - based on `rpm`/`tpm`.
|
||||
|
||||
Args:
|
||||
llm_router_instance: LitellmRouter instance
|
||||
healthy_deployments: List of healthy deployments
|
||||
model: Model name
|
||||
|
||||
Returns:
|
||||
Dict: A single healthy deployment
|
||||
"""
|
||||
|
||||
############## Check if 'weight' or 'rpm' or 'tpm' param set for a weighted pick #################
|
||||
for weight_by in ["weight", "rpm", "tpm"]:
|
||||
if any(m["litellm_params"].get(weight_by) is not None for m in healthy_deployments):
|
||||
weights = [m["litellm_params"].get(weight_by, 0) for m in healthy_deployments]
|
||||
verbose_router_logger.debug("\nweight %s", weights)
|
||||
total_weight = sum(weights)
|
||||
if total_weight <= 0:
|
||||
# All remaining candidates have weight 0 for this metric (e.g.
|
||||
# after a weighted-failover exclusion left only zero-weight
|
||||
# backups). Skip to the next metric (rpm/tpm) which may still
|
||||
# provide a meaningful weighted pick; if none do, we fall
|
||||
# through to the uniform random pick at the end.
|
||||
continue
|
||||
weights = [weight / total_weight for weight in weights]
|
||||
verbose_router_logger.debug("\n weights %s by %s", weights, weight_by)
|
||||
# Perform weighted random pick
|
||||
selected_index = random.choices(range(len(weights)), weights=weights)[0]
|
||||
verbose_router_logger.debug("\n selected index, %s", selected_index)
|
||||
deployment = healthy_deployments[selected_index]
|
||||
verbose_router_logger.info(
|
||||
"get_available_deployment for model: %s, Selected deployment: %s for model: %s",
|
||||
model,
|
||||
llm_router_instance.print_deployment(deployment) or deployment[0],
|
||||
model,
|
||||
)
|
||||
return deployment or deployment[0]
|
||||
|
||||
############## No RPM/TPM passed, we do a random pick #################
|
||||
item: Final = random.choice(healthy_deployments)
|
||||
return item or item[0]
|
||||
request_kwargs: Mapping[str, object] | None,
|
||||
) -> _DeploymentT:
|
||||
resolved_model: Final = resolve_model_alias(model) or model
|
||||
weight_sets: Final = chain(
|
||||
(_scoped_weights(healthy_deployments, resolved_model, request_kwargs),),
|
||||
(
|
||||
tuple(_metric_weight(deployment, metric) for deployment in healthy_deployments)
|
||||
for metric in ("weight", "rpm", "tpm")
|
||||
),
|
||||
)
|
||||
for weights in weight_sets:
|
||||
largest = max(weights, default=0.0)
|
||||
if largest <= 0:
|
||||
continue
|
||||
normalized = tuple(weight / largest for weight in weights)
|
||||
if sum(normalized) <= 0:
|
||||
continue
|
||||
selected = random.choices(healthy_deployments, weights=normalized)[0]
|
||||
_ROUTER_LOGGER.info("Selected deployment for model %s: %s", model, selected.get("model_info"))
|
||||
return selected
|
||||
return random.choice(healthy_deployments)
|
||||
|
|
|
|||
|
|
@ -13,13 +13,13 @@ where routing to a consistent deployment is still beneficial.
|
|||
"""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, cast
|
||||
|
||||
from typing_extensions import TypedDict
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.affinity_cache import claim_affinity_pin, claim_affinity_pin_in_memory, set_local_affinity_pin
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.constants import SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY, SESSION_ID_GENERATED_METADATA_KEY
|
||||
from litellm.integrations.custom_logger import CustomLogger, Span
|
||||
|
|
@ -28,8 +28,8 @@ from litellm.types.llms.openai import AllMessageValues
|
|||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
class DeploymentAffinityCacheValue(TypedDict):
|
||||
model_id: str
|
||||
class DeploymentAffinityCacheValue(TypedDict, closed=True):
|
||||
model_id: ReadOnly[str]
|
||||
|
||||
|
||||
VALID_MODEL_GROUP_AFFINITY_FLAGS: Final = frozenset(
|
||||
|
|
@ -60,19 +60,6 @@ def warn_on_unknown_model_group_affinity_flags(model_group_affinity_config: Mapp
|
|||
)
|
||||
|
||||
|
||||
_CLAIM_PIN_SCRIPT: Final = """
|
||||
local current = redis.call('GET', KEYS[1])
|
||||
if current == false then
|
||||
redis.call('SET', KEYS[1], ARGV[1], 'EX', ARGV[2])
|
||||
return ARGV[1]
|
||||
end
|
||||
if current == ARGV[1] then
|
||||
redis.call('EXPIRE', KEYS[1], ARGV[2])
|
||||
end
|
||||
return current
|
||||
"""
|
||||
|
||||
|
||||
class DeploymentAffinityCheck(CustomLogger):
|
||||
"""
|
||||
Router deployment affinity callback.
|
||||
|
|
@ -255,34 +242,33 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return f"{cls.CACHE_KEY_PREFIX}:session:{model_group}:{hashed_user_key}:{session_id}"
|
||||
|
||||
@staticmethod
|
||||
def _get_session_id_from_metadata_dict(metadata: dict) -> str | None:
|
||||
def _get_session_id_from_metadata_dict(metadata: Mapping[object, object]) -> str | None:
|
||||
session_id: Final = metadata.get("session_id")
|
||||
if session_id is None or metadata.get(SESSION_ID_GENERATED_METADATA_KEY):
|
||||
return None
|
||||
return str(session_id)
|
||||
|
||||
@staticmethod
|
||||
def _iter_metadata_dicts(request_kwargs: dict) -> list[dict]:
|
||||
def _iter_metadata_dicts(request_kwargs: Mapping[str, object]) -> tuple[Mapping[object, object], ...]:
|
||||
"""
|
||||
Return all metadata dicts available on the request.
|
||||
|
||||
Depending on the endpoint, Router may populate `metadata` or `litellm_metadata`.
|
||||
Users may also send one or both, so we check both (rather than using `or`).
|
||||
"""
|
||||
metadata_dicts: Final[list[dict]] = []
|
||||
for key in ("litellm_metadata", "metadata"):
|
||||
md = request_kwargs.get(key)
|
||||
if isinstance(md, dict):
|
||||
metadata_dicts.append(md)
|
||||
return metadata_dicts
|
||||
return tuple(
|
||||
cast(Mapping[object, object], metadata) # cast-ok: isinstance proves mapping shape; values remain opaque
|
||||
for key in ("litellm_metadata", "metadata")
|
||||
if isinstance(metadata := request_kwargs.get(key), dict)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _first_metadata_value(metadata_dicts: Sequence[dict], key: str) -> str | None:
|
||||
def _first_metadata_value(metadata_dicts: Sequence[Mapping[object, object]], key: str) -> str | None:
|
||||
value: Final = next((metadata[key] for metadata in metadata_dicts if metadata.get(key) is not None), None)
|
||||
return None if value is None else str(value)
|
||||
|
||||
@classmethod
|
||||
def _get_user_key_from_request_kwargs(cls, request_kwargs: dict) -> str | None:
|
||||
def get_user_key_from_request_kwargs(cls, request_kwargs: Mapping[str, object]) -> str | None:
|
||||
"""
|
||||
Extract a stable affinity key from request kwargs.
|
||||
|
||||
|
|
@ -334,74 +320,17 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return None
|
||||
|
||||
def _set_local_pin(self, cache_key: str, value: object, ttl_seconds: int) -> None:
|
||||
"""The one owner of authoritative local pin writes: a plain set keeps a live
|
||||
key's original expiry (`allow_ttl_override`), so the entry is replaced to make
|
||||
the TTL real. Every local pin write goes through here so the redis-winner sync
|
||||
and the pod-local claim can never disagree about expiry again."""
|
||||
self.cache.in_memory_cache.delete_cache(cache_key)
|
||||
self.cache.in_memory_cache.set_cache(cache_key, value, ttl=ttl_seconds)
|
||||
set_local_affinity_pin(self.cache, cache_key, value, ttl_seconds)
|
||||
|
||||
async def _claim_pin(self, cache_key: str, pin_value: DeploymentAffinityCacheValue, ttl_seconds: int) -> str | None:
|
||||
"""First-writer-wins pin write: store `pin_value` only when the key is absent and
|
||||
return the deployment id the key holds afterwards, so a caller learns whether it won
|
||||
by comparing against its own id, and None when the stored value is one no reader can
|
||||
interpret. Concurrent claimers converge on the
|
||||
first write instead of the last. Re-claiming with the stored value refreshes its
|
||||
TTL, the same keepalive the complexity router's model pin documents: an active
|
||||
session must not lose its pin mid-conversation just because it outlives the
|
||||
original write, so the affinity TTL (the Router's
|
||||
`deployment_affinity_ttl_seconds`, or a pre-routing hook's per-request
|
||||
`session_affinity_ttl_seconds` override) bounds idle time, not total
|
||||
session length. On Redis one Lua script does the get-or-set-or-refresh
|
||||
atomically (same registration seam the rate limiters use) and the in-memory
|
||||
tier is synchronized to the winner; without Redis, and whenever Redis is
|
||||
unreachable, the pod-local check-and-set below stands in and is atomic because it
|
||||
runs synchronously on the event loop. Degrading to a pod-local claim rather than
|
||||
propagating the fault is what keeps same-pod stickiness through a Redis blip: the
|
||||
caller only logs this result, so an escaping error would leave the session with no
|
||||
pin at all and reshuffle every turn for the outage, which is worse than losing
|
||||
cross-pod agreement. The redis tier is
|
||||
resolved per call because the proxy attaches it after Router construction
|
||||
(`Router._update_redis_cache`); the compiled script is cached per event loop
|
||||
underneath the registration seam.
|
||||
"""
|
||||
redis_cache: Final = self.cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
claim_script: Final = redis_cache.async_register_script(_CLAIM_PIN_SCRIPT)
|
||||
raw: Final = await claim_script(keys=(cache_key,), args=(json.dumps(pin_value), int(ttl_seconds)))
|
||||
decoded: Final = raw.decode("utf-8") if isinstance(raw, bytes) else raw
|
||||
if not isinstance(decoded, str):
|
||||
return pin_value["model_id"]
|
||||
try:
|
||||
winner: object = json.loads(decoded)
|
||||
except json.JSONDecodeError:
|
||||
winner = decoded
|
||||
self._set_local_pin(cache_key=cache_key, value=winner, ttl_seconds=ttl_seconds)
|
||||
return self._pinned_model_id(winner)
|
||||
except Exception as e: # noqa: BLE001 # any Redis/Lua failure degrades to the pod-local claim, never unpins
|
||||
verbose_router_logger.debug(
|
||||
"DeploymentAffinityCheck: redis pin claim failed, falling back to pod-local claim. error=%s", e
|
||||
)
|
||||
|
||||
return self._claim_pin_in_memory(cache_key=cache_key, pin_value=pin_value, ttl_seconds=ttl_seconds)
|
||||
winner: Final = await claim_affinity_pin(self.cache, cache_key, pin_value, ttl_seconds)
|
||||
return self._pinned_model_id(winner)
|
||||
|
||||
def _claim_pin_in_memory(
|
||||
self, cache_key: str, pin_value: DeploymentAffinityCacheValue, ttl_seconds: int
|
||||
) -> str | None:
|
||||
"""Pod-local half of the claim, used when no Redis tier is attached and as the
|
||||
fallback when the Redis claim fails. Mirrors the Lua script exactly, including
|
||||
the keepalive: re-claiming with the stored value slides the idle window through
|
||||
`_set_local_pin`. Both branches stay synchronous, hence atomic on the event
|
||||
loop."""
|
||||
existing: Final = self.cache.in_memory_cache.get_cache(cache_key)
|
||||
if existing is not None:
|
||||
existing_model_id: Final = self._pinned_model_id(existing)
|
||||
if existing_model_id == pin_value["model_id"]:
|
||||
self._set_local_pin(cache_key=cache_key, value=pin_value, ttl_seconds=ttl_seconds)
|
||||
return existing_model_id
|
||||
self._set_local_pin(cache_key=cache_key, value=pin_value, ttl_seconds=ttl_seconds)
|
||||
return pin_value["model_id"]
|
||||
winner: Final = claim_affinity_pin_in_memory(self.cache, cache_key, pin_value, ttl_seconds)
|
||||
return self._pinned_model_id(winner)
|
||||
|
||||
@staticmethod
|
||||
def _find_deployment_by_model_id(healthy_deployments: list[dict], model_id: str) -> dict | None:
|
||||
|
|
@ -465,7 +394,7 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
enable_session_id or self._get_marker_session_affinity_ttl(request_kwargs=request_kwargs) is not None
|
||||
)
|
||||
user_key: Final = (
|
||||
self._get_user_key_from_request_kwargs(request_kwargs=request_kwargs)
|
||||
self.get_user_key_from_request_kwargs(request_kwargs=request_kwargs)
|
||||
if (session_affinity_active or enable_user_key)
|
||||
else None
|
||||
)
|
||||
|
|
@ -580,7 +509,7 @@ class DeploymentAffinityCheck(CustomLogger):
|
|||
return None
|
||||
|
||||
user_key: Final = (
|
||||
self._get_user_key_from_request_kwargs(request_kwargs=kwargs)
|
||||
self.get_user_key_from_request_kwargs(request_kwargs=kwargs)
|
||||
if (enable_user_key or session_affinity_active)
|
||||
else None
|
||||
)
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ from litellm.exceptions import (
|
|||
ServiceUnavailableError,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger, Span
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
encrypted_content_of_block,
|
||||
strip_encrypted_reasoning_from_messages,
|
||||
|
|
@ -215,11 +216,11 @@ class EncryptedContentAffinityCheck(CustomLogger):
|
|||
@staticmethod
|
||||
def _encryption_boundary_key(
|
||||
litellm_params: object,
|
||||
) -> tuple | None:
|
||||
) -> tuple[object, object] | None:
|
||||
"""
|
||||
``(api_base, api_key)`` pair identifying an Azure resource. Two
|
||||
deployments sharing both are interchangeable for ``encrypted_content``
|
||||
follow-ups; Azure rejects content produced by any other resource.
|
||||
``(api_base, api_key)`` identifies an upstream encryption boundary.
|
||||
The values are resolved from the deployment and its named credential
|
||||
without modifying the deployment.
|
||||
|
||||
Accepts any object exposing dict-style ``.get(key, default)``: plain
|
||||
dicts (the common case in ``healthy_deployments``) as well as
|
||||
|
|
@ -234,9 +235,25 @@ class EncryptedContentAffinityCheck(CustomLogger):
|
|||
return None
|
||||
api_base: Final = getter("api_base")
|
||||
api_key: Final = getter("api_key")
|
||||
if not api_base or not api_key:
|
||||
credential_name: Final = getter("litellm_credential_name")
|
||||
credential_values: Final[Mapping[str, object] | None] = (
|
||||
CredentialAccessor.get_credential_values(credential_name)
|
||||
if isinstance(credential_name, str) and credential_name
|
||||
else None
|
||||
)
|
||||
effective_api_base: Final = (
|
||||
credential_values.get("api_base")
|
||||
if credential_values is not None and "api_base" in credential_values
|
||||
else api_base
|
||||
)
|
||||
effective_api_key: Final = (
|
||||
credential_values.get("api_key")
|
||||
if credential_values is not None and "api_key" in credential_values
|
||||
else api_key
|
||||
)
|
||||
if not effective_api_base or not effective_api_key:
|
||||
return None
|
||||
return (api_base, api_key)
|
||||
return (effective_api_base, effective_api_key)
|
||||
|
||||
def _find_deployments_on_same_encryption_boundary(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from typing import Any, Final, Literal
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, Literal, cast # noqa: TID251 # JSON chat rows have no typed constructor across roles
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import TypedDict
|
||||
|
|
@ -158,12 +159,21 @@ def coerce_stream_holdback_value(value: Any) -> int:
|
|||
return 0
|
||||
|
||||
|
||||
def structured_messages_from_response(value: object) -> Sequence[AllMessageValues] | None:
|
||||
if not isinstance(value, list):
|
||||
return None
|
||||
if not all(isinstance(message, Mapping) and isinstance(message.get("role"), str) for message in value):
|
||||
return None
|
||||
return cast("Sequence[AllMessageValues]", value) # cast-ok: JSON rows checked for a role, the same trust texts get
|
||||
|
||||
|
||||
class GenericGuardrailAPIResponse:
|
||||
"""Response model for the Generic Guardrail API"""
|
||||
|
||||
texts: list[str] | None
|
||||
images: list[str] | None
|
||||
tools: list[GuardrailToolParam] | None
|
||||
structured_messages: Sequence[AllMessageValues] | None
|
||||
action: str
|
||||
blocked_reason: str | None
|
||||
stream_holdback_chars: list[int] | None
|
||||
|
|
@ -176,12 +186,14 @@ class GenericGuardrailAPIResponse:
|
|||
images: list[str] | None = None,
|
||||
tools: list[GuardrailToolParam] | None = None,
|
||||
stream_holdback_chars: list[int] | None = None,
|
||||
structured_messages: Sequence[AllMessageValues] | None = None,
|
||||
) -> None:
|
||||
self.action = action
|
||||
self.blocked_reason = blocked_reason
|
||||
self.texts = texts
|
||||
self.images = images
|
||||
self.tools = tools
|
||||
self.structured_messages = structured_messages
|
||||
# Number of trailing chars, indexed the same as ``texts``, that the
|
||||
# framework must withhold from streaming emission until the next
|
||||
# processing round (word-boundary safety for text transformations).
|
||||
|
|
@ -200,4 +212,5 @@ class GenericGuardrailAPIResponse:
|
|||
images=data.get("images"),
|
||||
tools=data.get("tools"),
|
||||
stream_holdback_chars=stream_holdback_chars,
|
||||
structured_messages=structured_messages_from_response(data.get("structured_messages")),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,15 +1,18 @@
|
|||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel, field_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_UserTableWithKeyCount,
|
||||
NewUserRequest,
|
||||
UpdateUserRequest,
|
||||
UpdateUserRequestNoUserIDorEmail,
|
||||
)
|
||||
|
||||
MAX_BULK_NEW_USERS: Final = 500
|
||||
|
||||
|
||||
class InsensitiveContains(TypedDict):
|
||||
contains: ReadOnly[str]
|
||||
|
|
@ -83,3 +86,50 @@ class BulkUpdateUserResponse(BaseModel):
|
|||
total_requested: int
|
||||
successful_updates: int
|
||||
failed_updates: int
|
||||
|
||||
|
||||
class BulkNewUserItem(NewUserRequest):
|
||||
"""One row of `POST /management/v1/users/bulk`: the `/user/new` body, with keys opt-in and invite emails
|
||||
unsupported. Unknown fields are rejected, as on every `/management/v1` request body."""
|
||||
|
||||
model_config = ConfigDict(extra="forbid", protected_namespaces=())
|
||||
|
||||
auto_create_key: bool = False
|
||||
|
||||
@field_validator("send_invite_email")
|
||||
@classmethod
|
||||
def reject_invite_email(cls, value: bool | None) -> bool | None:
|
||||
if value:
|
||||
raise ValueError("send_invite_email is not supported on /management/v1/users/bulk; invite users separately")
|
||||
return value
|
||||
|
||||
|
||||
class BulkNewUserRequest(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
users: Sequence[BulkNewUserItem] = Field(min_length=1, max_length=MAX_BULK_NEW_USERS)
|
||||
|
||||
|
||||
class UserCreateResult(BaseModel):
|
||||
"""Outcome for one row of `POST /management/v1/users/bulk`. `teams` lists the teams the user was actually
|
||||
added to."""
|
||||
|
||||
user_id: str | None = None
|
||||
user_email: str | None = None
|
||||
success: bool
|
||||
teams: tuple[str, ...] | None = None
|
||||
key: str | None = None
|
||||
error: str | None = None
|
||||
|
||||
|
||||
class BulkNewUserMeta(BaseModel):
|
||||
total_requested: int
|
||||
created: int
|
||||
failed: int
|
||||
|
||||
|
||||
class BulkNewUserResponse(BaseModel):
|
||||
"""`data` holds one result per input row, in input order."""
|
||||
|
||||
data: tuple[UserCreateResult, ...]
|
||||
meta: BulkNewUserMeta
|
||||
|
|
|
|||
|
|
@ -1,9 +1,12 @@
|
|||
from datetime import datetime
|
||||
from typing import Any, Final, Literal
|
||||
from typing import Any, Final, Literal, TypeAlias
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, model_validator
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.models.verification_token import LiteLLM_VerificationToken
|
||||
from litellm.proxy._types import GenerateKeyRequest, RegenerateKeyRequest, UpdateKeyRequest
|
||||
from litellm.types.llms.base import LiteLLMPydanticObjectBase
|
||||
from litellm.types.proxy.management_endpoints.internal_user_endpoints import InsensitiveContains
|
||||
|
||||
|
||||
|
|
@ -123,3 +126,24 @@ class BulkUpdateTeamKeysRequest(BaseModel):
|
|||
if not has_key_ids and not self.all_keys_in_team:
|
||||
raise ValueError("Must provide either `key_ids` (non-empty) or `all_keys_in_team=True`.")
|
||||
return self
|
||||
|
||||
|
||||
CustomKeyPolicyOperation: TypeAlias = Literal["generate", "update", "regenerate"]
|
||||
|
||||
|
||||
class CustomKeyPolicyRequest(LiteLLMPydanticObjectBase):
|
||||
"""What `general_settings.custom_key_policy` receives.
|
||||
|
||||
`effective_key` is the verification token row as it will be written: the existing row overlaid with the
|
||||
requested changes, with `duration` resolved to `expires` and `budget_duration` to `budget_reset_at`. Values the
|
||||
proxy fills in after the policy stay at their defaults: `token`, `key_name`, `created_by`, `updated_by` and the
|
||||
soft-budget `budget_id` on generate, the rotated token on regenerate, and the `object_permission` relation on
|
||||
every operation (`object_permission_id` is set; read `request.object_permission` for the requested change).
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(protected_namespaces=(), frozen=True)
|
||||
|
||||
operation: CustomKeyPolicyOperation
|
||||
existing_key: LiteLLM_VerificationToken | None
|
||||
effective_key: LiteLLM_VerificationToken
|
||||
request: GenerateKeyRequest | UpdateKeyRequest | RegenerateKeyRequest
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_c
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
|
||||
from litellm.types.router_weights import RouterWeights
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.router import Router
|
||||
|
|
@ -146,6 +147,7 @@ class UpdateRouterConfig(BaseModel):
|
|||
context_window_fallbacks: list[dict] | None = None
|
||||
model_group_alias: dict[str, str | dict] | None = {}
|
||||
enable_tag_filtering: bool | None = None
|
||||
weights: RouterWeights | None = None
|
||||
tag_routing_prefix: str | None = None
|
||||
optional_pre_call_checks: OptionalPreCallChecks | None = None
|
||||
|
||||
|
|
|
|||
30
litellm/types/router_weights.py
Normal file
30
litellm/types/router_weights.py
Normal file
|
|
@ -0,0 +1,30 @@
|
|||
from collections.abc import Mapping
|
||||
from typing import Annotated, Final
|
||||
|
||||
from pydantic import AfterValidator, Field, TypeAdapter
|
||||
|
||||
|
||||
def _validate_positive_router_weights(weights: Mapping[str, Mapping[str, float]]) -> Mapping[str, Mapping[str, float]]:
|
||||
if any(group and not any(weight > 0 for weight in group.values()) for group in weights.values()):
|
||||
raise ValueError("Each nonempty weights group must contain at least one positive weight")
|
||||
return weights
|
||||
|
||||
|
||||
RouterWeightIdentifier = Annotated[str, Field(strict=True, min_length=1, pattern=r"\S")]
|
||||
RouterWeight = Annotated[float, Field(strict=True, ge=0, allow_inf_nan=False)]
|
||||
RouterWeights = Annotated[
|
||||
dict[RouterWeightIdentifier, dict[RouterWeightIdentifier, RouterWeight]],
|
||||
AfterValidator(_validate_positive_router_weights),
|
||||
]
|
||||
_ROUTER_WEIGHTS_ADAPTER: Final[TypeAdapter[RouterWeights | None]] = TypeAdapter(RouterWeights | None)
|
||||
_ROUTER_SETTINGS_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def validate_router_weights(value: object) -> RouterWeights | None:
|
||||
return _ROUTER_WEIGHTS_ADAPTER.validate_python(value)
|
||||
|
||||
|
||||
def validate_router_settings_dict(value: object) -> dict[str, object]:
|
||||
settings: Final = _ROUTER_SETTINGS_DICT_ADAPTER.validate_python(value)
|
||||
validate_router_weights(settings.get("weights"))
|
||||
return settings
|
||||
|
|
@ -3782,6 +3782,7 @@ all_litellm_params = (
|
|||
"id",
|
||||
"fallbacks",
|
||||
"routing_strategy",
|
||||
"_router_weights",
|
||||
"azure",
|
||||
"headers",
|
||||
"model_list",
|
||||
|
|
|
|||
|
|
@ -1208,30 +1208,13 @@ def _dispatch_success_logging(
|
|||
is_litellm_internal_call: bool,
|
||||
) -> None:
|
||||
if not is_litellm_internal_call:
|
||||
if getattr(logging_obj, "_defer_async_logging", False):
|
||||
|
||||
def _enqueue_deferred_logging() -> None:
|
||||
asyncio.create_task(
|
||||
_client_async_logging_helper(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
)
|
||||
)
|
||||
|
||||
logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging
|
||||
else:
|
||||
asyncio.create_task(
|
||||
_client_async_logging_helper(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
)
|
||||
)
|
||||
_schedule_async_success_logging(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
)
|
||||
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls(
|
||||
result=result,
|
||||
|
|
@ -1240,6 +1223,43 @@ def _dispatch_success_logging(
|
|||
)
|
||||
|
||||
|
||||
def _schedule_async_success_logging(
|
||||
logging_obj: LiteLLMLoggingObject,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
is_completion_with_fallbacks: bool,
|
||||
) -> None:
|
||||
"""Fire the async success log for ``result`` now, or park it on the logging object while
|
||||
the proxy defers logging past its post-call guardrails.
|
||||
|
||||
Nested @client wrappers (Anthropic Messages over the chat adapter, chat over the Responses
|
||||
bridge) each exit through here with the same logging object and their own shape of the same
|
||||
response. The immediate path already logs one request once, since the first task marks
|
||||
``has_logged_async_success`` and the later ones skip. The deferred slot keeps the same
|
||||
first-wins rule: the innermost wrapper's provider-shaped result is the one the spend log
|
||||
reads usage from, and a later wrapper never swaps in its client-shaped translation.
|
||||
"""
|
||||
|
||||
def _enqueue_async_logging() -> None:
|
||||
asyncio.create_task(
|
||||
_client_async_logging_helper(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
)
|
||||
)
|
||||
|
||||
if not getattr(logging_obj, "_defer_async_logging", False):
|
||||
_enqueue_async_logging()
|
||||
return
|
||||
if getattr(logging_obj, "_enqueue_deferred_logging", None) is not None:
|
||||
return
|
||||
logging_obj._enqueue_deferred_logging = _enqueue_async_logging
|
||||
|
||||
|
||||
async def _client_async_logging_helper(
|
||||
logging_obj: LiteLLMLoggingObject,
|
||||
result,
|
||||
|
|
|
|||
|
|
@ -38,6 +38,9 @@ longer signal it.
|
|||
### Fixed
|
||||
|
||||
- **key**: An update that changes `team_id` and fails because the key was already cascade-deleted along with its previous team now recovers by recreating the key under the new team, instead of aborting the apply. The key's absence is confirmed against the proxy first, so an unrelated failure still errors out, and a `team_id` change between two teams that both still exist stays a plain in-place update
|
||||
- **credential**: create now reports a `credential_name` collision as a clear error naming the `terraform import` command that adopts the existing credential, instead of surfacing the proxy's raw 500 with a Prisma `Unique constraint failed` message. New `adopt_existing` argument (default `false`) opts into taking the existing credential over during create, which makes `apply` idempotent again once state loses track of a credential that still exists on the proxy. Requires a proxy that answers 409 on the collision; older proxies are still detected by their 500 message
|
||||
- **credential**: credential names and `model_id` are now percent-encoded in request URLs, so a name containing `/`, `?`, `#` or spaces reaches the proxy intact instead of being cut at the first reserved character and read, updated or deleted as a different credential
|
||||
- **credential**: update now sends `model_id`, so a `model_id`-scoped credential keeps resolving its values from that deployment on update and on adoption instead of being overwritten with the literal `credential_values`; needs a proxy from 1.102.0, older proxies ignore the field
|
||||
- **team**: Read now decodes the `team_info` envelope `/team/info` actually returns, so team attributes refresh from the proxy instead of always falling back to the prior state
|
||||
- **key**: Read now unwraps the `info` envelope `/key/info` actually returns; previously reads mapped nothing back into state, so drift on a key was never detected
|
||||
- **key**: Read now picks up `model_rpm_limit`, `model_tpm_limit`, `guardrails`, `tags`, `enforced_params`, `allowed_passthrough_routes`, `rpm_limit_type`, `tpm_limit_type` and `prompts` from `info.metadata`, where the proxy actually stores them; previously they stayed empty in state, so a matching config showed a permanent phantom diff on them and out-of-band changes to them were never detected
|
||||
|
|
|
|||
|
|
@ -130,6 +130,7 @@ The following arguments are supported:
|
|||
* `credential_values` - (Required, Sensitive) Map of sensitive credential values such as API keys, tokens, etc.
|
||||
* `model_id` - (Optional) Model ID associated with this credential.
|
||||
* `credential_info` - (Optional) Map of additional non-sensitive information about the credential.
|
||||
* `adopt_existing` - (Optional, default `false`) Take over a credential of this name that already exists on the proxy instead of failing. Turning this on overwrites the existing credential's values with the ones in this configuration.
|
||||
|
||||
## Attributes Reference
|
||||
|
||||
|
|
|
|||
|
|
@ -39,6 +39,15 @@ func resourceLiteLLMCredential() *schema.Resource {
|
|||
Elem: &schema.Schema{Type: schema.TypeString},
|
||||
Description: "Sensitive credential values (API keys, tokens, etc.)",
|
||||
},
|
||||
"adopt_existing": {
|
||||
Type: schema.TypeBool,
|
||||
Optional: true,
|
||||
Default: false,
|
||||
Description: "Take over a credential of this name that already exists on the proxy instead of failing. " +
|
||||
"Off by default: create reports the conflict and points at `terraform import`, so an apply never " +
|
||||
"silently overwrites a credential it does not manage. Turning this on overwrites the existing " +
|
||||
"credential's values with the ones in this configuration.",
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,15 +1,23 @@
|
|||
package litellm
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema"
|
||||
)
|
||||
|
||||
const (
|
||||
endpointCredential = "/credentials/%s"
|
||||
endpointCredentialByName = "/credentials/by_name/%s"
|
||||
endpointCredentialByNameForModel = "/credentials/by_name/%s?model_id=%s"
|
||||
)
|
||||
|
||||
// retryCredentialRead attempts to read a credential with exponential backoff.
|
||||
// If the read path clears the ID (e.g., transient 404 right after create),
|
||||
// we treat it as retryable instead of accepting an empty state.
|
||||
|
|
@ -53,34 +61,28 @@ func retryCredentialRead(d *schema.ResourceData, m interface{}, maxRetries int)
|
|||
return err
|
||||
}
|
||||
|
||||
func resourceLiteLLMCredentialCreate(d *schema.ResourceData, m interface{}) error {
|
||||
client := m.(*Client)
|
||||
|
||||
credentialName := d.Get("credential_name").(string)
|
||||
modelID := d.Get("model_id").(string)
|
||||
credentialInfo := d.Get("credential_info").(map[string]interface{})
|
||||
credentialValues := d.Get("credential_values").(map[string]interface{})
|
||||
|
||||
// Convert credential_info to map[string]interface{} for JSON
|
||||
func credentialRequestFromResource(d *schema.ResourceData, credentialName string) CredentialRequest {
|
||||
credInfoMap := make(map[string]interface{})
|
||||
for k, v := range credentialInfo {
|
||||
for k, v := range d.Get("credential_info").(map[string]interface{}) {
|
||||
credInfoMap[k] = v
|
||||
}
|
||||
|
||||
// Convert credential_values to map[string]interface{} for JSON
|
||||
credValuesMap := make(map[string]interface{})
|
||||
for k, v := range credentialValues {
|
||||
for k, v := range d.Get("credential_values").(map[string]interface{}) {
|
||||
credValuesMap[k] = v
|
||||
}
|
||||
|
||||
credentialRequest := CredentialRequest{
|
||||
return CredentialRequest{
|
||||
CredentialName: credentialName,
|
||||
ModelID: modelID,
|
||||
ModelID: d.Get("model_id").(string),
|
||||
CredentialInfo: credInfoMap,
|
||||
CredentialValues: credValuesMap,
|
||||
}
|
||||
}
|
||||
|
||||
resp, err := MakeRequest(client, "POST", "/credentials", credentialRequest)
|
||||
func resourceLiteLLMCredentialCreate(d *schema.ResourceData, m interface{}) error {
|
||||
client := m.(*Client)
|
||||
credentialName := d.Get("credential_name").(string)
|
||||
|
||||
resp, err := MakeRequest(client, "POST", "/credentials", credentialRequestFromResource(d, credentialName))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create credential: %w", err)
|
||||
}
|
||||
|
|
@ -88,25 +90,51 @@ func resourceLiteLLMCredentialCreate(d *schema.ResourceData, m interface{}) erro
|
|||
|
||||
err = handleCredentialAPIResponse(resp, nil, client)
|
||||
if err != nil {
|
||||
if errors.Is(err, errCredentialConflict) {
|
||||
return handleCredentialNameConflict(d, m, credentialName)
|
||||
}
|
||||
return fmt.Errorf("failed to create credential: %w", err)
|
||||
}
|
||||
|
||||
// Set the resource ID to the credential name
|
||||
d.SetId(credentialName)
|
||||
|
||||
log.Printf("[INFO] Credential created with name %s. Starting retry mechanism to read the credential...", credentialName)
|
||||
return retryCredentialRead(d, m, 5)
|
||||
}
|
||||
|
||||
func handleCredentialNameConflict(d *schema.ResourceData, m interface{}, credentialName string) error {
|
||||
if !d.Get("adopt_existing").(bool) {
|
||||
return fmt.Errorf(
|
||||
"credential %q already exists on the proxy but is not in Terraform state. "+
|
||||
"Import it to manage it here:\n\n"+
|
||||
" terraform import litellm_credential.<this resource's name in your config> %s\n\n"+
|
||||
"The next apply then updates it to match this configuration. To take it over during "+
|
||||
"create instead, set adopt_existing = true on this resource, which overwrites the "+
|
||||
"existing credential's values with the ones configured here",
|
||||
credentialName, shellSingleQuote(credentialName),
|
||||
)
|
||||
}
|
||||
|
||||
log.Printf("[WARN] Credential %q already exists; adopt_existing is set, so taking it over and updating it to match configuration.", credentialName)
|
||||
d.SetId(credentialName)
|
||||
if err := patchCredential(m.(*Client), d, credentialName); err != nil {
|
||||
d.SetId("")
|
||||
return fmt.Errorf("failed to adopt existing credential %q: %w", credentialName, err)
|
||||
}
|
||||
return retryCredentialRead(d, m, 5)
|
||||
}
|
||||
|
||||
func shellSingleQuote(s string) string {
|
||||
return "'" + strings.ReplaceAll(s, "'", `'\''`) + "'"
|
||||
}
|
||||
|
||||
func resourceLiteLLMCredentialRead(d *schema.ResourceData, m interface{}) error {
|
||||
client := m.(*Client)
|
||||
credentialName := d.Id()
|
||||
|
||||
// Try to get credential by name first
|
||||
modelID := d.Get("model_id").(string)
|
||||
endpoint := fmt.Sprintf("/credentials/by_name/%s", credentialName)
|
||||
if modelID != "" {
|
||||
endpoint += fmt.Sprintf("?model_id=%s", modelID)
|
||||
endpoint := fmt.Sprintf(endpointCredentialByName, url.PathEscape(credentialName))
|
||||
if modelID := d.Get("model_id").(string); modelID != "" {
|
||||
endpoint = fmt.Sprintf(endpointCredentialByNameForModel, url.PathEscape(credentialName), url.QueryEscape(modelID))
|
||||
}
|
||||
|
||||
resp, err := MakeRequest(client, "GET", endpoint, nil)
|
||||
|
|
@ -138,42 +166,28 @@ func resourceLiteLLMCredentialRead(d *schema.ResourceData, m interface{}) error
|
|||
return nil
|
||||
}
|
||||
|
||||
func resourceLiteLLMCredentialUpdate(d *schema.ResourceData, m interface{}) error {
|
||||
client := m.(*Client)
|
||||
credentialName := d.Id()
|
||||
|
||||
credentialInfo := d.Get("credential_info").(map[string]interface{})
|
||||
credentialValues := d.Get("credential_values").(map[string]interface{})
|
||||
|
||||
// Convert credential_info to map[string]interface{} for JSON
|
||||
credInfoMap := make(map[string]interface{})
|
||||
for k, v := range credentialInfo {
|
||||
credInfoMap[k] = v
|
||||
}
|
||||
|
||||
// Convert credential_values to map[string]interface{} for JSON
|
||||
credValuesMap := make(map[string]interface{})
|
||||
for k, v := range credentialValues {
|
||||
credValuesMap[k] = v
|
||||
}
|
||||
|
||||
credentialRequest := CredentialRequest{
|
||||
CredentialName: credentialName,
|
||||
CredentialInfo: credInfoMap,
|
||||
CredentialValues: credValuesMap,
|
||||
}
|
||||
|
||||
endpoint := fmt.Sprintf("/credentials/%s", credentialName)
|
||||
resp, err := MakeRequest(client, "PATCH", endpoint, credentialRequest)
|
||||
func patchCredential(client *Client, d *schema.ResourceData, credentialName string) error {
|
||||
resp, err := MakeRequest(client, "PATCH", fmt.Sprintf(endpointCredential, url.PathEscape(credentialName)), credentialRequestFromResource(d, credentialName))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update credential: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
err = handleCredentialAPIResponse(resp, nil, client)
|
||||
if err != nil {
|
||||
if err := handleCredentialAPIResponse(resp, nil, client); err != nil {
|
||||
return fmt.Errorf("failed to update credential: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func resourceLiteLLMCredentialUpdate(d *schema.ResourceData, m interface{}) error {
|
||||
if !d.HasChangesExcept("adopt_existing") {
|
||||
return nil
|
||||
}
|
||||
|
||||
credentialName := d.Id()
|
||||
if err := patchCredential(m.(*Client), d, credentialName); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
log.Printf("[INFO] Credential updated with name %s. Starting retry mechanism to read the credential...", credentialName)
|
||||
return retryCredentialRead(d, m, 5)
|
||||
|
|
@ -183,8 +197,7 @@ func resourceLiteLLMCredentialDelete(d *schema.ResourceData, m interface{}) erro
|
|||
client := m.(*Client)
|
||||
credentialName := d.Id()
|
||||
|
||||
endpoint := fmt.Sprintf("/credentials/%s", credentialName)
|
||||
resp, err := MakeRequest(client, "DELETE", endpoint, nil)
|
||||
resp, err := MakeRequest(client, "DELETE", fmt.Sprintf(endpointCredential, url.PathEscape(credentialName)), nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete credential: %w", err)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,14 +1,18 @@
|
|||
package litellm
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/helper/schema"
|
||||
"github.com/hashicorp/terraform-plugin-sdk/v2/terraform"
|
||||
)
|
||||
|
||||
// newTestResourceData creates a *schema.ResourceData with the credential schema,
|
||||
|
|
@ -199,3 +203,394 @@ func TestRetryCredentialRead_ConnectionError(t *testing.T) {
|
|||
// Connection error should not be retried (not a "credential_not_found")
|
||||
fmt.Printf("connection error (expected): %v\n", err)
|
||||
}
|
||||
|
||||
type conflictBody struct {
|
||||
status int
|
||||
body string
|
||||
}
|
||||
|
||||
var (
|
||||
modernConflictBody = conflictBody{
|
||||
status: http.StatusConflict,
|
||||
body: `{"error":{"message":"Credential 'conflict-test' already exists. Update it with PATCH /credentials/conflict-test, or delete it first.","type":"internal_server_error","param":"None","code":"409"}}`,
|
||||
}
|
||||
legacyConflictBody = conflictBody{
|
||||
status: http.StatusInternalServerError,
|
||||
body: `{"error":{"message":"Unique constraint failed on the fields: (` + "`credential_name`" + `)","type":"internal_server_error","code":"500"}}`,
|
||||
}
|
||||
)
|
||||
|
||||
type conflictServerOptions struct {
|
||||
conflict conflictBody
|
||||
patchStatus int
|
||||
patchBody string
|
||||
getStatus int
|
||||
}
|
||||
|
||||
func conflictServer(t *testing.T, opts conflictServerOptions) (*httptest.Server, *int32, *int32, *[]byte) {
|
||||
t.Helper()
|
||||
var createCalls, patchCalls int32
|
||||
var capturedPatchBody []byte
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch {
|
||||
case r.Method == http.MethodPost && r.URL.Path == "/credentials":
|
||||
atomic.AddInt32(&createCalls, 1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(opts.conflict.status)
|
||||
w.Write([]byte(opts.conflict.body))
|
||||
case r.Method == http.MethodPatch:
|
||||
atomic.AddInt32(&patchCalls, 1)
|
||||
if r.URL.Path != "/credentials/conflict-test" {
|
||||
t.Errorf("PATCH went to %q, want /credentials/conflict-test", r.URL.Path)
|
||||
}
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
capturedPatchBody = body
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(opts.patchStatus)
|
||||
w.Write([]byte(opts.patchBody))
|
||||
case r.Method == http.MethodGet:
|
||||
if r.URL.Path != "/credentials/by_name/conflict-test" || r.URL.Query().Get("model_id") != "model-1" {
|
||||
t.Errorf("GET went to %q (query %q), want /credentials/by_name/conflict-test?model_id=model-1", r.URL.Path, r.URL.RawQuery)
|
||||
}
|
||||
if opts.getStatus != 0 && opts.getStatus != http.StatusOK {
|
||||
w.WriteHeader(opts.getStatus)
|
||||
w.Write([]byte(`{"error":{"message":"Internal Server Error"}}`))
|
||||
return
|
||||
}
|
||||
resp := CredentialResponse{CredentialName: "conflict-test", CredentialInfo: map[string]interface{}{}}
|
||||
body, _ := json.Marshal(resp)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write(body)
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
return srv, &createCalls, &patchCalls, &capturedPatchBody
|
||||
}
|
||||
|
||||
func adoptTestData(t *testing.T, adoptExisting bool) *schema.ResourceData {
|
||||
t.Helper()
|
||||
return schema.TestResourceDataRaw(t, resourceLiteLLMCredential().Schema, map[string]interface{}{
|
||||
"credential_name": "conflict-test",
|
||||
"model_id": "model-1",
|
||||
"credential_info": map[string]interface{}{"custom_llm_provider": "bedrock"},
|
||||
"credential_values": map[string]interface{}{"aws_access_key_id": "val"},
|
||||
"adopt_existing": adoptExisting,
|
||||
})
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMCredentialCreate_AdoptsOnConflictWhenOptedIn(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
conflict conflictBody
|
||||
}{
|
||||
{"typed 409", modernConflictBody},
|
||||
{"legacy 500 with unique-constraint message", legacyConflictBody},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
srv, createCalls, patchCalls, patchBody := conflictServer(t, conflictServerOptions{conflict: tc.conflict, patchStatus: http.StatusOK, patchBody: `{}`})
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := adoptTestData(t, true)
|
||||
|
||||
if err := resourceLiteLLMCredentialCreate(d, client); err != nil {
|
||||
t.Fatalf("expected create to adopt the existing credential, got error: %v", err)
|
||||
}
|
||||
if d.Id() != "conflict-test" {
|
||||
t.Fatalf("expected ID %q, got %q", "conflict-test", d.Id())
|
||||
}
|
||||
if got := atomic.LoadInt32(createCalls); got != 1 {
|
||||
t.Fatalf("expected exactly 1 POST /credentials call, got %d", got)
|
||||
}
|
||||
if got := atomic.LoadInt32(patchCalls); got != 1 {
|
||||
t.Fatalf("expected the conflict to trigger exactly 1 PATCH (adopt-and-update), got %d", got)
|
||||
}
|
||||
|
||||
var sent map[string]interface{}
|
||||
if err := json.Unmarshal(*patchBody, &sent); err != nil {
|
||||
t.Fatalf("PATCH body was not valid JSON: %v (%s)", err, *patchBody)
|
||||
}
|
||||
if sent["credential_name"] != "conflict-test" {
|
||||
t.Errorf("PATCH body credential_name = %v, want conflict-test", sent["credential_name"])
|
||||
}
|
||||
if sent["model_id"] != "model-1" {
|
||||
t.Errorf("PATCH body model_id = %v, want model-1 (adoption must not drop model-based credential resolution)", sent["model_id"])
|
||||
}
|
||||
credInfo, _ := sent["credential_info"].(map[string]interface{})
|
||||
if credInfo["custom_llm_provider"] != "bedrock" {
|
||||
t.Errorf("PATCH body credential_info = %v, want custom_llm_provider=bedrock", sent["credential_info"])
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMCredentialCreate_ConflictWithoutOptInFailsWithImportHint(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
conflict conflictBody
|
||||
}{
|
||||
{"typed 409", modernConflictBody},
|
||||
{"legacy 500 with unique-constraint message", legacyConflictBody},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
srv, createCalls, patchCalls, _ := conflictServer(t, conflictServerOptions{conflict: tc.conflict, patchStatus: http.StatusOK, patchBody: `{}`})
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := adoptTestData(t, false)
|
||||
|
||||
err := resourceLiteLLMCredentialCreate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected create to fail on the conflict when adopt_existing is unset, got nil")
|
||||
}
|
||||
if got := atomic.LoadInt32(createCalls); got != 1 {
|
||||
t.Fatalf("expected exactly 1 POST /credentials call, got %d", got)
|
||||
}
|
||||
if got := atomic.LoadInt32(patchCalls); got != 0 {
|
||||
t.Fatalf("expected no PATCH without adopt_existing - create must not overwrite an unmanaged credential - got %d", got)
|
||||
}
|
||||
if d.Id() != "" {
|
||||
t.Fatalf("resource ID must stay empty when create refuses the conflict, got %q", d.Id())
|
||||
}
|
||||
for _, want := range []string{
|
||||
"already exists",
|
||||
`terraform import litellm_credential.<this resource's name in your config> 'conflict-test'`,
|
||||
"adopt_existing = true",
|
||||
} {
|
||||
if !strings.Contains(err.Error(), want) {
|
||||
t.Errorf("error must tell the operator how to proceed; missing %q in: %v", want, err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMCredentialCreate_FailedAdoptDoesNotTaint(t *testing.T) {
|
||||
srv, createCalls, patchCalls, _ := conflictServer(t, conflictServerOptions{
|
||||
conflict: modernConflictBody,
|
||||
patchStatus: http.StatusInternalServerError,
|
||||
patchBody: `{"error":{"message":"Internal Server Error"}}`,
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := adoptTestData(t, true)
|
||||
|
||||
err := resourceLiteLLMCredentialCreate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected an error when the adopt PATCH fails, got nil")
|
||||
}
|
||||
if got := atomic.LoadInt32(createCalls); got != 1 {
|
||||
t.Fatalf("expected exactly 1 POST /credentials call, got %d", got)
|
||||
}
|
||||
if got := atomic.LoadInt32(patchCalls); got != 1 {
|
||||
t.Fatalf("expected exactly 1 PATCH attempt, got %d", got)
|
||||
}
|
||||
if d.Id() != "" {
|
||||
t.Fatalf("resource ID must stay empty after a failed adopt, got %q (a tainted entry would be destroyed on the next apply)", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMCredentialCreate_NonConflictErrorDoesNotAdopt(t *testing.T) {
|
||||
var createCalls, patchCalls int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch {
|
||||
case r.Method == http.MethodPost && r.URL.Path == "/credentials":
|
||||
atomic.AddInt32(&createCalls, 1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
w.Write([]byte(`{"error":{"message":"Internal Server Error","type":"internal_server_error"}}`))
|
||||
case r.Method == http.MethodPatch:
|
||||
atomic.AddInt32(&patchCalls, 1)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte(`{}`))
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMCredential().Schema, map[string]interface{}{
|
||||
"credential_name": "some-cred",
|
||||
"credential_info": map[string]interface{}{},
|
||||
"credential_values": map[string]interface{}{"key": "val"},
|
||||
"adopt_existing": true,
|
||||
})
|
||||
|
||||
err := resourceLiteLLMCredentialCreate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected an error for a non-conflict failure, got nil")
|
||||
}
|
||||
if got := atomic.LoadInt32(&patchCalls); got != 0 {
|
||||
t.Fatalf("expected no PATCH attempt for a non-conflict error, got %d", got)
|
||||
}
|
||||
if d.Id() != "" {
|
||||
t.Fatalf("resource ID must stay empty on a non-conflict failure, got %q", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMCredentialCreate_AdoptKeepsIDWhenPostPatchReadFails(t *testing.T) {
|
||||
srv, _, patchCalls, _ := conflictServer(t, conflictServerOptions{
|
||||
conflict: modernConflictBody,
|
||||
patchStatus: http.StatusOK,
|
||||
patchBody: `{}`,
|
||||
getStatus: http.StatusInternalServerError,
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := adoptTestData(t, true)
|
||||
|
||||
err := resourceLiteLLMCredentialCreate(d, client)
|
||||
if err == nil {
|
||||
t.Fatal("expected the failed post-adopt read to surface as an error, got nil")
|
||||
}
|
||||
if got := atomic.LoadInt32(patchCalls); got != 1 {
|
||||
t.Fatalf("expected exactly 1 PATCH, got %d", got)
|
||||
}
|
||||
if d.Id() != "conflict-test" {
|
||||
t.Fatalf("the PATCH already overwrote the remote credential, so the ID must stay set for Terraform to track it; got %q", d.Id())
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMCredentialImportHintQuotesTheNameForTheShell(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
want string
|
||||
}{
|
||||
{"my cred", `'my cred'`},
|
||||
{"it's $HOME `id` \"x\"", `'it'\''s $HOME ` + "`id`" + ` "x"'`},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusConflict)
|
||||
w.Write([]byte(`{"error":{"message":"already exists","code":"409"}}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMCredential().Schema, map[string]interface{}{
|
||||
"credential_name": tc.name,
|
||||
"credential_info": map[string]interface{}{},
|
||||
"credential_values": map[string]interface{}{"key": "val"},
|
||||
})
|
||||
|
||||
err := resourceLiteLLMCredentialCreate(d, NewClient(srv.URL, "test-key", true))
|
||||
if err == nil {
|
||||
t.Fatal("expected the conflict to fail create, got nil")
|
||||
}
|
||||
want := "terraform import litellm_credential.<this resource's name in your config> " + tc.want
|
||||
if !strings.Contains(err.Error(), want) {
|
||||
t.Fatalf("import hint must single-quote the name for the shell; missing %q in: %v", want, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCredentialRequestsEscapeReservedCharactersInTheName(t *testing.T) {
|
||||
const name = "team/a?b c"
|
||||
var paths []string
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
paths = append(paths, r.Method+" "+r.URL.EscapedPath()+"?"+r.URL.RawQuery)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte(`{"credential_name":"` + name + `","credential_info":{}}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
client := NewClient(srv.URL, "test-key", true)
|
||||
d := schema.TestResourceDataRaw(t, resourceLiteLLMCredential().Schema, map[string]interface{}{
|
||||
"credential_name": name,
|
||||
"model_id": "m&1",
|
||||
"credential_info": map[string]interface{}{},
|
||||
"credential_values": map[string]interface{}{"key": "val"},
|
||||
})
|
||||
d.SetId(name)
|
||||
|
||||
if err := resourceLiteLLMCredentialRead(d, client); err != nil {
|
||||
t.Fatalf("read failed: %v", err)
|
||||
}
|
||||
if err := patchCredential(client, d, name); err != nil {
|
||||
t.Fatalf("patch failed: %v", err)
|
||||
}
|
||||
if err := resourceLiteLLMCredentialDelete(d, client); err != nil {
|
||||
t.Fatalf("delete failed: %v", err)
|
||||
}
|
||||
|
||||
want := []string{
|
||||
"GET /credentials/by_name/team%2Fa%3Fb%20c?model_id=m%261",
|
||||
"PATCH /credentials/team%2Fa%3Fb%20c?",
|
||||
"DELETE /credentials/team%2Fa%3Fb%20c?",
|
||||
}
|
||||
if strings.Join(paths, "\n") != strings.Join(want, "\n") {
|
||||
t.Fatalf("request paths:\n%s\nwant:\n%s", strings.Join(paths, "\n"), strings.Join(want, "\n"))
|
||||
}
|
||||
}
|
||||
|
||||
func TestResourceLiteLLMCredentialUpdate_TogglingAdoptExistingSendsNoPatch(t *testing.T) {
|
||||
var patchCalls int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodPatch {
|
||||
atomic.AddInt32(&patchCalls, 1)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write([]byte(`{"credential_name":"cred-1","credential_info":{}}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
res := resourceLiteLLMCredential()
|
||||
priorData := schema.TestResourceDataRaw(t, res.Schema, map[string]interface{}{
|
||||
"credential_name": "cred-1",
|
||||
"credential_info": map[string]interface{}{},
|
||||
"credential_values": map[string]interface{}{"api_key": "sk-secret"},
|
||||
"adopt_existing": false,
|
||||
})
|
||||
priorData.SetId("cred-1")
|
||||
prior := priorData.State()
|
||||
|
||||
toggled := terraform.NewResourceConfigRaw(map[string]interface{}{
|
||||
"credential_name": "cred-1",
|
||||
"credential_info": map[string]interface{}{},
|
||||
"credential_values": map[string]interface{}{"api_key": "sk-secret"},
|
||||
"adopt_existing": true,
|
||||
})
|
||||
diff, err := res.Diff(context.Background(), prior, toggled, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("diff failed: %v", err)
|
||||
}
|
||||
d, err := schema.InternalMap(res.Schema).Data(prior, diff)
|
||||
if err != nil {
|
||||
t.Fatalf("data failed: %v", err)
|
||||
}
|
||||
if err := resourceLiteLLMCredentialUpdate(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("update failed: %v", err)
|
||||
}
|
||||
if got := atomic.LoadInt32(&patchCalls); got != 0 {
|
||||
t.Fatalf("flipping adopt_existing alone must not rewrite the credential's secrets; got %d PATCH calls", got)
|
||||
}
|
||||
|
||||
rotated := terraform.NewResourceConfigRaw(map[string]interface{}{
|
||||
"credential_name": "cred-1",
|
||||
"credential_info": map[string]interface{}{},
|
||||
"credential_values": map[string]interface{}{"api_key": "sk-rotated"},
|
||||
"adopt_existing": true,
|
||||
})
|
||||
diff, err = res.Diff(context.Background(), prior, rotated, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("diff failed: %v", err)
|
||||
}
|
||||
d, err = schema.InternalMap(res.Schema).Data(prior, diff)
|
||||
if err != nil {
|
||||
t.Fatalf("data failed: %v", err)
|
||||
}
|
||||
if err := resourceLiteLLMCredentialUpdate(d, NewClient(srv.URL, "test-key", true)); err != nil {
|
||||
t.Fatalf("update failed: %v", err)
|
||||
}
|
||||
if got := atomic.LoadInt32(&patchCalls); got != 1 {
|
||||
t.Fatalf("a real value change must still PATCH; got %d PATCH calls", got)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import (
|
|||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
|
|
@ -202,6 +203,23 @@ func isCredentialNotFoundError(errResp ErrorResponse) bool {
|
|||
return false
|
||||
}
|
||||
|
||||
var errCredentialConflict = errors.New("credential_conflict")
|
||||
|
||||
func isLegacyCredentialConflictError(errResp ErrorResponse) bool {
|
||||
isConflict := func(msg string) bool {
|
||||
return strings.Contains(msg, "Unique constraint failed") && strings.Contains(msg, "credential_name")
|
||||
}
|
||||
if msg, ok := errResp.Error.Message.(string); ok && isConflict(msg) {
|
||||
return true
|
||||
}
|
||||
if msgMap, ok := errResp.Error.Message.(map[string]interface{}); ok {
|
||||
if errStr, ok := msgMap["error"].(string); ok && isConflict(errStr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return isConflict(errResp.Detail.Error)
|
||||
}
|
||||
|
||||
// handleCredentialAPIResponse handles API responses specifically for credential operations
|
||||
func handleCredentialAPIResponse(resp *http.Response, result interface{}, client *Client) error {
|
||||
bodyBytes, err := io.ReadAll(resp.Body)
|
||||
|
|
@ -213,12 +231,19 @@ func handleCredentialAPIResponse(resp *http.Response, result interface{}, client
|
|||
return fmt.Errorf("credential_not_found")
|
||||
}
|
||||
|
||||
if resp.StatusCode == http.StatusConflict {
|
||||
return errCredentialConflict
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusCreated {
|
||||
var errResp ErrorResponse
|
||||
if err := json.Unmarshal(bodyBytes, &errResp); err == nil {
|
||||
if isCredentialNotFoundError(errResp) {
|
||||
return fmt.Errorf("credential_not_found")
|
||||
}
|
||||
if isLegacyCredentialConflictError(errResp) {
|
||||
return errCredentialConflict
|
||||
}
|
||||
}
|
||||
return fmt.Errorf("API request failed: Status: %s, Response: %s",
|
||||
resp.Status, client.redactSensitiveData(string(bodyBytes)))
|
||||
|
|
|
|||
|
|
@ -170,6 +170,7 @@ class StreamingResponse(BaseModel):
|
|||
# the consumed body is elided, so this is the only place they surface.
|
||||
stream_error: str | None = None
|
||||
stream_done: bool = False
|
||||
stream_done_positions: tuple[int, ...] = ()
|
||||
|
||||
@property
|
||||
def ok(self) -> bool:
|
||||
|
|
@ -647,6 +648,7 @@ def streaming_outcome(
|
|||
stream_events=[payload for payload, _ in events],
|
||||
stream_event_arrivals=[arrived for _, arrived in events],
|
||||
stream_done=any(payload == _SSE_DONE for payload, _ in payloads),
|
||||
stream_done_positions=tuple(index for index, (payload, _) in enumerate(payloads) if payload == _SSE_DONE),
|
||||
stream_error=next(
|
||||
(line.decode(errors="replace")[:300] for line, _ in stamped if _is_stream_error_line(line)),
|
||||
None,
|
||||
|
|
|
|||
|
|
@ -1,51 +1,90 @@
|
|||
"""Vendor §12.3: chat completions streaming SSE contract (LIT-4778).
|
||||
|
||||
Asserts a streamed /chat/completions response is SSE, carries content chunks,
|
||||
and terminates with the OpenAI [DONE] sentinel.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from e2e_config import unique_marker
|
||||
from e2e_config import provider_edge_base, unique_marker
|
||||
from e2e_http import require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatBody, ChatMessage, LiteLLMParamsBody
|
||||
from models import ChatBody, ChatMessage, ChatStreamOptions, LiteLLMParamsBody, Usage
|
||||
from proxy_client import ProxyClient
|
||||
from pydantic import BaseModel
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
pytestmark = [pytest.mark.e2e, pytest.mark.replayable]
|
||||
|
||||
|
||||
class _Delta(BaseModel):
|
||||
content: str | None = None
|
||||
|
||||
|
||||
class _Choice(BaseModel):
|
||||
index: int
|
||||
delta: _Delta
|
||||
finish_reason: str | None = None
|
||||
|
||||
|
||||
class _Chunk(BaseModel):
|
||||
choices: tuple[_Choice, ...]
|
||||
usage: Usage | None = None
|
||||
|
||||
|
||||
class TestChatStreamContract:
|
||||
@pytest.mark.covers("llm.chat_completions.openai.basic.stream.works")
|
||||
def test_chat_stream_is_sse_and_ends_with_done(self, proxy: ProxyClient, resources: ResourceManager) -> None:
|
||||
model = f"e2e-chat-stream-{unique_marker()}"
|
||||
model_id = proxy.create_model(
|
||||
model: Final = f"e2e-chat-stream-{unique_marker()}"
|
||||
base: Final = provider_edge_base("openai")
|
||||
model_id: Final = proxy.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(model="openai/gpt-4o-mini", api_key="os.environ/OPENAI_API_KEY"),
|
||||
LiteLLMParamsBody(
|
||||
model="openai/gpt-5.6",
|
||||
api_key="os.environ/OPENAI_API_KEY",
|
||||
api_base=f"{base}/v1" if base else None,
|
||||
),
|
||||
)
|
||||
resources.defer(lambda: proxy.delete_model(model_id))
|
||||
key = resources.key()
|
||||
|
||||
result = proxy.chat_stream(
|
||||
key: Final = resources.key()
|
||||
expected: Final = "The amber kite crosses the quiet lake."
|
||||
result: Final = proxy.chat_stream(
|
||||
key,
|
||||
ChatBody(
|
||||
model=model,
|
||||
messages=[
|
||||
ChatMessage(
|
||||
role="user",
|
||||
content=f"Reply with the single word ok. {unique_marker()}",
|
||||
role="user", content=f"Repeat exactly this sentence, with no additional text: {expected}"
|
||||
)
|
||||
],
|
||||
stream=True,
|
||||
max_completion_tokens=32,
|
||||
temperature=0.0,
|
||||
stream_options=ChatStreamOptions(include_usage=True),
|
||||
max_completion_tokens=256,
|
||||
reasoning_effort="none",
|
||||
),
|
||||
)
|
||||
require_successful_call(result)
|
||||
assert result.is_streaming, f"expected SSE content-type, got {result.content_type!r}"
|
||||
assert result.stream_events, "stream returned no data events"
|
||||
assert result.stream_done, (
|
||||
f"stream must terminate with [DONE]; "
|
||||
f"chunks={result.chunks} done={result.stream_done} events={len(result.stream_events)}"
|
||||
assert not result.stream_error, f"stream errored: {result.stream_error}"
|
||||
assert result.stream_done, "stream must terminate with [DONE]"
|
||||
assert result.stream_done_positions == (len(result.stream_events),), "[DONE] must occur once after all events"
|
||||
chunks: Final = tuple(_Chunk.model_validate_json(event) for event in result.stream_events)
|
||||
text_positions: Final = tuple(
|
||||
i for i, chunk in enumerate(chunks) if any(c.delta.content for c in chunk.choices)
|
||||
)
|
||||
terminal_positions: Final = tuple(
|
||||
i for i, chunk in enumerate(chunks) if any(c.finish_reason is not None for c in chunk.choices)
|
||||
)
|
||||
assert text_positions, "stream completed without meaningful text"
|
||||
assert len(terminal_positions) == 1, "expected exactly one terminal choice"
|
||||
assert text_positions[0] < terminal_positions[0], "meaningful text must arrive before termination"
|
||||
assert text_positions[-1] <= terminal_positions[0], "text arrived after termination"
|
||||
assert all(c.index == 0 for chunk in chunks for c in chunk.choices)
|
||||
assert tuple(c.finish_reason for c in chunks[terminal_positions[0]].choices) == ("stop",)
|
||||
text: Final = "".join(c.delta.content or "" for chunk in chunks for c in chunk.choices)
|
||||
assert text.strip() == expected, f"streamed answer was altered or incomplete: {text!r}"
|
||||
usage_positions: Final = tuple(i for i, chunk in enumerate(chunks) if chunk.usage is not None)
|
||||
assert usage_positions == (len(chunks) - 1,), "expected one final usage chunk"
|
||||
assert terminal_positions[0] < usage_positions[0], "usage must follow the terminal choice"
|
||||
usage: Final = chunks[-1].usage
|
||||
assert usage is not None
|
||||
assert usage.prompt_tokens is not None and usage.prompt_tokens > 0
|
||||
assert usage.completion_tokens is not None and usage.completion_tokens > 0
|
||||
assert usage.total_tokens == usage.prompt_tokens + usage.completion_tokens
|
||||
|
|
|
|||
|
|
@ -21,7 +21,12 @@ from e2e_http import assert_client_error, require_successful_call, unwrap
|
|||
from endpoints_client import EndpointsClient, MessagesResult
|
||||
from lifecycle import ResourceManager
|
||||
from models import (
|
||||
AnthropicAssistantTurn,
|
||||
AnthropicContentBlock,
|
||||
AnthropicCustomTool,
|
||||
AnthropicToolChoice,
|
||||
AnthropicToolResultBlock,
|
||||
AnthropicToolResultTurn,
|
||||
AnthropicMessagesBody,
|
||||
ChatMessage,
|
||||
JsonSchemaProperty,
|
||||
|
|
@ -29,7 +34,7 @@ from models import (
|
|||
SpendLogRow,
|
||||
ToolInputSchema,
|
||||
)
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
pytestmark = [pytest.mark.e2e, pytest.mark.replayable]
|
||||
|
||||
|
|
@ -284,8 +289,139 @@ class TestAnthropicMessages:
|
|||
result = endpoints_client.proxy.transport.send(
|
||||
"/v1/messages",
|
||||
headers=endpoints_client.proxy.transport.bearer(key),
|
||||
json=_OptionalMessagesBody(
|
||||
messages=[ChatMessage(role="user", content="hi")], max_tokens=50
|
||||
),
|
||||
json=_OptionalMessagesBody(messages=[ChatMessage(role="user", content="hi")], max_tokens=50),
|
||||
)
|
||||
assert_client_error(result, "messages missing model")
|
||||
|
||||
|
||||
class _BridgeDelta(BaseModel):
|
||||
type: str | None = None
|
||||
partial_json: str | None = None
|
||||
stop_reason: str | None = None
|
||||
|
||||
|
||||
class _BridgeEvent(BaseModel):
|
||||
type: str
|
||||
index: int | None = None
|
||||
content_block: AnthropicContentBlock | None = None
|
||||
delta: _BridgeDelta | None = None
|
||||
|
||||
|
||||
class _ParcelInput(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid", strict=True)
|
||||
parcel: str
|
||||
shelf: int
|
||||
|
||||
|
||||
def _tool_from_stream(events: tuple[_BridgeEvent, ...]) -> AnthropicContentBlock:
|
||||
starts: Final = tuple(
|
||||
event
|
||||
for event in events
|
||||
if event.type == "content_block_start"
|
||||
and event.content_block is not None
|
||||
and event.content_block.type == "tool_use"
|
||||
)
|
||||
assert len(starts) == 1, "expected exactly one tool call"
|
||||
start: Final = starts[0]
|
||||
block: Final = start.content_block
|
||||
assert block is not None and block.id and start.index is not None
|
||||
fragments: Final = tuple(
|
||||
event
|
||||
for event in events
|
||||
if event.type == "content_block_delta" and event.delta is not None and event.delta.type == "input_json_delta"
|
||||
)
|
||||
assert fragments, "tool stream contained no argument fragments"
|
||||
assert all(event.index == start.index for event in fragments), "tool fragments changed index"
|
||||
positions: Final = tuple(i for i, event in enumerate(events) if event in fragments)
|
||||
stops: Final = tuple(
|
||||
i for i, event in enumerate(events) if event.type == "content_block_stop" and event.index == start.index
|
||||
)
|
||||
assert len(stops) == 1 and events.index(start) < positions[0] <= positions[-1] < stops[0]
|
||||
assert tuple(
|
||||
event.delta.stop_reason for event in events if event.type == "message_delta" and event.delta is not None
|
||||
) == ("tool_use",)
|
||||
terminal_positions: Final = tuple(i for i, event in enumerate(events) if event.type == "message_delta")
|
||||
assert len(terminal_positions) == 1 and stops[0] < terminal_positions[0] < len(events) - 1
|
||||
assert tuple(i for i, event in enumerate(events) if event.type == "message_stop") == (len(events) - 1,), (
|
||||
"tool stream did not terminate exactly once"
|
||||
)
|
||||
arguments: Final = _ParcelInput.model_validate_json(
|
||||
"".join(event.delta.partial_json or "" for event in fragments if event.delta is not None)
|
||||
)
|
||||
return AnthropicContentBlock(type="tool_use", id=block.id, name=block.name, input=arguments.model_dump())
|
||||
|
||||
|
||||
def _parcel_result(tool: AnthropicContentBlock, result: AnthropicToolResultBlock) -> AnthropicToolResultTurn:
|
||||
assert tool.id and result.tool_use_id == tool.id, "tool result ID does not match the emitted call"
|
||||
return AnthropicToolResultTurn(content=[result])
|
||||
|
||||
|
||||
def _request_tool(
|
||||
client: EndpointsClient, key: str, request: AnthropicMessagesBody, stream: bool
|
||||
) -> AnthropicContentBlock:
|
||||
if stream:
|
||||
response: Final = client.proxy.messages_stream(key, request)
|
||||
require_successful_call(response)
|
||||
assert response.is_streaming and not response.stream_error
|
||||
return _tool_from_stream(tuple(_BridgeEvent.model_validate_json(event) for event in response.stream_events))
|
||||
response_body: Final = unwrap(client.proxy.messages(key, request))
|
||||
blocks: Final = tuple(block for block in response_body.content or () if block.type == "tool_use")
|
||||
assert len(blocks) == 1
|
||||
return blocks[0]
|
||||
|
||||
|
||||
class TestOpenAIMessagesToolContinuation:
|
||||
@pytest.mark.parametrize("stream", [True, False], ids=["stream", "nonstream"])
|
||||
def test_required_tool_arguments_and_correlated_result(
|
||||
self, endpoints_client: EndpointsClient, resources: ResourceManager, stream: bool
|
||||
) -> None:
|
||||
model: Final = f"e2e-bridge-tool-{unique_marker()}"
|
||||
base: Final = provider_edge_base("openai")
|
||||
model_id: Final = endpoints_client.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(
|
||||
model="openai/gpt-5.6", api_key="os.environ/OPENAI_API_KEY", api_base=f"{base}/v1" if base else None
|
||||
),
|
||||
)
|
||||
resources.defer(lambda: endpoints_client.delete_model(model_id))
|
||||
key: Final = resources.key(models=[model])
|
||||
tool: Final = AnthropicCustomTool(
|
||||
name="locate_parcel",
|
||||
description="Look up the receipt for a parcel on a shelf. Return the receipt verbatim.",
|
||||
input_schema=ToolInputSchema(
|
||||
properties={"parcel": JsonSchemaProperty(type="string"), "shelf": JsonSchemaProperty(type="integer")},
|
||||
required=["parcel", "shelf"],
|
||||
),
|
||||
)
|
||||
question: Final = ChatMessage(
|
||||
role="user",
|
||||
content="Call locate_parcel with parcel exactly amber-kite and shelf exactly 7. After the tool result, reply with only the receipt returned by the tool.",
|
||||
)
|
||||
request: Final = AnthropicMessagesBody(
|
||||
model=model,
|
||||
max_tokens=2048,
|
||||
messages=[question],
|
||||
tools=[tool],
|
||||
tool_choice=AnthropicToolChoice(type="tool", name=tool.name),
|
||||
stream=stream,
|
||||
)
|
||||
emitted: Final = _request_tool(endpoints_client, key, request, stream)
|
||||
assert emitted.id and emitted.name == "locate_parcel"
|
||||
assert emitted.input == {"parcel": "amber-kite", "shelf": 7}, "required tool arguments were lost or changed"
|
||||
receipt: Final = f"receipt-{unique_marker()}"
|
||||
result_turn: Final = _parcel_result(emitted, AnthropicToolResultBlock(tool_use_id=emitted.id, content=receipt))
|
||||
continuation: Final = unwrap(
|
||||
endpoints_client.proxy.messages(
|
||||
key,
|
||||
AnthropicMessagesBody(
|
||||
model=model,
|
||||
max_tokens=2048,
|
||||
tools=[tool],
|
||||
tool_choice=AnthropicToolChoice(type="none"),
|
||||
messages=[question, AnthropicAssistantTurn(content=[emitted]), result_turn],
|
||||
),
|
||||
)
|
||||
)
|
||||
answer: Final = "".join(block.text or "" for block in continuation.content or ())
|
||||
assert answer.strip() == receipt, "continuation did not consume the correlated tool result"
|
||||
assert all(block.type != "tool_use" for block in continuation.content or ())
|
||||
|
|
|
|||
|
|
@ -283,10 +283,15 @@ class ChatToolResultTurn(BaseModel):
|
|||
type ChatTurn = ChatMessage | ChatAssistantTurn | ChatToolResultTurn
|
||||
|
||||
|
||||
class ChatStreamOptions(BaseModel):
|
||||
include_usage: bool
|
||||
|
||||
|
||||
class ChatBody(BaseModel):
|
||||
model: str
|
||||
messages: Sequence[ChatTurn]
|
||||
stream: bool = False
|
||||
stream_options: ChatStreamOptions | None = None
|
||||
max_tokens: int | None = None
|
||||
max_completion_tokens: int | None = None
|
||||
temperature: float | None = None
|
||||
|
|
@ -488,12 +493,18 @@ class AnthropicToolResultTurn(BaseModel):
|
|||
type AnthropicMessage = ChatMessage | AnthropicAssistantTurn | AnthropicToolResultTurn
|
||||
|
||||
|
||||
class AnthropicToolChoice(BaseModel):
|
||||
type: Literal["auto", "any", "tool", "none"]
|
||||
name: str | None = None
|
||||
|
||||
|
||||
class AnthropicMessagesBody(BaseModel):
|
||||
model: str
|
||||
messages: list[AnthropicMessage]
|
||||
max_tokens: int
|
||||
stream: bool | None = None
|
||||
tools: list[AnthropicTool] | None = None
|
||||
tool_choice: AnthropicToolChoice | None = None
|
||||
guardrails: list[str] | None = None
|
||||
cache: dict[str, bool] | None = {"no-cache": True}
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,113 @@
|
|||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from math import isclose
|
||||
from typing import Final
|
||||
|
||||
from e2e_config import provider_edge_base, unique_marker
|
||||
from e2e_http import unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatBody, ChatMessage, ChatResponse, KeyGenerateBody, LiteLLMParamsBody, TeamNewBody
|
||||
from spend_e2e_client import SpendClient
|
||||
|
||||
INPUT_RATE: Final = 0.00004
|
||||
OUTPUT_RATE: Final = 0.00008
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TeamTraffic:
|
||||
team_id: str
|
||||
key: str
|
||||
responses: tuple[ChatResponse, ...]
|
||||
|
||||
@property
|
||||
def prompt_tokens(self) -> int:
|
||||
return sum(response.usage.prompt_tokens or 0 for response in self.responses if response.usage)
|
||||
|
||||
@property
|
||||
def completion_tokens(self) -> int:
|
||||
return sum(response.usage.completion_tokens or 0 for response in self.responses if response.usage)
|
||||
|
||||
@property
|
||||
def spend(self) -> float:
|
||||
return self.prompt_tokens * INPUT_RATE + self.completion_tokens * OUTPUT_RATE
|
||||
|
||||
|
||||
def create_traffic(client: SpendClient, resources: ResourceManager) -> tuple[TeamTraffic, ...]:
|
||||
base: Final = provider_edge_base("openai")
|
||||
model: Final = f"e2e-reconciliation-{unique_marker()}"
|
||||
model_id: Final = client.proxy.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(
|
||||
model="openai/gpt-5.6-luna",
|
||||
api_key="os.environ/OPENAI_API_KEY",
|
||||
api_base=None if base is None else f"{base}/v1",
|
||||
input_cost_per_token=INPUT_RATE,
|
||||
output_cost_per_token=OUTPUT_RATE,
|
||||
),
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_model(model_id))
|
||||
|
||||
def team_traffic() -> TeamTraffic:
|
||||
team: Final = client.proxy.create_team(TeamNewBody(team_alias=f"e2e-spend-{unique_marker()}"))
|
||||
resources.defer(lambda: client.proxy.delete_team(team))
|
||||
key: Final = client.proxy.generate_key(KeyGenerateBody(team_id=team, models=[model]))
|
||||
resources.defer(lambda: client.proxy.delete_key(key))
|
||||
|
||||
prompts: Final = tuple(f"Reply with one word. {index} {unique_marker()}" for index in range(7))
|
||||
|
||||
def call(index: int) -> ChatResponse:
|
||||
response: Final = unwrap(
|
||||
client.proxy.chat(
|
||||
key,
|
||||
ChatBody(
|
||||
model=model,
|
||||
messages=[ChatMessage(role="user", content=prompts[index])],
|
||||
max_completion_tokens=128,
|
||||
),
|
||||
)
|
||||
)
|
||||
assert response.id, "successful response must have an ID"
|
||||
assert response.usage is not None, "successful response must have usage"
|
||||
assert response.usage.prompt_tokens is not None and response.usage.prompt_tokens > 0
|
||||
assert response.usage.completion_tokens is not None and response.usage.completion_tokens > 0
|
||||
assert response.usage.total_tokens == response.usage.prompt_tokens + response.usage.completion_tokens
|
||||
assert not response.usage.cache_creation_input_tokens
|
||||
assert not response.usage.cache_read_input_tokens
|
||||
assert not response.usage.prompt_tokens_details or not response.usage.prompt_tokens_details.cached_tokens
|
||||
return response
|
||||
|
||||
sequential: Final = call(0)
|
||||
with ThreadPoolExecutor(max_workers=6) as pool:
|
||||
concurrent: Final = tuple(pool.map(call, range(1, 7)))
|
||||
return TeamTraffic(team, key, (sequential, *concurrent))
|
||||
|
||||
return tuple(team_traffic() for _ in range(2))
|
||||
|
||||
|
||||
def assert_logs_match(client: SpendClient, traffic: TeamTraffic) -> None:
|
||||
expected_ids: Final = frozenset(response.id for response in traffic.responses)
|
||||
assert len(expected_ids) == len(traffic.responses), "responses must have distinct IDs"
|
||||
rows: Final = client.poll_logs_for_key(
|
||||
traffic.key,
|
||||
min_rows=len(traffic.responses),
|
||||
predicate=lambda values: frozenset(row.request_id for row in values) == expected_ids,
|
||||
)
|
||||
assert frozenset(row.request_id for row in rows) == expected_ids, "stored IDs must equal returned response IDs"
|
||||
assert len(rows) == len(traffic.responses), "expected exactly one scoped spend row per response"
|
||||
by_id: Final = {row.request_id: row for row in rows}
|
||||
|
||||
def assert_response(response: ChatResponse) -> None:
|
||||
row: Final = by_id[response.id]
|
||||
usage: Final = response.usage
|
||||
assert usage is not None and usage.prompt_tokens is not None and usage.completion_tokens is not None
|
||||
assert row.team_id == traffic.team_id
|
||||
assert row.status == "success"
|
||||
assert row.cache_hit != "True"
|
||||
assert row.prompt_tokens == usage.prompt_tokens
|
||||
assert row.completion_tokens == usage.completion_tokens
|
||||
assert row.total_tokens == usage.total_tokens
|
||||
expected_cost: Final = usage.prompt_tokens * INPUT_RATE + usage.completion_tokens * OUTPUT_RATE
|
||||
assert row.spend is not None and isclose(row.spend, expected_cost, rel_tol=1e-6, abs_tol=1e-9)
|
||||
|
||||
for response in traffic.responses:
|
||||
assert_response(response)
|
||||
|
|
@ -17,13 +17,13 @@ fails the test; a pricing or token-count drift does not.
|
|||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from math import isclose
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from e2e_http import Result, Success
|
||||
from e2e_http import Success
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatResponse, LiteLLMParamsBody, SpendLogs, SpendLogsParams
|
||||
from models import LiteLLMParamsBody, SpendLogs, SpendLogsParams
|
||||
from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
|
@ -280,51 +280,22 @@ def test_key_spend_equals_sum_of_logs(client: SpendClient, scoped_key: str) -> N
|
|||
), f"key aggregate {key_spend} != sum of logs {logs_total}; rows: {_summarize(rows)}"
|
||||
|
||||
|
||||
@pytest.mark.replayable
|
||||
@pytest.mark.covers("quota_management.spend_tracking.concurrent_burst.loses_no_spend")
|
||||
def test_burst_of_concurrent_calls_loses_no_spend(
|
||||
client: SpendClient, scoped_key: str
|
||||
client: SpendClient, resources: ResourceManager
|
||||
) -> None:
|
||||
"""Six concurrent calls on one key: every call lands its own spend row under a
|
||||
distinct request_id and the key aggregate equals the sum of the rows.
|
||||
Sequential accuracy is covered by test_key_spend_equals_sum_of_logs; this pins
|
||||
the concurrent increment path (parallel writers racing on one key's counter),
|
||||
where a lost update can never be reproduced by sequential calls."""
|
||||
burst = 6
|
||||
from spend_reconciliation import TeamTraffic, assert_logs_match, create_traffic
|
||||
|
||||
def call(idx: int) -> Result[ChatResponse]:
|
||||
return client.chat(
|
||||
scoped_key,
|
||||
"gemini-2.5-flash",
|
||||
f"burst call {idx} {unique_marker()}",
|
||||
max_tokens=16,
|
||||
)
|
||||
traffic: Final = create_traffic(client, resources)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=burst) as pool:
|
||||
results = tuple(pool.map(call, range(burst)))
|
||||
failed = [r for r in results if not is_ok(r)]
|
||||
assert not failed, f"{len(failed)}/{burst} burst calls failed; first: {failed[0]}"
|
||||
def assert_team(team: TeamTraffic) -> None:
|
||||
assert_logs_match(client, team)
|
||||
key_spend: Final = client.poll_key_spend(team.key, minimum=team.spend * 0.999999)
|
||||
assert isclose(key_spend, team.spend, rel_tol=1e-6, abs_tol=1e-9)
|
||||
|
||||
rows = client.poll_logs_for_key(
|
||||
scoped_key,
|
||||
min_rows=burst,
|
||||
predicate=lambda rs: len([r for r in rs if (r.spend or 0) > 0]) >= burst,
|
||||
)
|
||||
costed = [r for r in rows if (r.spend or 0) > 0]
|
||||
assert len(costed) >= burst, (
|
||||
f"only {len(costed)}/{burst} burst calls produced a costed row - "
|
||||
f"rows lost under concurrency: {_summarize(rows)}"
|
||||
)
|
||||
request_ids = [r.request_id for r in costed]
|
||||
assert len(set(request_ids)) == len(request_ids), (
|
||||
f"concurrent rows collapsed onto shared request_ids: {_summarize(rows)}"
|
||||
)
|
||||
|
||||
logs_total = sum((r.spend or 0) for r in rows)
|
||||
key_spend = client.poll_key_spend(scoped_key, minimum=logs_total * 0.999)
|
||||
assert _approx_equal(key_spend, logs_total), (
|
||||
f"key aggregate {key_spend} != sum of {len(rows)} rows {logs_total} - "
|
||||
f"spend increments lost under concurrency: {_summarize(rows)}"
|
||||
)
|
||||
for team in traffic:
|
||||
assert_team(team)
|
||||
|
||||
|
||||
@pytest.mark.covers("quota_management.spend_tracking.pagination.keeps_total")
|
||||
|
|
|
|||
|
|
@ -7,13 +7,18 @@ missing start/end dates are rejected.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from math import isclose
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from e2e_http import ProbeResult
|
||||
from models import DateRangeParams
|
||||
from lifecycle import ResourceManager
|
||||
from proxy_client import Converged, await_converged
|
||||
from pydantic import BaseModel
|
||||
from spend_e2e_client import SpendClient
|
||||
from spend_reconciliation import TeamTraffic, assert_logs_match, create_traffic
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
|
@ -24,22 +29,45 @@ class TeamDailyActivityParams(BaseModel):
|
|||
start_date: str | None = None
|
||||
end_date: str | None = None
|
||||
page: int = 1
|
||||
page_size: int = 1
|
||||
team_ids: str | None = None
|
||||
|
||||
|
||||
class TeamDailyActivityRow(BaseModel):
|
||||
date: str
|
||||
metrics: TeamDailyActivityMetrics
|
||||
breakdown: TeamDailyActivityBreakdown
|
||||
|
||||
|
||||
class TeamDailyActivityMetrics(BaseModel):
|
||||
spend: float
|
||||
total_tokens: int
|
||||
prompt_tokens: int
|
||||
completion_tokens: int
|
||||
api_requests: int
|
||||
successful_requests: int
|
||||
failed_requests: int
|
||||
|
||||
|
||||
class TeamDailyActivityEntity(BaseModel):
|
||||
metrics: TeamDailyActivityMetrics
|
||||
|
||||
|
||||
class TeamDailyActivityBreakdown(BaseModel):
|
||||
entities: dict[str, TeamDailyActivityEntity]
|
||||
|
||||
|
||||
class TeamDailyActivityMetadata(BaseModel):
|
||||
page: int
|
||||
total_pages: int
|
||||
has_more: bool
|
||||
total_spend: float
|
||||
total_prompt_tokens: int
|
||||
total_completion_tokens: int
|
||||
total_tokens: int
|
||||
total_api_requests: int
|
||||
total_successful_requests: int
|
||||
total_failed_requests: int
|
||||
|
||||
|
||||
class TeamDailyActivityResponse(BaseModel):
|
||||
|
|
@ -47,32 +75,128 @@ class TeamDailyActivityResponse(BaseModel):
|
|||
metadata: TeamDailyActivityMetadata
|
||||
|
||||
|
||||
def _range_days(days: int) -> DateRangeParams:
|
||||
end = datetime.now(timezone.utc).date()
|
||||
start = end - timedelta(days=days)
|
||||
return DateRangeParams(start_date=start.isoformat(), end_date=end.isoformat())
|
||||
|
||||
|
||||
def _probe(client: SpendClient, params: BaseModel) -> ProbeResult:
|
||||
return client.proxy.transport.probe(ROUTE, params=params)
|
||||
|
||||
|
||||
class TestTeamDailyActivity:
|
||||
@pytest.mark.replayable
|
||||
@pytest.mark.covers("mgmt.team.daily_activity.happy_path")
|
||||
@pytest.mark.parametrize("days", [1, 7, 30])
|
||||
def test_valid_date_range_returns_results_and_metadata(self, client: SpendClient, days: int) -> None:
|
||||
result = _probe(client, _range_days(days))
|
||||
assert result.status_code == 200, (
|
||||
f"{ROUTE} range={days}d must be 200, got {result.status_code}: {result.body[:600]}"
|
||||
def test_valid_date_range_returns_results_and_metadata(
|
||||
self, client: SpendClient, resources: ResourceManager
|
||||
) -> None:
|
||||
started: Final = datetime.now(timezone.utc).date()
|
||||
traffic: Final = create_traffic(client, resources)
|
||||
for team in traffic:
|
||||
assert_logs_match(client, team)
|
||||
ended: Final = datetime.now(timezone.utc).date()
|
||||
team_ids: Final = ",".join(team.team_id for team in traffic)
|
||||
|
||||
def fetch(
|
||||
page: int, start: str = (started - timedelta(days=1)).isoformat(), end: str = ended.isoformat()
|
||||
) -> TeamDailyActivityResponse:
|
||||
result: Final = _probe(
|
||||
client,
|
||||
TeamDailyActivityParams(
|
||||
start_date=start,
|
||||
end_date=end,
|
||||
page=page,
|
||||
page_size=1,
|
||||
team_ids=team_ids,
|
||||
),
|
||||
)
|
||||
assert result.status_code == 200, f"daily activity failed: {result.status_code} {result.body[:300]}"
|
||||
return TeamDailyActivityResponse.model_validate_json(result.body)
|
||||
|
||||
def pages() -> tuple[TeamDailyActivityResponse, ...]:
|
||||
first: Final = fetch(1)
|
||||
assert first.metadata.total_pages <= len(traffic) * 2, "unexpected extra scoped daily groups"
|
||||
return (first, *(fetch(page) for page in range(2, first.metadata.total_pages + 1)))
|
||||
|
||||
outcome: Final = await_converged(
|
||||
pages,
|
||||
converged=lambda values: (
|
||||
sum(page.metadata.total_api_requests for page in values) >= sum(len(team.responses) for team in traffic)
|
||||
),
|
||||
timeout=client.proxy.poll_timeout,
|
||||
interval=client.proxy.poll_interval,
|
||||
now=time.monotonic,
|
||||
sleep=time.sleep,
|
||||
)
|
||||
parsed = TeamDailyActivityResponse.model_validate_json(result.body)
|
||||
assert parsed.metadata.page == 1
|
||||
assert parsed.metadata.total_pages >= 1
|
||||
if parsed.results:
|
||||
first = parsed.results[0]
|
||||
assert first.date
|
||||
assert first.metrics.spend >= 0
|
||||
assert first.metrics.total_tokens >= 0
|
||||
observed: Final = outcome.result if isinstance(outcome, Converged) else outcome.last_result
|
||||
assert observed is not None, "daily aggregation must return a response before the deadline"
|
||||
|
||||
assert len(observed) >= 2, "two teams must exercise a page boundary"
|
||||
|
||||
def assert_page(index: int, page: TeamDailyActivityResponse) -> None:
|
||||
assert page.metadata.page == index
|
||||
assert page.metadata.total_pages == len(observed)
|
||||
assert page.metadata.has_more == (index < len(observed))
|
||||
assert len(page.results) == 1, "each fetched daily group must appear in results"
|
||||
row: Final = page.results[0]
|
||||
assert started <= datetime.fromisoformat(row.date).date() <= ended
|
||||
assert len(row.breakdown.entities) == 1
|
||||
assert row.metrics.total_tokens == page.metadata.total_tokens
|
||||
assert row.metrics.prompt_tokens == page.metadata.total_prompt_tokens
|
||||
assert row.metrics.completion_tokens == page.metadata.total_completion_tokens
|
||||
assert row.metrics.api_requests == page.metadata.total_api_requests
|
||||
assert row.metrics.successful_requests == page.metadata.total_successful_requests
|
||||
assert row.metrics.failed_requests == page.metadata.total_failed_requests
|
||||
assert isclose(row.metrics.spend, page.metadata.total_spend, rel_tol=1e-6, abs_tol=1e-9)
|
||||
|
||||
for index, page in enumerate(observed, 1):
|
||||
assert_page(index, page)
|
||||
|
||||
entities: Final = tuple(
|
||||
(team_id, entity.metrics)
|
||||
for page in observed
|
||||
for row in page.results
|
||||
for team_id, entity in row.breakdown.entities.items()
|
||||
)
|
||||
assert frozenset(team_id for team_id, _ in entities) == frozenset(team.team_id for team in traffic)
|
||||
|
||||
def assert_team(team: TeamTraffic) -> None:
|
||||
metrics: Final = tuple(metrics for team_id, metrics in entities if team_id == team.team_id)
|
||||
assert sum(m.api_requests for m in metrics) == len(team.responses)
|
||||
assert sum(m.successful_requests for m in metrics) == len(team.responses)
|
||||
assert sum(m.failed_requests for m in metrics) == 0
|
||||
assert sum(m.prompt_tokens for m in metrics) == team.prompt_tokens
|
||||
assert sum(m.completion_tokens for m in metrics) == team.completion_tokens
|
||||
assert sum(m.total_tokens for m in metrics) == team.prompt_tokens + team.completion_tokens
|
||||
assert isclose(sum(m.spend for m in metrics), team.spend, rel_tol=1e-6, abs_tol=1e-9)
|
||||
|
||||
for team in traffic:
|
||||
assert_team(team)
|
||||
|
||||
assert isclose(
|
||||
sum(page.metadata.total_spend for page in observed),
|
||||
sum(team.spend for team in traffic),
|
||||
rel_tol=1e-6,
|
||||
abs_tol=1e-9,
|
||||
)
|
||||
assert sum(page.metadata.total_tokens for page in observed) == sum(
|
||||
team.prompt_tokens + team.completion_tokens for team in traffic
|
||||
)
|
||||
|
||||
for days in (7, 30):
|
||||
assert (
|
||||
tuple(fetch(page, (started - timedelta(days=days)).isoformat()) for page in range(1, len(observed) + 1))
|
||||
== observed
|
||||
), f"{days}-day activity must preserve the same isolated groups and totals"
|
||||
|
||||
empty_date: Final = (started - timedelta(days=7)).isoformat()
|
||||
empty: Final = fetch(1, empty_date, empty_date)
|
||||
assert empty.results == []
|
||||
assert empty.metadata.total_pages == 0
|
||||
assert empty.metadata.page == 1
|
||||
assert not empty.metadata.has_more
|
||||
assert empty.metadata.total_spend == 0
|
||||
assert empty.metadata.total_tokens == 0
|
||||
assert empty.metadata.total_api_requests == 0
|
||||
assert empty.metadata.total_prompt_tokens == 0
|
||||
assert empty.metadata.total_completion_tokens == 0
|
||||
assert empty.metadata.total_successful_requests == 0
|
||||
assert empty.metadata.total_failed_requests == 0
|
||||
|
||||
@pytest.mark.covers("mgmt.team.daily_activity.missing_start_date_rejected")
|
||||
def test_missing_start_date_is_rejected(self, client: SpendClient) -> None:
|
||||
|
|
|
|||
|
|
@ -6,12 +6,15 @@ including the logging handler, cost tracking, and WebSocket message processing.
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, Mock, patch, MagicMock
|
||||
from typing import Dict, List, Any, Optional
|
||||
|
||||
import pytest
|
||||
import httpx
|
||||
import litellm
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
# Add the parent directory to the system path
|
||||
|
||||
|
|
@ -22,10 +25,16 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
|
|||
PassThroughEndpointLogging,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.types.utils import CostBreakdown, LlmProviders, Usage
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
|
||||
class _LiveTurn(TypedDict):
|
||||
prompt: ReadOnly[tuple[int, int]]
|
||||
candidates: ReadOnly[tuple[int, int]]
|
||||
candidate_audio_token_count_missing: NotRequired[ReadOnly[bool]]
|
||||
|
||||
|
||||
class TestVertexAILivePassthroughLoggingHandler:
|
||||
"""Test the Vertex AI Live Passthrough Logging Handler"""
|
||||
|
||||
|
|
@ -39,6 +48,7 @@ class TestVertexAILivePassthroughLoggingHandler:
|
|||
"""Create a mock logging object"""
|
||||
mock = MagicMock(spec=LiteLLMLoggingObj)
|
||||
mock.model_call_details = {}
|
||||
mock._response_cost_calculator.return_value = None
|
||||
return mock
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -201,88 +211,490 @@ class TestVertexAILivePassthroughLoggingHandler:
|
|||
assert text_prompt["tokenCount"] == 10
|
||||
assert audio_prompt["tokenCount"] == 10
|
||||
|
||||
@patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
|
||||
)
|
||||
def test_calculate_cost_basic(self, mock_get_model_info, handler):
|
||||
"""Test basic cost calculation"""
|
||||
mock_get_model_info.return_value = {
|
||||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000002,
|
||||
}
|
||||
def test_usage_carries_every_modality(self, handler):
|
||||
"""Regression: the Usage object reported only TEXT, so audio and image billed as nothing.
|
||||
|
||||
prompt_tokens must be the full count and the details must name each modality,
|
||||
because the cost calculator prices audio and image from *_tokens_details.
|
||||
"""
|
||||
usage_metadata = {
|
||||
"promptTokenCount": 100,
|
||||
"candidatesTokenCount": 50,
|
||||
"totalTokenCount": 150,
|
||||
}
|
||||
|
||||
cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata)
|
||||
|
||||
# The cost calculation may include additional factors, so we check it's reasonable
|
||||
expected_min_cost = (100 * 0.000001) + (50 * 0.000002)
|
||||
assert cost >= expected_min_cost
|
||||
assert cost > 0
|
||||
|
||||
@patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
|
||||
)
|
||||
def test_calculate_cost_with_audio(self, mock_get_model_info, handler):
|
||||
"""Test cost calculation with audio tokens"""
|
||||
mock_get_model_info.return_value = {
|
||||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000002,
|
||||
"input_cost_per_audio_token": 0.0001,
|
||||
"output_cost_per_audio_token": 0.0002,
|
||||
}
|
||||
|
||||
usage_metadata = {
|
||||
"promptTokenCount": 100,
|
||||
"candidatesTokenCount": 50,
|
||||
"totalTokenCount": 150,
|
||||
"promptTokenCount": 1300,
|
||||
"candidatesTokenCount": 124,
|
||||
"totalTokenCount": 1424,
|
||||
"promptTokensDetails": [
|
||||
{"modality": "TEXT", "tokenCount": 80},
|
||||
{"modality": "AUDIO", "tokenCount": 20},
|
||||
{"modality": "TEXT", "tokenCount": 13},
|
||||
{"modality": "AUDIO", "tokenCount": 127},
|
||||
{"modality": "IMAGE", "tokenCount": 1160},
|
||||
],
|
||||
"candidatesTokensDetails": [
|
||||
{"modality": "TEXT", "tokenCount": 30},
|
||||
{"modality": "AUDIO", "tokenCount": 20},
|
||||
{"modality": "TEXT", "tokenCount": 29},
|
||||
{"modality": "AUDIO", "tokenCount": 95},
|
||||
],
|
||||
}
|
||||
|
||||
cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata)
|
||||
usage = handler._create_usage_object_from_metadata(
|
||||
usage_metadata=usage_metadata, model="gemini-live-2.5-flash"
|
||||
)
|
||||
|
||||
# Should include both text and audio costs
|
||||
assert cost > 0
|
||||
assert cost > (100 * 0.000001) + (
|
||||
50 * 0.000002
|
||||
) # Should be higher due to audio
|
||||
assert usage.prompt_tokens == 1300, "the full prompt count must survive, not just its text share"
|
||||
assert usage.completion_tokens == 124
|
||||
assert usage.prompt_tokens_details.text_tokens == 13
|
||||
assert usage.prompt_tokens_details.audio_tokens == 127
|
||||
assert usage.prompt_tokens_details.image_tokens == 1160
|
||||
assert usage.completion_tokens_details.text_tokens == 29
|
||||
assert usage.completion_tokens_details.audio_tokens == 95
|
||||
|
||||
@patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
|
||||
def test_usage_sums_repeated_modality_entries(self, handler):
|
||||
"""A modality can appear more than once across aggregated turns; sum, don't overwrite."""
|
||||
usage = handler._create_usage_object_from_metadata(
|
||||
usage_metadata={
|
||||
"promptTokenCount": 40,
|
||||
"candidatesTokenCount": 0,
|
||||
"promptTokensDetails": [
|
||||
{"modality": "IMAGE", "tokenCount": 10},
|
||||
{"modality": "IMAGE", "tokenCount": 25},
|
||||
{"modality": "TEXT", "tokenCount": 5},
|
||||
],
|
||||
},
|
||||
model="gemini-live-2.5-flash",
|
||||
)
|
||||
assert usage.prompt_tokens_details.image_tokens == 35
|
||||
assert usage.prompt_tokens_details.text_tokens == 5
|
||||
|
||||
NATIVE_AUDIO_MODEL = "gemini-live-2.5-flash-preview-native-audio-09-2025"
|
||||
|
||||
# A four-turn native-audio session. Google charges per turn for the whole session context
|
||||
# window, so the prompt side repeats the accumulated audio while the candidates side reports
|
||||
# only that turn's own response. The last turn names AUDIO and omits its tokenCount, which is
|
||||
# the shape Live really emits at the end of a spoken answer.
|
||||
AUDIO_SESSION: tuple[_LiveTurn, ...] = (
|
||||
{"prompt": (14, 122), "candidates": (8, 20)},
|
||||
{"prompt": (21, 182), "candidates": (5, 50)},
|
||||
{"prompt": (24, 203), "candidates": (13, 27)},
|
||||
{"prompt": (24, 203), "candidates": (0, 3), "candidate_audio_token_count_missing": True},
|
||||
)
|
||||
def test_calculate_cost_with_web_search(self, mock_get_model_info, handler):
|
||||
"""Test cost calculation with web search (tool use)"""
|
||||
mock_get_model_info.return_value = {
|
||||
"input_cost_per_token": 0.000001,
|
||||
"output_cost_per_token": 0.000002,
|
||||
"web_search_cost_per_request": 0.01,
|
||||
}
|
||||
|
||||
usage_metadata = {
|
||||
"promptTokenCount": 100,
|
||||
"candidatesTokenCount": 50,
|
||||
"totalTokenCount": 150,
|
||||
"toolUsePromptTokenCount": 10,
|
||||
}
|
||||
@staticmethod
|
||||
def _live_messages(turns: Sequence[_LiveTurn]) -> list[dict[str, object]]:
|
||||
"""Wrap (text, audio) prompt/candidate pairs as the server messages a Live session emits."""
|
||||
return [{"type": "session.created", "session": {"id": "s"}}] + [
|
||||
{
|
||||
"type": "response.done",
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": sum(turn["prompt"]),
|
||||
"candidatesTokenCount": sum(turn["candidates"]),
|
||||
"totalTokenCount": sum(turn["prompt"]) + sum(turn["candidates"]),
|
||||
"promptTokensDetails": [
|
||||
{"modality": "TEXT", "tokenCount": turn["prompt"][0]},
|
||||
{"modality": "AUDIO", "tokenCount": turn["prompt"][1]},
|
||||
],
|
||||
"candidatesTokensDetails": (
|
||||
[{"modality": "AUDIO"}]
|
||||
if turn.get("candidate_audio_token_count_missing")
|
||||
else [
|
||||
{"modality": "TEXT", "tokenCount": turn["candidates"][0]},
|
||||
{"modality": "AUDIO", "tokenCount": turn["candidates"][1]},
|
||||
]
|
||||
),
|
||||
},
|
||||
}
|
||||
for turn in turns
|
||||
]
|
||||
|
||||
cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata)
|
||||
@staticmethod
|
||||
def _session_usage(
|
||||
handler: VertexAILivePassthroughLoggingHandler,
|
||||
mock_logging_obj: MagicMock,
|
||||
messages: list[dict[str, object]],
|
||||
model: str,
|
||||
) -> Usage:
|
||||
result = handler.vertex_ai_live_passthrough_handler(
|
||||
websocket_messages=messages,
|
||||
logging_obj=mock_logging_obj,
|
||||
url_route="/vertex_ai/live",
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
request_body={},
|
||||
model=model,
|
||||
)
|
||||
assert result["result"] is not None, "the handler must produce a usage-bearing response to bill"
|
||||
return result["result"].usage
|
||||
|
||||
# Should include web search cost
|
||||
expected_base_cost = (100 * 0.000001) + (50 * 0.000002)
|
||||
# The web search cost might be handled differently, so just check it's reasonable
|
||||
assert cost >= expected_base_cost
|
||||
assert cost > 0
|
||||
@classmethod
|
||||
def _session_cost(
|
||||
cls,
|
||||
handler: VertexAILivePassthroughLoggingHandler,
|
||||
mock_logging_obj: MagicMock,
|
||||
messages: list[dict[str, object]],
|
||||
model: str,
|
||||
) -> float:
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
usage = cls._session_usage(handler, mock_logging_obj, messages, model)
|
||||
return completion_cost(
|
||||
completion_response=ModelResponse(
|
||||
id="x", object="chat.completion", created=0, model=model, usage=usage, choices=[]
|
||||
),
|
||||
model=f"vertex_ai/{model}",
|
||||
custom_llm_provider="vertex_ai",
|
||||
call_type="acompletion",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _expected_session_cost(cls, turns: Sequence[_LiveTurn]) -> float:
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
info = get_model_info(model=cls.NATIVE_AUDIO_MODEL, custom_llm_provider="vertex_ai")
|
||||
return (
|
||||
sum(turn["prompt"][0] for turn in turns) * info["input_cost_per_token"]
|
||||
+ sum(turn["prompt"][1] for turn in turns) * info["input_cost_per_audio_token"]
|
||||
+ sum(turn["candidates"][0] for turn in turns) * info["output_cost_per_token"]
|
||||
+ sum(turn["candidates"][1] for turn in turns) * info["output_cost_per_audio_token"]
|
||||
)
|
||||
|
||||
def test_every_turn_of_a_session_is_billed(self, handler, mock_logging_obj):
|
||||
"""Google charges per turn for the whole context window, so every turn adds to the bill.
|
||||
|
||||
Billing one snapshot instead gives away all the other turns: on this session the
|
||||
largest single turn is well under the session total, and its share of the audio is
|
||||
priced 6x the text rate, so the gap is money rather than rounding.
|
||||
"""
|
||||
turns = self.AUDIO_SESSION[:3]
|
||||
cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
|
||||
|
||||
assert cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9)
|
||||
widest_single_turn = max(self._expected_session_cost([turn]) for turn in turns)
|
||||
assert cost > widest_single_turn, "billing one snapshot drops every other turn of the session"
|
||||
|
||||
def test_audio_named_without_a_token_count_bills_at_the_audio_rate(self, handler, mock_logging_obj):
|
||||
"""Live can name the modality carrying the rest of a turn and omit its tokenCount.
|
||||
|
||||
Reading the absent key as zero left those tokens inside candidatesTokenCount but outside
|
||||
the breakdown, so the calculator charged real speech at the text output rate. At this
|
||||
entry's rates the last turn's 3 audio tokens are $0.0000360 rather than $0.0000060.
|
||||
"""
|
||||
turns = self.AUDIO_SESSION
|
||||
usage = self._session_usage(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
|
||||
|
||||
assert usage.completion_tokens_details.audio_tokens == 100, "the unpriced entry takes the turn's residual"
|
||||
assert usage.completion_tokens_details.text_tokens == 26
|
||||
assert usage.completion_tokens == 126
|
||||
|
||||
cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
|
||||
assert cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9)
|
||||
|
||||
TOOL_USE_PER_TURN = (100, 250, 400)
|
||||
|
||||
def _grounded_messages(self):
|
||||
"""The three-turn session again, with each turn's own toolUsePromptTokenCount attached."""
|
||||
messages = self._live_messages(self.AUDIO_SESSION[:3])
|
||||
head, turns = messages[0], messages[1:]
|
||||
return [head] + [
|
||||
{**message, "usageMetadata": {**message["usageMetadata"], "toolUsePromptTokenCount": tool_use}}
|
||||
for message, tool_use in zip(turns, self.TOOL_USE_PER_TURN)
|
||||
]
|
||||
|
||||
def test_server_side_tool_use_prompt_tokens_are_summed_over_the_session(self, handler, mock_logging_obj):
|
||||
"""toolUsePromptTokenCount rode the unknown-key pass-through, so it took the first turn only.
|
||||
|
||||
Every other total beside it is summed across the session, and the first turn is the
|
||||
smallest number in the series, so a grounded session logged far fewer tool-use tokens
|
||||
than it used. This session's turns are deliberately distinct, so 750 can only come from
|
||||
summing: first-turn selection gives 100, last-turn or max gives 400.
|
||||
"""
|
||||
grounded = self._grounded_messages()
|
||||
|
||||
usage = self._session_usage(handler, mock_logging_obj, grounded, self.NATIVE_AUDIO_MODEL)
|
||||
assert usage.prompt_tokens_details.tool_use_tokens == sum(self.TOOL_USE_PER_TURN)
|
||||
|
||||
@staticmethod
|
||||
def _grounding_frame(metadata: dict[str, object]) -> dict[str, object]:
|
||||
"""One server frame carrying grounding metadata, the way Live reports it."""
|
||||
return {"type": "response.done", "serverContent": {"groundingMetadata": metadata}}
|
||||
|
||||
def test_web_grounding_is_counted_so_it_can_be_billed(self, handler, mock_logging_obj):
|
||||
"""Live reports grounding in the server frames and never in usageMetadata.
|
||||
|
||||
Nothing read those frames, so web_search_requests stayed unset and the cost path's only
|
||||
trigger for the per-query grounding charge never fired. Google bills a grounded Live
|
||||
prompt on top of its tokens, so the whole fee was missing from the bill.
|
||||
"""
|
||||
messages = [
|
||||
self._grounding_frame(
|
||||
{
|
||||
"webSearchQueries": ["who won the 2026 world cup final"],
|
||||
"groundingChunks": [{"web": {"uri": "https://example.com"}}],
|
||||
}
|
||||
),
|
||||
*self._live_messages(self.AUDIO_SESSION[:1]),
|
||||
]
|
||||
|
||||
usage = self._session_usage(handler, mock_logging_obj, messages, self.NATIVE_AUDIO_MODEL)
|
||||
|
||||
assert usage.prompt_tokens_details.web_search_requests == 1, "a grounded turn must report its query"
|
||||
assert getattr(usage.prompt_tokens_details, "google_maps_grounding_requests", None) is None
|
||||
|
||||
def test_maps_grounding_is_counted_under_its_own_sku(self, handler, mock_logging_obj):
|
||||
"""Maps grounding is a separate SKU from web search, so it needs its own counter.
|
||||
|
||||
A maps-only turn carries grounding chunks but no webSearchQueries, so counting queries
|
||||
alone would report nothing and bill nothing.
|
||||
"""
|
||||
messages = [
|
||||
self._grounding_frame({"groundingChunks": [{"maps": {"placeId": "abc123"}}]}),
|
||||
*self._live_messages(self.AUDIO_SESSION[:1]),
|
||||
]
|
||||
|
||||
usage = self._session_usage(handler, mock_logging_obj, messages, self.NATIVE_AUDIO_MODEL)
|
||||
|
||||
assert usage.prompt_tokens_details.google_maps_grounding_requests == 1
|
||||
assert getattr(usage.prompt_tokens_details, "web_search_requests", None) is None
|
||||
|
||||
def test_an_ungrounded_session_reports_no_grounding(self, handler, mock_logging_obj):
|
||||
"""The counters must stay absent when no tool ran, or every session pays a grounding fee."""
|
||||
usage = self._session_usage(
|
||||
handler, mock_logging_obj, self._live_messages(self.AUDIO_SESSION[:1]), self.NATIVE_AUDIO_MODEL
|
||||
)
|
||||
|
||||
assert getattr(usage.prompt_tokens_details, "web_search_requests", None) is None
|
||||
assert getattr(usage.prompt_tokens_details, "google_maps_grounding_requests", None) is None
|
||||
|
||||
def test_grounding_adds_its_query_fee_to_the_session_bill(self, handler, mock_logging_obj):
|
||||
"""The counter only matters if it reaches the bill, so assert against the cost, not the field.
|
||||
|
||||
Same tokens either way: the difference between the two sessions is the grounding fee alone.
|
||||
"""
|
||||
turns = self.AUDIO_SESSION[:1]
|
||||
plain = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
|
||||
grounded = self._session_cost(
|
||||
handler,
|
||||
mock_logging_obj,
|
||||
[self._grounding_frame({"webSearchQueries": ["q"]}), *self._live_messages(turns)],
|
||||
self.NATIVE_AUDIO_MODEL,
|
||||
)
|
||||
|
||||
assert grounded > plain, "a grounded session must cost more than the same tokens ungrounded"
|
||||
|
||||
def _priced_logging_obj(self) -> LiteLLMLoggingObj:
|
||||
"""A real logging object, since the session's price is handed to it turn by turn."""
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
model=self.NATIVE_AUDIO_MODEL,
|
||||
messages=[],
|
||||
stream=True,
|
||||
call_type="pass_through_endpoint",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="live-session",
|
||||
function_id="live",
|
||||
)
|
||||
logging_obj.update_environment_variables(
|
||||
model=self.NATIVE_AUDIO_MODEL,
|
||||
user="u",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
call_type="pass_through_endpoint",
|
||||
)
|
||||
logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai"
|
||||
return logging_obj
|
||||
|
||||
def _billed_session(
|
||||
self, handler: VertexAILivePassthroughLoggingHandler, messages: list[dict[str, object]]
|
||||
) -> tuple[float, CostBreakdown]:
|
||||
logging_obj = self._priced_logging_obj()
|
||||
result = handler.vertex_ai_live_passthrough_handler(
|
||||
websocket_messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
url_route="/vertex_ai/live",
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
request_body={},
|
||||
model=self.NATIVE_AUDIO_MODEL,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
assert result["result"] is not None, "the handler must produce a usage-bearing response to bill"
|
||||
assert logging_obj.cost_breakdown is not None, "the session's price must reach the logging object"
|
||||
return result["result"]._hidden_params["response_cost"], logging_obj.cost_breakdown
|
||||
|
||||
def test_each_grounded_turn_pays_its_own_query_fee(self, handler):
|
||||
"""Google charges the grounding fee per grounded prompt, not per session.
|
||||
|
||||
Summing the session into one usage collapsed two grounded turns into one query, so the
|
||||
second question was answered for free. The bill now grows by one fee per grounded turn.
|
||||
"""
|
||||
head, turn = self._live_messages(self.AUDIO_SESSION[:1])
|
||||
grounding = self._grounding_frame({"webSearchQueries": ["q"]})
|
||||
|
||||
plain_cost, _ = self._billed_session(handler, [head, turn, turn])
|
||||
one_cost, one_breakdown = self._billed_session(handler, [head, grounding, turn, turn])
|
||||
two_cost, two_breakdown = self._billed_session(handler, [head, grounding, turn, grounding, turn])
|
||||
|
||||
fee = one_cost - plain_cost
|
||||
assert fee > 0, "a grounded turn must cost more than the same tokens ungrounded"
|
||||
assert two_cost - plain_cost == pytest.approx(2 * fee), "two grounded turns must pay the fee twice"
|
||||
assert two_breakdown["total_cost"] == pytest.approx(two_cost)
|
||||
assert two_breakdown["tool_usage_cost"] == pytest.approx(2 * one_breakdown["tool_usage_cost"])
|
||||
|
||||
def test_a_query_repeated_across_turns_is_reported_once_per_turn(self, handler):
|
||||
"""The reported query count must agree with the bill, which charges every grounded turn.
|
||||
|
||||
The session usage collapsed duplicate query strings across turns while the price was
|
||||
per turn, so two turns asking the same question paid two fees yet reported one query.
|
||||
Duplicates within one turn still collapse, since that turn ran one search.
|
||||
"""
|
||||
head, turn = self._live_messages(self.AUDIO_SESSION[:1])
|
||||
grounding = self._grounding_frame({"webSearchQueries": ["q"]})
|
||||
logging_obj = self._priced_logging_obj()
|
||||
|
||||
result = handler.vertex_ai_live_passthrough_handler(
|
||||
websocket_messages=[head, grounding, turn, grounding, turn],
|
||||
logging_obj=logging_obj,
|
||||
url_route="/vertex_ai/live",
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
request_body={},
|
||||
model=self.NATIVE_AUDIO_MODEL,
|
||||
custom_llm_provider="vertex_ai",
|
||||
)
|
||||
_, one_breakdown = self._billed_session(handler, [head, grounding, turn])
|
||||
repeated_within_turn = handler._session_usage(
|
||||
[head, self._grounding_frame({"webSearchQueries": ["q", "q"]}), turn], self.NATIVE_AUDIO_MODEL
|
||||
)
|
||||
|
||||
assert result["result"].usage.prompt_tokens_details.web_search_requests == 2
|
||||
assert logging_obj.cost_breakdown["tool_usage_cost"] == pytest.approx(2 * one_breakdown["tool_usage_cost"])
|
||||
assert repeated_within_turn.prompt_tokens_details.web_search_requests == 1
|
||||
|
||||
def test_the_fixed_cost_margin_is_charged_once_per_session(self, handler):
|
||||
"""A fixed cost margin is a flat per-request fee, and a Live session is one spend row.
|
||||
|
||||
Pricing each turn on its own applied the fixed margin per turn, so a two-turn session paid it
|
||||
twice. The session now carries the fixed margin once no matter how many turns it billed.
|
||||
"""
|
||||
head, turn = self._live_messages(self.AUDIO_SESSION[:1])
|
||||
grounding = self._grounding_frame({"webSearchQueries": ["q"]})
|
||||
messages = [head, grounding, turn, grounding, turn]
|
||||
|
||||
plain_cost, _ = self._billed_session(handler, messages)
|
||||
|
||||
fixed_amount = 0.01
|
||||
with patch.object(litellm, "cost_margin_config", {"vertex_ai": {"fixed_amount": fixed_amount}}):
|
||||
margined_cost, breakdown = self._billed_session(handler, messages)
|
||||
|
||||
assert margined_cost - plain_cost == pytest.approx(
|
||||
fixed_amount
|
||||
), "a two-turn session must add the fixed margin once, not once per billed turn"
|
||||
assert breakdown["margin_fixed_amount"] == pytest.approx(fixed_amount)
|
||||
assert breakdown["margin_total_amount"] == pytest.approx(fixed_amount)
|
||||
|
||||
def test_reporting_tool_use_tokens_does_not_move_the_bill(self, handler, mock_logging_obj):
|
||||
"""Deliberate boundary: these tokens are reported here, and priced nowhere.
|
||||
|
||||
generic_cost_per_token reads the input bill out of prompt_tokens_details, and falls
|
||||
back to prompt_tokens only when the details carry no text or a cache hit overlaps them,
|
||||
so adding tool-use tokens to prompt_tokens is worth nothing on an ordinary Live turn and
|
||||
over-charges against the cache-overlap correction when it is not. Pricing them belongs
|
||||
in the shared input-cost path, beside the modality terms that already read the details.
|
||||
"""
|
||||
turns = self.AUDIO_SESSION[:3]
|
||||
plain_cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
|
||||
grounded_cost = self._session_cost(
|
||||
handler, mock_logging_obj, self._grounded_messages(), self.NATIVE_AUDIO_MODEL
|
||||
)
|
||||
|
||||
assert plain_cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9)
|
||||
assert grounded_cost == pytest.approx(plain_cost, rel=1e-9), "reporting tool use must not move the bill"
|
||||
|
||||
def test_a_malformed_details_entry_does_not_cost_the_whole_session(self, handler, mock_logging_obj):
|
||||
"""A ``*TokensDetails`` value that is not a list of objects must not take the session down.
|
||||
|
||||
The handler's only error path returns no result at all, so one odd frame used to throw
|
||||
while reading it and the whole session billed nothing. The good turns still bill.
|
||||
"""
|
||||
turns = self.AUDIO_SESSION[:3]
|
||||
messages = self._live_messages(turns)
|
||||
mangled = [dict(message) for message in messages]
|
||||
mangled[1]["usageMetadata"] = {**mangled[1]["usageMetadata"], "promptTokensDetails": "TEXT"}
|
||||
|
||||
usage = self._session_usage(handler, mock_logging_obj, mangled, self.NATIVE_AUDIO_MODEL)
|
||||
|
||||
surviving = turns[1:]
|
||||
assert usage.prompt_tokens_details.audio_tokens == sum(turn["prompt"][1] for turn in surviving)
|
||||
assert usage.prompt_tokens_details.text_tokens == sum(turn["prompt"][0] for turn in surviving)
|
||||
assert usage.prompt_tokens == sum(sum(turn["prompt"]) for turn in turns), "the totals still cover every turn"
|
||||
|
||||
direct = handler._create_usage_object_from_metadata(
|
||||
usage_metadata={
|
||||
"promptTokenCount": 40,
|
||||
"candidatesTokenCount": 12,
|
||||
"promptTokensDetails": [{"modality": "AUDIO", "tokenCount": 40}, "AUDIO"],
|
||||
"candidatesTokensDetails": {"modality": "TEXT", "tokenCount": 12},
|
||||
},
|
||||
model=self.NATIVE_AUDIO_MODEL,
|
||||
)
|
||||
assert direct.prompt_tokens_details.audio_tokens == 40, "the well-formed entry beside a bad one still counts"
|
||||
assert direct.completion_tokens == 12
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"label,prompt_details,candidate_details",
|
||||
[
|
||||
("text only", [("TEXT", 6)], [("TEXT", 2)]),
|
||||
("audio in", [("TEXT", 13), ("AUDIO", 127)], [("TEXT", 18)]),
|
||||
("image in", [("TEXT", 10), ("IMAGE", 258)], [("TEXT", 24)]),
|
||||
("frames in", [("TEXT", 11), ("IMAGE", 1032)], [("TEXT", 26)]),
|
||||
("audio both ways", [("TEXT", 13), ("AUDIO", 127)], [("TEXT", 29), ("AUDIO", 95)]),
|
||||
],
|
||||
)
|
||||
def test_live_session_bills_each_modality_at_its_own_rate(self, handler, label, prompt_details, candidate_details):
|
||||
"""Every payload here is a real Vertex Live session's usageMetadata.
|
||||
|
||||
Before the fix these billed the text share only, from 1x (text) to 55x under.
|
||||
The expected amount is derived from the entry's own rates rather than hardcoded,
|
||||
so this stays correct as prices move, and it is asserted exactly, so dropping a
|
||||
modality and double-charging one both fail.
|
||||
"""
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
model = self.NATIVE_AUDIO_MODEL
|
||||
info = get_model_info(model=model, custom_llm_provider="vertex_ai")
|
||||
|
||||
text_in = info["input_cost_per_token"]
|
||||
audio_in = info.get("input_cost_per_audio_token") or text_in
|
||||
image_in = info.get("input_cost_per_image_token") or text_in
|
||||
text_out = info["output_cost_per_token"]
|
||||
audio_out = info.get("output_cost_per_audio_token") or text_out
|
||||
rate_in = {"TEXT": text_in, "AUDIO": audio_in, "IMAGE": image_in}
|
||||
rate_out = {"TEXT": text_out, "AUDIO": audio_out}
|
||||
|
||||
expected = sum(c * rate_in[m] for m, c in prompt_details) + sum(c * rate_out[m] for m, c in candidate_details)
|
||||
|
||||
usage = handler._create_usage_object_from_metadata(
|
||||
usage_metadata={
|
||||
"promptTokenCount": sum(c for _, c in prompt_details),
|
||||
"candidatesTokenCount": sum(c for _, c in candidate_details),
|
||||
"promptTokensDetails": [{"modality": m, "tokenCount": c} for m, c in prompt_details],
|
||||
"candidatesTokensDetails": [{"modality": m, "tokenCount": c} for m, c in candidate_details],
|
||||
},
|
||||
model=model,
|
||||
)
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=ModelResponse(
|
||||
id="x", object="chat.completion", created=0, model=model, usage=usage, choices=[]
|
||||
),
|
||||
model=f"vertex_ai/{model}",
|
||||
custom_llm_provider="vertex_ai",
|
||||
call_type="acompletion",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(expected, rel=1e-9), label
|
||||
|
||||
text_only = sum(c for m, c in prompt_details if m == "TEXT") * text_in + sum(
|
||||
c for m, c in candidate_details if m == "TEXT"
|
||||
) * text_out
|
||||
if any(m != "TEXT" for m, _ in prompt_details + candidate_details) and audio_in != text_in:
|
||||
assert cost > text_only, f"{label}: non-text modalities must add cost"
|
||||
|
||||
def test_vertex_ai_live_passthrough_handler_integration(
|
||||
self, handler, mock_logging_obj, sample_websocket_messages
|
||||
|
|
@ -376,6 +788,7 @@ class TestVertexAILivePassthroughIntegration:
|
|||
"""Create a mock logging object"""
|
||||
mock = MagicMock(spec=LiteLLMLoggingObj)
|
||||
mock.model_call_details = {}
|
||||
mock._response_cost_calculator.return_value = None
|
||||
return mock
|
||||
|
||||
@patch(
|
||||
|
|
@ -509,6 +922,7 @@ class TestVertexAILivePassthroughErrorHandling:
|
|||
"""Create a mock logging object"""
|
||||
mock = MagicMock(spec=LiteLLMLoggingObj)
|
||||
mock.model_call_details = {}
|
||||
mock._response_cost_calculator.return_value = None
|
||||
return mock
|
||||
|
||||
def test_invalid_websocket_messages_format(self):
|
||||
|
|
@ -540,25 +954,24 @@ class TestVertexAILivePassthroughErrorHandling:
|
|||
result = handler._extract_usage_metadata_from_websocket_messages(messages)
|
||||
assert result is None
|
||||
|
||||
@patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
|
||||
)
|
||||
def test_cost_calculation_with_missing_model_info(self, mock_get_model_info):
|
||||
"""Test cost calculation when model info is missing"""
|
||||
def test_usage_without_modality_details(self):
|
||||
"""Older payloads carry only the totals; fall back to them rather than reporting zero."""
|
||||
handler = VertexAILivePassthroughLoggingHandler()
|
||||
|
||||
# Mock missing model info
|
||||
mock_get_model_info.return_value = {}
|
||||
usage = handler._create_usage_object_from_metadata(
|
||||
usage_metadata={
|
||||
"promptTokenCount": 100,
|
||||
"candidatesTokenCount": 50,
|
||||
"totalTokenCount": 150,
|
||||
},
|
||||
model="unknown-model",
|
||||
)
|
||||
|
||||
usage_metadata = {
|
||||
"promptTokenCount": 100,
|
||||
"candidatesTokenCount": 50,
|
||||
"totalTokenCount": 150,
|
||||
}
|
||||
|
||||
# Should not raise an exception, should return 0 or handle gracefully
|
||||
cost = handler._calculate_live_api_cost("unknown-model", usage_metadata)
|
||||
assert cost == 0.0
|
||||
assert usage.prompt_tokens == 100
|
||||
assert usage.completion_tokens == 50
|
||||
assert usage.total_tokens == 150
|
||||
assert usage.prompt_tokens_details.audio_tokens is None
|
||||
assert usage.prompt_tokens_details.image_tokens is None
|
||||
|
||||
def test_handler_with_none_websocket_messages(self, mock_logging_obj):
|
||||
"""Test handler with None websocket messages"""
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import litellm
|
|||
from unittest.mock import patch, MagicMock, AsyncMock
|
||||
from create_mock_standard_logging_payload import create_standard_logging_payload
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
|
||||
from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo
|
||||
from litellm.constants import DEFAULT_AUTO_ROUTER_MAX_INPUT_CHARS
|
||||
|
||||
|
||||
|
|
@ -630,10 +630,12 @@ def test_deployment_callback_respects_cooldown_time(model_list):
|
|||
|
||||
|
||||
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
|
||||
def test_log_retry(model_list, metadata_key):
|
||||
"""log_retry appends one flat record per failed attempt and copies neither the request kwargs nor
|
||||
the request metadata into it"""
|
||||
def test_log_retry(model_list: list[DeploymentTypedDict], metadata_key: str) -> None:
|
||||
"""log_retry appends one flat record per failed attempt, copies neither the request kwargs nor the
|
||||
request metadata into it, counts every failed attempt of the request independently of the
|
||||
per-hop attempted_retries, and never trusts a negative count planted before the first failure"""
|
||||
router = Router(model_list=model_list)
|
||||
rate_limit_error = litellm.RateLimitError(message="slow down", llm_provider="openai", model="gpt-3.5-turbo")
|
||||
new_kwargs = router.log_retry(
|
||||
kwargs={
|
||||
"model": "gpt-3.5-turbo",
|
||||
|
|
@ -641,7 +643,7 @@ def test_log_retry(model_list, metadata_key):
|
|||
"messages": [{"role": "user", "content": "hi"}],
|
||||
metadata_key: {"model_info": {"id": "deployment-1"}, "attempted_retries": 2, "user_api_key": "sk-proxy"},
|
||||
},
|
||||
e=litellm.RateLimitError(message="slow down", llm_provider="openai", model="gpt-3.5-turbo"),
|
||||
e=rate_limit_error,
|
||||
)
|
||||
assert json.loads(json.dumps(new_kwargs[metadata_key]["previous_models"])) == [
|
||||
{
|
||||
|
|
@ -652,6 +654,10 @@ def test_log_retry(model_list, metadata_key):
|
|||
"attempted_retries": 2,
|
||||
}
|
||||
]
|
||||
assert new_kwargs[metadata_key]["request_retry_count"] == 1
|
||||
assert router.log_retry(kwargs=new_kwargs, e=rate_limit_error)[metadata_key]["request_retry_count"] == 2
|
||||
planted_kwargs = {"model": "gpt-3.5-turbo", metadata_key: {"request_retry_count": -100}}
|
||||
assert router.log_retry(kwargs=planted_kwargs, e=rate_limit_error)[metadata_key]["request_retry_count"] == 1
|
||||
|
||||
|
||||
def test_update_usage(model_list):
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import json
|
|||
from datetime import datetime, timezone
|
||||
|
||||
import pytest
|
||||
from collections.abc import Mapping
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import litellm
|
||||
|
|
@ -4141,7 +4142,7 @@ def test_billed_token_rates_follow_the_token_tier_the_breakdown_bills_at(monkeyp
|
|||
cache_read_input_token_cost=6e-7,
|
||||
cache_read_input_audio_token_cost=6e-7,
|
||||
cache_creation_input_token_cost=7.5e-6,
|
||||
cache_creation_input_token_cost_above_1hr=0.0,
|
||||
cache_creation_input_token_cost_above_1hr=7.5e-6,
|
||||
output_cost_per_reasoning_token=3e-5,
|
||||
)
|
||||
assert breakdown.cache_read_cost == pytest.approx(200_000 * rates.cache_read_input_token_cost)
|
||||
|
|
@ -5436,3 +5437,72 @@ def test_realtime_models_bill_cached_text_and_audio_at_their_cache_read_rates(
|
|||
|
||||
prompt_cost, _ = generic_cost_per_token(model=model, usage=usage, custom_llm_provider=custom_llm_provider)
|
||||
assert prompt_cost == pytest.approx(expected_prompt_cost)
|
||||
|
||||
|
||||
def test_generic_cost_per_token_bills_cache_creation_at_the_input_rate_without_a_write_price():
|
||||
"""Azure and OpenAI publish no cache-write price and bill cache writes as ordinary input.
|
||||
A deployment priced with only input, output, and cache-read rates must bill the creation
|
||||
tokens the provider reports at the input rate, never at 0. The numbers are a cold 7,336-token
|
||||
prompt on a deployment that reports all but 3 of them as cache creation."""
|
||||
model_info = {
|
||||
"input_cost_per_token": 2e-7,
|
||||
"output_cost_per_token": 1.25e-6,
|
||||
"cache_read_input_token_cost": 2e-8,
|
||||
}
|
||||
usage = Usage(
|
||||
prompt_tokens=7336,
|
||||
completion_tokens=23,
|
||||
total_tokens=7359,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=0, cache_creation_tokens=7333),
|
||||
)
|
||||
|
||||
prompt_cost, completion_cost = generic_cost_per_token(
|
||||
model="custom-priced-deployment", usage=usage, custom_llm_provider="azure", model_info=model_info
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(7336 * 2e-7)
|
||||
assert completion_cost == pytest.approx(23 * 1.25e-6)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("cache_rates", "current_time", "expected_creation", "expected_creation_1h"),
|
||||
(
|
||||
pytest.param({}, None, 2e-7, 2e-7, id="no-write-price-uses-the-input-rate"),
|
||||
pytest.param({"cache_creation_input_token_cost": 2.5e-7}, None, 2.5e-7, 2.5e-7, id="no-1h-price-uses-the-write-price"),
|
||||
pytest.param({"cache_creation_input_token_cost": 0.0}, None, 0.0, 0.0, id="explicit-zero-stays-zero"),
|
||||
pytest.param(
|
||||
{"off_peak_pricing": {"hours_utc": "00:00-23:59", "input_cost_per_token": 1e-7}},
|
||||
datetime(2026, 9, 14, 12, tzinfo=timezone.utc),
|
||||
1e-7,
|
||||
1e-7,
|
||||
id="no-write-price-uses-the-off-peak-input-rate",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"off_peak_pricing": {
|
||||
"hours_utc": "00:00-23:59",
|
||||
"input_cost_per_token": 1e-7,
|
||||
"cache_creation_input_token_cost": 3e-7,
|
||||
}
|
||||
},
|
||||
datetime(2026, 9, 14, 12, tzinfo=timezone.utc),
|
||||
3e-7,
|
||||
3e-7,
|
||||
id="no-1h-price-uses-the-off-peak-write-price",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_get_token_base_cost_resolves_missing_cache_write_rates_like_the_tiered_path(
|
||||
cache_rates: Mapping[str, float | Mapping[str, float | str]],
|
||||
current_time: datetime | None,
|
||||
expected_creation: float,
|
||||
expected_creation_1h: float,
|
||||
) -> None:
|
||||
model_info = {"input_cost_per_token": 2e-7, "output_cost_per_token": 1.25e-6, **cache_rates}
|
||||
usage = Usage(prompt_tokens=10, completion_tokens=1, total_tokens=11)
|
||||
|
||||
_, _, creation, creation_1h, _ = _get_token_base_cost(model_info, usage, current_time=current_time)
|
||||
|
||||
assert creation == pytest.approx(expected_creation)
|
||||
assert creation_1h == pytest.approx(expected_creation_1h)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
|
||||
import httpx
|
||||
import openai
|
||||
import pytest
|
||||
|
|
@ -178,9 +177,7 @@ class TestExceptionCheckers:
|
|||
]
|
||||
|
||||
for error_str in error_strings:
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(
|
||||
error_str
|
||||
)
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str)
|
||||
assert result is True, f"Should detect policy violation in: {error_str}"
|
||||
|
||||
def test_is_azure_content_policy_violation_error_case_insensitive(self):
|
||||
|
|
@ -194,12 +191,8 @@ class TestExceptionCheckers:
|
|||
]
|
||||
|
||||
for error_str in error_strings:
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(
|
||||
error_str
|
||||
)
|
||||
assert (
|
||||
result is True
|
||||
), f"Should detect policy violation in uppercase: {error_str}"
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str)
|
||||
assert result is True, f"Should detect policy violation in uppercase: {error_str}"
|
||||
|
||||
def test_is_azure_content_policy_violation_error_with_non_policy_errors(self):
|
||||
"""Test that non-policy violation errors are not detected as policy violations"""
|
||||
|
|
@ -216,12 +209,8 @@ class TestExceptionCheckers:
|
|||
]
|
||||
|
||||
for error_str in error_strings:
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(
|
||||
error_str
|
||||
)
|
||||
assert (
|
||||
result is False
|
||||
), f"Should NOT detect policy violation in: {error_str}"
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str)
|
||||
assert result is False, f"Should NOT detect policy violation in: {error_str}"
|
||||
|
||||
def test_is_azure_content_policy_violation_error_with_partial_matches(self):
|
||||
"""Test that partial keyword matches work correctly"""
|
||||
|
|
@ -234,9 +223,7 @@ class TestExceptionCheckers:
|
|||
]
|
||||
|
||||
for error_str in positive_cases:
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(
|
||||
error_str
|
||||
)
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str)
|
||||
assert result is True, f"Should detect policy violation in: {error_str}"
|
||||
|
||||
# These should not match even though they contain similar words
|
||||
|
|
@ -248,12 +235,8 @@ class TestExceptionCheckers:
|
|||
]
|
||||
|
||||
for error_str in negative_cases:
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(
|
||||
error_str
|
||||
)
|
||||
assert (
|
||||
result is False
|
||||
), f"Should NOT detect policy violation in: {error_str}"
|
||||
result = ExceptionCheckers.is_azure_content_policy_violation_error(error_str)
|
||||
assert result is False, f"Should NOT detect policy violation in: {error_str}"
|
||||
|
||||
|
||||
gemini_context_window_test_cases = [
|
||||
|
|
@ -271,12 +254,8 @@ gemini_context_window_test_cases = [
|
|||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error_message, should_raise_context_window", gemini_context_window_test_cases
|
||||
)
|
||||
def test_gemini_context_window_error_mapping(
|
||||
error_message, should_raise_context_window
|
||||
):
|
||||
@pytest.mark.parametrize("error_message, should_raise_context_window", gemini_context_window_test_cases)
|
||||
def test_gemini_context_window_error_mapping(error_message, should_raise_context_window):
|
||||
"""
|
||||
Tests that the exception_type function correctly maps Gemini's
|
||||
context window exceeded errors to litellm.ContextWindowExceededError.
|
||||
|
|
@ -421,9 +400,7 @@ vertex_rate_limit_test_cases = [
|
|||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error_message, should_raise_rate_limit", vertex_rate_limit_test_cases
|
||||
)
|
||||
@pytest.mark.parametrize("error_message, should_raise_rate_limit", vertex_rate_limit_test_cases)
|
||||
def test_vertex_ai_rate_limit_error_mapping(error_message, should_raise_rate_limit):
|
||||
"""
|
||||
Tests that the exception_type function correctly maps Vertex AI's
|
||||
|
|
@ -458,10 +435,7 @@ class TestGetBodyErrorCode:
|
|||
"""Unit tests for _get_body_error_code helper."""
|
||||
|
||||
def test_parses_int_code(self):
|
||||
body = (
|
||||
'{"error":{"message":"high demand","type":"upstream_error",'
|
||||
'"param":"","code":429}}'
|
||||
)
|
||||
body = '{"error":{"message":"high demand","type":"upstream_error","param":"","code":429}}'
|
||||
assert _get_body_error_code(body) == 429
|
||||
|
||||
def test_parses_string_code(self):
|
||||
|
|
@ -498,8 +472,7 @@ gemini_body_code_429_test_cases = [
|
|||
),
|
||||
(
|
||||
503,
|
||||
'{"error":{"message":"upstream unavailable","type":"upstream_error",'
|
||||
'"param":"","code":429}}',
|
||||
'{"error":{"message":"upstream unavailable","type":"upstream_error","param":"","code":429}}',
|
||||
litellm.RateLimitError,
|
||||
"HTTP 503 envelope with body code:429 -> RateLimitError",
|
||||
),
|
||||
|
|
@ -769,9 +742,7 @@ class _UpstreamHTTPError(Exception):
|
|||
self.message = "upstream failure"
|
||||
self.status_code = status_code
|
||||
self.request = httpx.Request("POST", "https://api.example.com/v1/chat/completions")
|
||||
self.response = httpx.Response(
|
||||
status_code=status_code, request=self.request, text="upstream failure"
|
||||
)
|
||||
self.response = httpx.Response(status_code=status_code, request=self.request, text="upstream failure")
|
||||
|
||||
|
||||
UPSTREAM_STATUS_CODES = (400, 401, 403, 404, 408, 422, 429, 500, 503)
|
||||
|
|
@ -892,15 +863,13 @@ PROVIDERS_WITHOUT_A_HANDLER = tuple(
|
|||
|
||||
MINIMAX_401_BODY = (
|
||||
'{"type":"error","error":{"type":"authorized_error","message":"login fail: Please carry the API secret key '
|
||||
"in the 'Authorization' field of the request header (1004)\",\"http_code\":\"401\"},"
|
||||
'in the \'Authorization\' field of the request header (1004)","http_code":"401"},'
|
||||
'"request_id":"06ddc9ba97ee6340e38f10e09787f547"}'
|
||||
)
|
||||
|
||||
|
||||
def _expected_for(provider: str, status_code: int) -> tuple[type[Exception], int]:
|
||||
return DEVIATIONS_FROM_THE_OPENAI_SHAPE.get(provider, {}).get(
|
||||
status_code, OPENAI_SHAPED[status_code]
|
||||
)
|
||||
return DEVIATIONS_FROM_THE_OPENAI_SHAPE.get(provider, {}).get(status_code, OPENAI_SHAPED[status_code])
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -910,9 +879,7 @@ def quiet_exception_mapping(monkeypatch):
|
|||
|
||||
@pytest.mark.parametrize("status_code", UPSTREAM_STATUS_CODES)
|
||||
@pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER)
|
||||
def test_an_upstream_status_maps_to_one_exception_per_provider(
|
||||
provider, status_code, quiet_exception_mapping
|
||||
):
|
||||
def test_an_upstream_status_maps_to_one_exception_per_provider(provider, status_code, quiet_exception_mapping):
|
||||
expected_class, expected_status = _expected_for(provider, status_code)
|
||||
|
||||
with pytest.raises(openai.APIError) as raised:
|
||||
|
|
@ -928,9 +895,7 @@ def test_an_upstream_status_maps_to_one_exception_per_provider(
|
|||
|
||||
@pytest.mark.parametrize("status_code", UPSTREAM_STATUS_CODES)
|
||||
@pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER)
|
||||
def test_a_mapped_exception_keeps_the_provider_and_model_it_came_from(
|
||||
provider, status_code, quiet_exception_mapping
|
||||
):
|
||||
def test_a_mapped_exception_keeps_the_provider_and_model_it_came_from(provider, status_code, quiet_exception_mapping):
|
||||
with pytest.raises(openai.APIError) as raised:
|
||||
exception_type(
|
||||
model="test-model",
|
||||
|
|
@ -943,12 +908,8 @@ def test_a_mapped_exception_keeps_the_provider_and_model_it_came_from(
|
|||
|
||||
|
||||
@pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER)
|
||||
def test_an_already_mapped_litellm_exception_passes_through_untouched(
|
||||
provider, quiet_exception_mapping
|
||||
):
|
||||
already_mapped = litellm.RateLimitError(
|
||||
message="already mapped", llm_provider=provider, model="test-model"
|
||||
)
|
||||
def test_an_already_mapped_litellm_exception_passes_through_untouched(provider, quiet_exception_mapping):
|
||||
already_mapped = litellm.RateLimitError(message="already mapped", llm_provider=provider, model="test-model")
|
||||
|
||||
returned = exception_type(
|
||||
model="test-model",
|
||||
|
|
@ -961,9 +922,7 @@ def test_an_already_mapped_litellm_exception_passes_through_untouched(
|
|||
|
||||
@pytest.mark.parametrize("status_code", UPSTREAM_STATUS_CODES)
|
||||
@pytest.mark.parametrize("provider", PROVIDERS_WITHOUT_A_HANDLER)
|
||||
def test_a_provider_without_a_handler_maps_by_the_upstream_status(
|
||||
provider, status_code, quiet_exception_mapping
|
||||
):
|
||||
def test_a_provider_without_a_handler_maps_by_the_upstream_status(provider, status_code, quiet_exception_mapping):
|
||||
expected_class, expected_status = STATUS_KEYED[status_code]
|
||||
|
||||
with pytest.raises(openai.APIError) as raised:
|
||||
|
|
@ -1015,9 +974,7 @@ def test_an_unmapped_exception_with_no_model_or_provider_is_a_connection_error(q
|
|||
assert "boom" in raised.value.message
|
||||
|
||||
|
||||
def _raise_and_map(
|
||||
model: str | None, original_exception: Exception, custom_llm_provider: str | None
|
||||
) -> None:
|
||||
def _raise_and_map(model: str | None, original_exception: Exception, custom_llm_provider: str | None) -> None:
|
||||
"""Calls exception_type() from inside the except block, as litellm/main.py does,
|
||||
so traceback.format_exc() has a real stack."""
|
||||
try:
|
||||
|
|
@ -1058,9 +1015,7 @@ def test_an_unmapped_exception_with_no_model_or_provider_message_keeps_traceback
|
|||
|
||||
|
||||
CONTEXT_WINDOW_MESSAGE = "This model's maximum context length is 4096 tokens."
|
||||
CONTENT_POLICY_MESSAGE = (
|
||||
'{"error": {"type": "invalid_request_error", "code": "content_policy_violation"}}'
|
||||
)
|
||||
CONTENT_POLICY_MESSAGE = '{"error": {"type": "invalid_request_error", "code": "content_policy_violation"}}'
|
||||
TIMEOUT_MESSAGE = "Request timed out."
|
||||
|
||||
PROVIDERS_THAT_RECOGNISE_A_FULL_CONTEXT_WINDOW = (
|
||||
|
|
@ -1103,15 +1058,11 @@ class _UpstreamErrorWithMessage(_UpstreamHTTPError):
|
|||
super().__init__(status_code=status_code)
|
||||
self.args = (message,)
|
||||
self.message = message
|
||||
self.response = httpx.Response(
|
||||
status_code=status_code, request=self.request, text=message
|
||||
)
|
||||
self.response = httpx.Response(status_code=status_code, request=self.request, text=message)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER)
|
||||
def test_a_full_context_window_reaches_the_caller_as_the_router_needs_it(
|
||||
provider, quiet_exception_mapping
|
||||
):
|
||||
def test_a_full_context_window_reaches_the_caller_as_the_router_needs_it(provider, quiet_exception_mapping):
|
||||
if provider in PROVIDERS_THAT_RECOGNISE_A_FULL_CONTEXT_WINDOW:
|
||||
expected_class, expected_status = litellm.ContextWindowExceededError, 400
|
||||
else:
|
||||
|
|
@ -1129,9 +1080,7 @@ def test_a_full_context_window_reaches_the_caller_as_the_router_needs_it(
|
|||
|
||||
|
||||
@pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER)
|
||||
def test_a_content_policy_block_reaches_the_caller_as_the_router_needs_it(
|
||||
provider, quiet_exception_mapping
|
||||
):
|
||||
def test_a_content_policy_block_reaches_the_caller_as_the_router_needs_it(provider, quiet_exception_mapping):
|
||||
if provider in PROVIDERS_THAT_RECOGNISE_A_CONTENT_POLICY_BLOCK:
|
||||
expected_class, expected_status = litellm.ContentPolicyViolationError, 400
|
||||
else:
|
||||
|
|
@ -1149,9 +1098,7 @@ def test_a_content_policy_block_reaches_the_caller_as_the_router_needs_it(
|
|||
|
||||
|
||||
@pytest.mark.parametrize("provider", PROVIDERS_WITH_A_HANDLER)
|
||||
def test_a_timed_out_request_is_a_timeout_for_every_provider(
|
||||
provider, quiet_exception_mapping
|
||||
):
|
||||
def test_a_timed_out_request_is_a_timeout_for_every_provider(provider, quiet_exception_mapping):
|
||||
with pytest.raises(litellm.Timeout) as raised:
|
||||
exception_type(
|
||||
model="test-model",
|
||||
|
|
@ -1409,3 +1356,97 @@ def test_bedrock_timeout_mapping_keeps_retry_after_readable(status_code):
|
|||
exception_headers = _get_response_headers(original_exception=exc_info.value)
|
||||
assert exception_headers is not None
|
||||
assert litellm.utils._get_retry_after_from_exception_header(response_headers=exception_headers) == 7
|
||||
|
||||
|
||||
_GUARDRAIL_BLOCK_ERROR = {
|
||||
"message": "Content blocked: secret_project_codename pattern detected",
|
||||
"param": "None",
|
||||
"code": "400",
|
||||
"provider_specific_fields": {
|
||||
"error": "Content blocked: secret_project_codename pattern detected",
|
||||
"pattern": "secret_project_codename",
|
||||
"guardrail_name": "block-secret-project",
|
||||
"guardrail_mode": "pre_call",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _openai_handler_error(
|
||||
error_type: str,
|
||||
headers: dict[str, str] | list[tuple[str, str]],
|
||||
status_code: int = 400,
|
||||
message: str = _GUARDRAIL_BLOCK_ERROR["message"],
|
||||
) -> OpenAIError:
|
||||
wire_error = {**_GUARDRAIL_BLOCK_ERROR, "type": error_type, "code": str(status_code), "message": message}
|
||||
return OpenAIError(
|
||||
status_code=status_code,
|
||||
message=f"Error code: {status_code} - {{'error': {wire_error}}}",
|
||||
headers=httpx.Headers(headers),
|
||||
body=wire_error,
|
||||
)
|
||||
|
||||
|
||||
_PROXY_HEADERS = {"x-litellm-call-id": "call-guardrail", "x-litellm-applied-guardrails": "block-secret-project"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("error_type", "status_code"), [("None", 400), ("invalid_request_error", 400), ("None", 422)])
|
||||
def test_litellm_proxy_guardrail_block_keeps_body_and_headers(error_type: str, status_code: int):
|
||||
with pytest.raises(litellm.BadRequestError) as exc_info:
|
||||
exception_type(
|
||||
model="claude-haiku-4-5",
|
||||
original_exception=_openai_handler_error(error_type, _PROXY_HEADERS, status_code=status_code),
|
||||
custom_llm_provider="litellm_proxy",
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
|
||||
assert exc_info.value.body["provider_specific_fields"]["guardrail_name"] == "block-secret-project"
|
||||
assert exc_info.value.body["type"] == error_type
|
||||
assert dict(exc_info.value.response.headers) == _PROXY_HEADERS
|
||||
|
||||
|
||||
@pytest.mark.parametrize("relayed_class", [litellm.BadRequestError, litellm.ContentPolicyViolationError])
|
||||
def test_litellm_proxy_relayed_litellm_error_keeps_body_and_headers(relayed_class: type[litellm.BadRequestError]):
|
||||
message = f"litellm.{relayed_class.__name__}: {_GUARDRAIL_BLOCK_ERROR['message']}"
|
||||
|
||||
with pytest.raises(relayed_class) as exc_info:
|
||||
exception_type(
|
||||
model="claude-haiku-4-5",
|
||||
original_exception=_openai_handler_error("None", _PROXY_HEADERS, message=message),
|
||||
custom_llm_provider="litellm_proxy",
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
|
||||
assert type(exc_info.value) is relayed_class
|
||||
assert exc_info.value.body["provider_specific_fields"]["guardrail_name"] == "block-secret-project"
|
||||
assert dict(exc_info.value.response.headers) == _PROXY_HEADERS
|
||||
|
||||
|
||||
def test_openai_compatible_vendor_400_keeps_body_but_not_headers():
|
||||
with pytest.raises(litellm.BadRequestError) as exc_info:
|
||||
exception_type(
|
||||
model="gpt-5.4-mini",
|
||||
original_exception=_openai_handler_error("vendor_specific_error", {"openai-organization": "org-1"}),
|
||||
custom_llm_provider="openai",
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
|
||||
assert exc_info.value.body["type"] == "vendor_specific_error"
|
||||
assert not exc_info.value.response.headers
|
||||
|
||||
|
||||
def test_litellm_proxy_repeated_response_header_keeps_each_value():
|
||||
repeated = [("x-litellm-call-id", "call-guardrail"), ("set-cookie", "a=1"), ("set-cookie", "b=2")]
|
||||
|
||||
with pytest.raises(litellm.BadRequestError) as exc_info:
|
||||
exception_type(
|
||||
model="claude-haiku-4-5",
|
||||
original_exception=_openai_handler_error("None", repeated),
|
||||
custom_llm_provider="litellm_proxy",
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
|
||||
assert exc_info.value.response.headers.multi_items() == repeated
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import pytest
|
|||
|
||||
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import StreamingScanKey
|
||||
from litellm.llms.anthropic.chat.guardrail_translation.handler import (
|
||||
AnthropicMessagesHandler,
|
||||
|
|
@ -635,14 +636,19 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.inputs is not None
|
||||
assert guardrail.inputs["texts"] == ["safe text", "prohibited correction"]
|
||||
assert guardrail.inputs["texts"] == [
|
||||
"trusted top-level system prompt",
|
||||
"safe text",
|
||||
"prohibited correction",
|
||||
]
|
||||
structured = guardrail.inputs["structured_messages"]
|
||||
assert [m["role"] for m in structured] == ["system", "user", "system"]
|
||||
assert structured[0]["content"] == "trusted top-level system prompt"
|
||||
assert data["system"] == "trusted top-level system prompt"
|
||||
assert data["messages"][1]["content"] == "[MASKED]"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_masking_slice_is_unavailable_when_top_level_system_is_included(
|
||||
async def test_bedrock_masking_slice_lines_up_when_top_level_system_is_included(
|
||||
self,
|
||||
):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
|
||||
|
|
@ -668,25 +674,25 @@ class TestAnthropicMessagesHandlerInputProcessing:
|
|||
structured = guardrail.inputs["structured_messages"]
|
||||
|
||||
bedrock = BedrockGuardrail(guardrailIdentifier="gi", guardrailVersion="1")
|
||||
assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts) + 1
|
||||
assert sum(bedrock._count_message_texts(m) for m in structured) == len(texts)
|
||||
latest_user_index = bedrock._find_latest_message_index(structured, target_role="user")
|
||||
assert (
|
||||
bedrock._locate_message_texts_slice(
|
||||
structured_messages=structured,
|
||||
target_index=latest_user_index,
|
||||
texts=texts,
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert (
|
||||
bedrock._merge_masked_texts(
|
||||
masked_texts=["{MASKED}"],
|
||||
texts=texts,
|
||||
scanned_slice=None,
|
||||
scanned_role_subset=True,
|
||||
)
|
||||
== texts
|
||||
scanned_slice = bedrock._locate_message_texts_slice(
|
||||
structured_messages=structured,
|
||||
target_index=latest_user_index,
|
||||
texts=texts,
|
||||
)
|
||||
assert scanned_slice == (3, 1)
|
||||
assert bedrock._merge_masked_texts(
|
||||
masked_texts=["{MASKED}"],
|
||||
texts=texts,
|
||||
scanned_slice=scanned_slice,
|
||||
scanned_role_subset=True,
|
||||
) == [
|
||||
"trusted top-level system prompt",
|
||||
"safe text",
|
||||
"prohibited correction",
|
||||
"{MASKED}",
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("skip_system_message_in_guardrail", [True, None])
|
||||
|
|
@ -1611,7 +1617,8 @@ class TestAnthropicMessagesIncrementalScan:
|
|||
)
|
||||
assert mock_api.call_count == 1
|
||||
assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [
|
||||
"What is the capital of France?"
|
||||
"You are a helpful geography assistant.",
|
||||
"What is the capital of France?",
|
||||
]
|
||||
mock_api.reset_mock()
|
||||
await handler.process_input_messages(
|
||||
|
|
@ -2150,6 +2157,213 @@ class TestAnthropicMessagesScanOnlyToolResults:
|
|||
assert guardrail.captured_inputs.get("images") == ["TOOL_IMG"]
|
||||
|
||||
|
||||
class ToolCallArgumentsMaskingGuardrail(InputsRecordingGuardrail):
|
||||
"""Masks the canary inside tool-call arguments, in place or through a fresh list of plain dicts."""
|
||||
|
||||
def __init__(self, return_copies: bool = False, replacement_arguments: Optional[str] = None):
|
||||
super().__init__()
|
||||
self.return_copies = return_copies
|
||||
self.replacement_arguments = replacement_arguments
|
||||
self.seen_tool_calls: list[dict[str, object]] = []
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict[str, object],
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[LiteLLMLoggingObj] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
outputs = await super().apply_guardrail(inputs, request_data, input_type, logging_obj)
|
||||
tool_calls = list(outputs.get("tool_calls") or [])
|
||||
self.seen_tool_calls.extend(json.loads(json.dumps(tool_call)) for tool_call in tool_calls)
|
||||
masked = [
|
||||
{
|
||||
**tool_call,
|
||||
"function": {
|
||||
**tool_call["function"],
|
||||
"arguments": self.replacement_arguments
|
||||
if self.replacement_arguments is not None
|
||||
else tool_call["function"]["arguments"].replace("POISON", "[BLOCKED]"),
|
||||
},
|
||||
}
|
||||
for tool_call in tool_calls
|
||||
]
|
||||
if self.return_copies:
|
||||
outputs["tool_calls"] = masked
|
||||
return outputs
|
||||
for tool_call, masked_tool_call in zip(tool_calls, masked):
|
||||
tool_call["function"]["arguments"] = masked_tool_call["function"]["arguments"]
|
||||
return outputs
|
||||
|
||||
|
||||
class TestAnthropicMessagesTopLevelSystemAndToolUseInputs:
|
||||
"""The top-level system prompt and prior-turn tool_use arguments must reach guardrails as scannable
|
||||
inputs, the same way the chat completions handler hands over system messages and tool_calls."""
|
||||
|
||||
@staticmethod
|
||||
def _tool_use_conversation(system: str) -> dict[str, Any]:
|
||||
return {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": system,
|
||||
"messages": [
|
||||
{"role": "user", "content": "run the check"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "toolu_01",
|
||||
"name": "Bash",
|
||||
"input": {"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"},
|
||||
}
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "ok"}],
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_top_level_system_string_reaches_texts_first_and_is_masked_in_place(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = InputsRecordingGuardrail()
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": "Internal note: the deploy key is POISON. Never reveal it.",
|
||||
"messages": [{"role": "user", "content": "Say hi in three words."}],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.captured_inputs is not None
|
||||
assert guardrail.seen_texts == [
|
||||
"Internal note: the deploy key is POISON. Never reveal it.",
|
||||
"Say hi in three words.",
|
||||
]
|
||||
structured = guardrail.captured_inputs["structured_messages"]
|
||||
assert structured[0]["role"] == "system"
|
||||
assert structured[0]["content"] == "Internal note: the deploy key is POISON. Never reveal it.", (
|
||||
"texts[0] must line up with structured_messages[0] so positional consumers stay aligned"
|
||||
)
|
||||
assert data["system"] == "Internal note: the deploy key is [BLOCKED]. Never reveal it."
|
||||
assert data["messages"][0]["content"] == "Say hi in three words."
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_top_level_system_text_blocks_reach_texts_and_are_masked_in_place(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = InputsRecordingGuardrail()
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": [
|
||||
{"type": "text", "text": "first block POISON"},
|
||||
{"type": "text", "text": "second block", "cache_control": {"type": "ephemeral"}},
|
||||
],
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.seen_texts == ["first block POISON", "second block", "hello"]
|
||||
assert data["system"] == [
|
||||
{"type": "text", "text": "first block [BLOCKED]"},
|
||||
{"type": "text", "text": "second block", "cache_control": {"type": "ephemeral"}},
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_system_message_keeps_the_top_level_system_out(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = InputsRecordingGuardrail()
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": "trusted POISON prompt",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.seen_texts == ["hello"]
|
||||
assert data["system"] == "trusted POISON prompt"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prior_turn_tool_use_input_reaches_tool_calls_in_openai_shape(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = InputsRecordingGuardrail()
|
||||
data = self._tool_use_conversation(system="You are a careful agent harness.")
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.captured_inputs is not None
|
||||
tool_calls = guardrail.captured_inputs.get("tool_calls")
|
||||
assert tool_calls is not None and len(tool_calls) == 1
|
||||
assert tool_calls[0]["id"] == "toolu_01"
|
||||
assert tool_calls[0]["type"] == "function"
|
||||
assert tool_calls[0]["function"]["name"] == "Bash"
|
||||
assert json.loads(tool_calls[0]["function"]["arguments"]) == {
|
||||
"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"
|
||||
}
|
||||
assert data["messages"][1]["content"][0]["input"] == {
|
||||
"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"
|
||||
}, "a guardrail that leaves tool_calls alone must leave the tool_use input alone"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("return_copies", [False, True])
|
||||
async def test_masked_tool_call_arguments_write_back_into_the_tool_use_input(self, return_copies: bool):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = ToolCallArgumentsMaskingGuardrail(return_copies=return_copies)
|
||||
data = self._tool_use_conversation(system="You are a careful agent harness.")
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert [tool_call["function"]["name"] for tool_call in guardrail.seen_tool_calls] == ["Bash"]
|
||||
tool_use = data["messages"][1]["content"][0]
|
||||
assert tool_use == {
|
||||
"type": "tool_use",
|
||||
"id": "toolu_01",
|
||||
"name": "Bash",
|
||||
"input": {"cmd": "AWS_ACCESS_KEY_ID=[BLOCKED] aws sts get-caller-identity"},
|
||||
}
|
||||
assert data["messages"][2]["content"][0]["tool_use_id"] == "toolu_01"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_json_rewritten_arguments_are_rejected_by_name(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = ToolCallArgumentsMaskingGuardrail(replacement_arguments="[REDACTED]")
|
||||
data = self._tool_use_conversation(system="Internal note: the deploy key is POISON. Never reveal it.")
|
||||
data["messages"][2]["content"][0]["content"] = "fetched POISON page"
|
||||
original = json.loads(json.dumps(data))
|
||||
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "scan-only-capture"
|
||||
assert data["system"] == original["system"], "a rejected rewrite must leave the request untouched"
|
||||
assert data["messages"] == original["messages"], "a rejected rewrite must leave the request untouched"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scan_only_tool_results_keeps_system_and_tool_use_out(self):
|
||||
handler = AnthropicMessagesHandler()
|
||||
guardrail = InputsRecordingGuardrail()
|
||||
guardrail.scan_only_tool_results = True
|
||||
data = self._tool_use_conversation(system="trusted POISON prompt")
|
||||
data["messages"][2]["content"][0]["content"] = "fetched POISON page"
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.seen_texts == ["fetched POISON page"]
|
||||
assert guardrail.captured_inputs is not None
|
||||
assert guardrail.captured_inputs.get("tool_calls") is None
|
||||
assert data["system"] == "trusted POISON prompt"
|
||||
assert data["messages"][1]["content"][0]["input"] == {
|
||||
"cmd": "AWS_ACCESS_KEY_ID=POISON aws sts get-caller-identity"
|
||||
}
|
||||
assert data["messages"][2]["content"][0]["content"] == "fetched [BLOCKED] page"
|
||||
|
||||
|
||||
class TestStructuredWriteBackKeepsToolResults:
|
||||
"""A guardrail rewrite must never leave a tool_use without its tool_result (Claude Code ToolSearch, LIT-6103)."""
|
||||
|
||||
|
|
@ -2272,6 +2486,116 @@ class TestAnthropicMessagesHandlerStreamingScanKey:
|
|||
assert ended_key != open_key
|
||||
|
||||
|
||||
class PerRowTextGuardrail(CustomGuardrail):
|
||||
"""Answers one redacted text per chat row it was shown, the way a guardrail
|
||||
that scans per message does, and hands back only texts."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="per-row-redactor")
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
rows = inputs.get("structured_messages") or []
|
||||
return {**inputs, "texts": [str(row.get("content")).replace("123-45-6789", "<US_SSN>") for row in rows]}
|
||||
|
||||
|
||||
class PerSlotTextGuardrail(CustomGuardrail):
|
||||
"""Answers one redacted text per text slot of every chat row it was shown, the
|
||||
way a guardrail that counts slots per message does, and hands back only texts."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="per-slot-redactor")
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import message_slot_texts
|
||||
|
||||
rows = inputs.get("structured_messages") or []
|
||||
return {
|
||||
**inputs,
|
||||
"texts": [text.replace("123-45-6789", "<US_SSN>") for row in rows for text in message_slot_texts(row)],
|
||||
}
|
||||
|
||||
|
||||
class TestPerMessageTextWriteBack:
|
||||
"""Texts that no longer pair one-to-one with what the handler extracted must be
|
||||
rejected by name instead of sliding onto the wrong messages."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_text_per_row_over_a_system_prompt_is_applied(self):
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": "Reply with exactly the SSN you were given.",
|
||||
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
|
||||
}
|
||||
|
||||
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail())
|
||||
|
||||
assert data["system"] == "Reply with exactly the SSN you were given."
|
||||
assert data["messages"] == [{"role": "user", "content": "My SSN is <US_SSN>."}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_text_per_row_over_a_multi_block_system_prompt_is_rejected_by_name(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": [
|
||||
{"type": "text", "text": "Reply with exactly the SSN you were given."},
|
||||
{"type": "text", "text": "Never apologize."},
|
||||
],
|
||||
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
|
||||
}
|
||||
original = json.loads(json.dumps(data))
|
||||
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail())
|
||||
|
||||
assert excinfo.value.guardrail_name == "per-row-redactor"
|
||||
assert data["system"] == original["system"], "a rejected rewrite must leave the request untouched"
|
||||
assert data["messages"] == original["messages"], "a rejected rewrite must leave the request untouched"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_text_per_slot_over_a_system_prompt_with_an_empty_block_is_applied(self):
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"system": [
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "text", "text": "Reply with exactly the SSN you were given."},
|
||||
],
|
||||
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
|
||||
}
|
||||
|
||||
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerSlotTextGuardrail())
|
||||
|
||||
assert data["system"] == [
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "text", "text": "Reply with exactly the SSN you were given."},
|
||||
]
|
||||
assert data["messages"] == [{"role": "user", "content": "My SSN is <US_SSN>."}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_text_per_row_without_a_system_prompt_is_applied(self):
|
||||
data = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "My SSN is 123-45-6789."}],
|
||||
}
|
||||
|
||||
await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=PerRowTextGuardrail())
|
||||
|
||||
assert data["messages"] == [{"role": "user", "content": "My SSN is <US_SSN>."}]
|
||||
|
||||
|
||||
class TestAnthropicMessagesHandlerPostCallHookResponse:
|
||||
def test_openai_shaped_stream_assembly_reaches_the_hook_as_a_messages_response(self):
|
||||
from litellm.types.utils import Choices, Message, ModelResponse, Usage
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ Without the fix, the AnthropicStreamWrapper silently dropped these
|
|||
arguments, causing tool_use blocks to arrive with empty input {}.
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
from typing import List
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
|
@ -139,9 +141,7 @@ async def test_async_stream_emits_input_json_delta_for_bundled_tool_args():
|
|||
|
||||
# Verify the delta carries the tool arguments
|
||||
delta_event = events[input_json_delta_idx]
|
||||
assert delta_event["delta"][
|
||||
"partial_json"
|
||||
], "input_json_delta should have non-empty partial_json"
|
||||
assert json.loads(delta_event["delta"]["partial_json"]) == {"location": "Boston"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -300,7 +300,7 @@ def test_sync_stream_emits_input_json_delta_for_bundled_tool_args():
|
|||
assert (
|
||||
input_json_delta_idx == tool_start_idx + 1
|
||||
), "input_json_delta should immediately follow the tool_use content_block_start"
|
||||
assert events[input_json_delta_idx]["delta"]["partial_json"]
|
||||
assert json.loads(events[input_json_delta_idx]["delta"]["partial_json"]) == {"location": "Boston"}
|
||||
|
||||
|
||||
def test_sync_stream_no_extra_delta_when_tool_args_empty():
|
||||
|
|
|
|||
|
|
@ -32,9 +32,14 @@ action.
|
|||
import base64
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.llms.bedrock.base_aws_llm import WebIdentitySessionPolicy, _SessionPolicyStatement
|
||||
|
||||
# Actions the Claude Platform on AWS service is documented to call.
|
||||
# Source: AWS IAM action reference + the #27678 surface area.
|
||||
|
|
@ -49,9 +54,9 @@ _CLAUDE_PLATFORM_ACTIONS = {
|
|||
}
|
||||
|
||||
|
||||
def _captured_policy() -> dict:
|
||||
"""Run _auth_with_web_identity_token under mocks + return the parsed
|
||||
Policy dict that was actually sent to STS."""
|
||||
def _captured_policy_document() -> str:
|
||||
"""Run _auth_with_web_identity_token under mocks + return the Policy
|
||||
JSON document that was actually sent to STS."""
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
base = BaseAWSLLM()
|
||||
|
|
@ -84,11 +89,21 @@ def _captured_policy() -> dict:
|
|||
|
||||
mock_sts.assume_role_with_web_identity.assert_called_once()
|
||||
kwargs = mock_sts.assume_role_with_web_identity.call_args.kwargs
|
||||
policy_str = kwargs["Policy"]
|
||||
return json.loads(policy_str)
|
||||
return kwargs["Policy"]
|
||||
|
||||
|
||||
def _statement_by_sid(policy: dict, sid: str) -> dict:
|
||||
_SESSION_POLICY_ADAPTER: Final = TypeAdapter(WebIdentitySessionPolicy)
|
||||
|
||||
|
||||
def _captured_policy() -> WebIdentitySessionPolicy:
|
||||
return _SESSION_POLICY_ADAPTER.validate_python(json.loads(_captured_policy_document()))
|
||||
|
||||
|
||||
def _granted_actions(policy: WebIdentitySessionPolicy) -> frozenset[str]:
|
||||
return frozenset(action for stmt in policy["Statement"] for action in stmt["Action"])
|
||||
|
||||
|
||||
def _statement_by_sid(policy: WebIdentitySessionPolicy, sid: str) -> _SessionPolicyStatement:
|
||||
for stmt in policy["Statement"]:
|
||||
if stmt.get("Sid") == sid:
|
||||
return stmt
|
||||
|
|
@ -102,7 +117,6 @@ class TestWebIdentitySessionPolicyShape:
|
|||
def test_policy_parses_as_valid_iam_document(self):
|
||||
policy = _captured_policy()
|
||||
assert policy["Version"] == "2012-10-17"
|
||||
assert isinstance(policy["Statement"], list)
|
||||
assert len(policy["Statement"]) >= 2
|
||||
|
||||
def test_bedrock_statement_actions_preserved(self):
|
||||
|
|
@ -137,16 +151,7 @@ class TestClaudePlatformActionsCovered:
|
|||
|
||||
@pytest.mark.parametrize("action", sorted(_CLAUDE_PLATFORM_ACTIONS))
|
||||
def test_claude_platform_action_present(self, action: str):
|
||||
policy = _captured_policy()
|
||||
# Action may live in any Statement — search across all.
|
||||
all_actions: set = set()
|
||||
for stmt in policy["Statement"]:
|
||||
stmt_actions = stmt.get("Action")
|
||||
if isinstance(stmt_actions, str):
|
||||
all_actions.add(stmt_actions)
|
||||
elif isinstance(stmt_actions, list):
|
||||
all_actions.update(stmt_actions)
|
||||
assert action in all_actions, (
|
||||
assert action in _granted_actions(_captured_policy()), (
|
||||
f"{action} missing from session policy — "
|
||||
f"bedrock/claude_platform/* requests will 403 on OIDC auth"
|
||||
)
|
||||
|
|
@ -179,15 +184,7 @@ class TestBedrockMantleActionsCovered:
|
|||
action" even when the role's identity policy grants it."""
|
||||
|
||||
def test_bedrock_mantle_create_inference_present(self):
|
||||
policy = _captured_policy()
|
||||
all_actions: set = set()
|
||||
for stmt in policy["Statement"]:
|
||||
stmt_actions = stmt.get("Action")
|
||||
if isinstance(stmt_actions, str):
|
||||
all_actions.add(stmt_actions)
|
||||
elif isinstance(stmt_actions, list):
|
||||
all_actions.update(stmt_actions)
|
||||
assert "bedrock-mantle:CreateInference" in all_actions, (
|
||||
assert "bedrock-mantle:CreateInference" in _granted_actions(_captured_policy()), (
|
||||
"bedrock-mantle:CreateInference missing from session policy — "
|
||||
"bedrock_mantle/* requests will 403 on OIDC/WIF auth"
|
||||
)
|
||||
|
|
@ -233,7 +230,7 @@ class TestInvalidIdentityTokenSurfacesAudience:
|
|||
operator can diagnose the mismatch without enabling LITELLM_LOG=DEBUG on a
|
||||
prod instance."""
|
||||
|
||||
_AUD = "https://guidepoint.litellm-prod.ai"
|
||||
_AUD = "https://gateway.example.com"
|
||||
_ISS = "https://accounts.google.com"
|
||||
_STS_MESSAGE = (
|
||||
"An error occurred (InvalidIdentityToken) when calling the "
|
||||
|
|
@ -308,3 +305,44 @@ class TestPolicyTransportConditions:
|
|||
"ClaudePlatformLiteLLM must require aws:SecureTransport=true "
|
||||
"to keep parity with the bedrock statement"
|
||||
)
|
||||
|
||||
|
||||
_STS_SESSION_POLICY_PLAINTEXT_LIMIT: Final = 2048
|
||||
|
||||
_BEDROCK_ROUTE_ACTIONS: Final = MappingProxyType(
|
||||
{
|
||||
"model/{model_id}/invoke": "bedrock:InvokeModel",
|
||||
"model/{model_id}/invoke-with-response-stream": "bedrock:InvokeModelWithResponseStream",
|
||||
"model/{model_id}/converse": "bedrock:InvokeModel",
|
||||
"model/{model_id}/converse-stream": "bedrock:InvokeModelWithResponseStream",
|
||||
"model/{model_id}/count-tokens": "bedrock:CountTokens",
|
||||
"guardrail/{guardrail_id}/version/{version}/apply": "bedrock:ApplyGuardrail",
|
||||
"rerank": "bedrock:Rerank",
|
||||
"knowledgebases/{knowledge_base_id}/retrieve": "bedrock:Retrieve",
|
||||
"knowledgebases": "bedrock:ListKnowledgeBases",
|
||||
"agents/{agent_id}/agentAliases/{alias_id}/sessions/{session_id}/text": "bedrock:InvokeAgent",
|
||||
"runtimes/{agent_runtime_arn}/invocations": "bedrock-agentcore:InvokeAgentRuntime",
|
||||
"runtimes/{agent_runtime_arn}/invocations with X-Amzn-Bedrock-AgentCore-Runtime-User-Id": (
|
||||
"bedrock-agentcore:InvokeAgentRuntimeForUser"
|
||||
),
|
||||
"mcp": "bedrock-agentcore:InvokeGateway",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class TestSessionPolicyGrantsEveryBedrockRoute:
|
||||
"""LIT-7348: ``/rerank`` authorizes against ``bedrock:Rerank``, which the
|
||||
ceiling never granted, so rerank 403d on web identity auth while static
|
||||
credentials and IRSA worked. Each route the bedrock package signs with the
|
||||
web identity session maps to the IAM action it authorizes against, and the
|
||||
ceiling must grant every one of them."""
|
||||
|
||||
@pytest.mark.parametrize(("route", "action"), sorted(_BEDROCK_ROUTE_ACTIONS.items()))
|
||||
def test_route_action_is_granted_by_the_ceiling(self, route: str, action: str):
|
||||
assert action in _granted_actions(_captured_policy()), (
|
||||
f"/{route} authorizes against {action}, which the session policy does not grant, "
|
||||
"so it 403s on web identity auth"
|
||||
)
|
||||
|
||||
def test_policy_document_fits_the_sts_plaintext_limit(self):
|
||||
assert len(_captured_policy_document()) <= _STS_SESSION_POLICY_PLAINTEXT_LIMIT
|
||||
|
|
|
|||
|
|
@ -773,6 +773,12 @@ class TestBedrockMantleCodexAdditionalTools:
|
|||
assert body["input"] == codex_agentic_items
|
||||
assert "tools" not in body
|
||||
|
||||
def test_input_without_additional_tools_sanitizes_tools_on_the_caller_params_object(self):
|
||||
params = {"tools": [{"type": "function", "name": "wait", "parameters": '{"type": "object"}'}]}
|
||||
body = self._transform(input=[self._USER_MESSAGE], params=params)
|
||||
assert body["tools"][0]["parameters"] == {"type": "object"}
|
||||
assert params["tools"][0]["parameters"] == {"type": "object"}
|
||||
|
||||
def test_malformed_additional_tools_item_without_tools_list_is_stripped(self):
|
||||
body = self._transform(
|
||||
input=[
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import cast
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
|
@ -6,6 +8,7 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig
|
||||
from litellm.types.llms.gemini import BidiGenerateContentServerMessage
|
||||
|
||||
|
||||
def test_gemini_realtime_transformation_session_created():
|
||||
|
|
@ -2178,3 +2181,71 @@ def test_unbilled_usage_on_session_close_flushes_trailing_audio(patch_gemini_tra
|
|||
}
|
||||
assert usage == expected
|
||||
assert config.unbilled_usage_on_session_close("gemini-3.5-transcribe-live") is None
|
||||
|
||||
|
||||
def _grounded_live_frame(grounding_metadata: Mapping[str, object] | None) -> Mapping[str, object]:
|
||||
"""One Live server frame. Grounding metadata and usageMetadata arrive together, as Vertex sends them."""
|
||||
from typing import Final
|
||||
|
||||
server_content: Final = {
|
||||
"turnComplete": True,
|
||||
**({} if grounding_metadata is None else {"groundingMetadata": grounding_metadata}),
|
||||
}
|
||||
return {
|
||||
"serverContent": server_content,
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 19,
|
||||
"candidatesTokenCount": 157,
|
||||
"totalTokenCount": 176,
|
||||
"promptTokensDetails": ({"modality": "TEXT", "tokenCount": 19},),
|
||||
"candidatesTokensDetails": ({"modality": "AUDIO", "tokenCount": 157},),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _response_done_input_details(message: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""The ``input_tokens_details`` a ``response.done`` event carries, read off the emitted event."""
|
||||
from typing import Final
|
||||
|
||||
config: Final = GeminiRealtimeConfig()
|
||||
event: Final = config.transform_response_done_event(
|
||||
message=cast( # cast-ok: a test fixture stands in for the server frame TypedDict
|
||||
BidiGenerateContentServerMessage, message
|
||||
),
|
||||
current_response_id="resp_grounding",
|
||||
current_conversation_id="conv_grounding",
|
||||
output_items=None,
|
||||
)
|
||||
usage: Final = event["response"]["usage"]
|
||||
assert usage, "response.done must carry a usage object"
|
||||
return usage.get("input_tokens_details") or {}
|
||||
|
||||
|
||||
def test_gemini_realtime_response_done_counts_web_grounding():
|
||||
"""Regression: Live reports grounding in the server frames and never in usageMetadata.
|
||||
|
||||
Nothing read those frames on the realtime path, so web_search_requests stayed unset and the
|
||||
cost path's only trigger for Google's per-query grounding charge never fired.
|
||||
|
||||
The counter is read off the emitted event, which is what the cost path is handed, so this covers
|
||||
the grounding read and the usage bridge that carries it together
|
||||
"""
|
||||
input_details = _response_done_input_details(
|
||||
_grounded_live_frame(
|
||||
{
|
||||
"webSearchQueries": ["who won the 2026 world cup final"],
|
||||
"groundingChunks": [{"web": {"uri": "https://example.com"}}],
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert input_details.get("web_search_requests") == 1, "a grounded turn must report its query"
|
||||
assert input_details.get("text_tokens") == 19, "the modality breakdown must survive alongside it"
|
||||
|
||||
|
||||
def test_gemini_realtime_response_done_reports_no_grounding_when_none_ran():
|
||||
"""The counter must stay unset on an ordinary turn, or every session pays a grounding fee."""
|
||||
input_details = _response_done_input_details(_grounded_live_frame(None))
|
||||
|
||||
assert input_details.get("web_search_requests") is None
|
||||
assert input_details.get("google_maps_grounding_requests") is None
|
||||
|
|
|
|||
|
|
@ -1893,6 +1893,183 @@ class TestScanOnlyToolResults:
|
|||
assert data["messages"][4]["content"] == "and then?"
|
||||
|
||||
|
||||
class TestNoScannableContentRecordsNotRun:
|
||||
"""LIT-6314: a guardrail whose scoping leaves nothing to scan must still persist an evaluation record"""
|
||||
|
||||
def _system_only_data(self) -> dict:
|
||||
return {"messages": [{"role": "system", "content": "SYSTEM-PROMPT"}]}
|
||||
|
||||
def _recorded_entries(self, data: dict) -> list:
|
||||
metadata = data.get("metadata") or data.get("litellm_metadata") or {}
|
||||
return metadata.get("standard_logging_guardrail_information") or []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skipped_scan_records_not_run_entry(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="skip-system-guardrail")
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = self._system_only_data()
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.last_inputs is None, "nothing survived scoping, apply_guardrail must not run"
|
||||
entries = self._recorded_entries(data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_name"] == "skip-system-guardrail"
|
||||
assert entries[0]["guardrail_status"] == "not_run"
|
||||
assert entries[0]["guardrail_response"] == "no scannable content after message scoping"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("skip_system", [False, True])
|
||||
async def test_empty_content_does_not_blame_scoping(self, skip_system: bool):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="unscoped-guardrail")
|
||||
guardrail.skip_system_message_in_guardrail = skip_system
|
||||
data = {"messages": [{"role": "user", "content": None}]}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.last_inputs is None
|
||||
entries = self._recorded_entries(data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_status"] == "not_run"
|
||||
assert entries[0]["guardrail_response"] == "no scannable content"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_self_recording_guardrail_is_left_alone(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="self-recording-guardrail")
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
guardrail.records_own_guardrail_information = True
|
||||
data = self._system_only_data()
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.last_inputs is None
|
||||
assert self._recorded_entries(data) == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scannable_content_records_no_extra_entry(self):
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="normal-guardrail")
|
||||
data = {"messages": [{"role": "user", "content": "hello"}]}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.last_inputs is not None
|
||||
assert all(e.get("guardrail_status") != "not_run" for e in self._recorded_entries(data))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_image_only_content_is_not_reported_as_not_run(self):
|
||||
"""Images are only scanned alongside text, so an image-only request is a
|
||||
pre-existing scan gap, not a message-scoping skip, and must not be labelled one"""
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="image-guardrail")
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"messages": [
|
||||
{"role": "system", "content": "SYSTEM-PROMPT"},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [{"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}}],
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert self._recorded_entries(data) == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scoped_out_image_only_message_is_not_reported_as_not_run(self):
|
||||
"""An image in a skipped role must behave like any other image-only request"""
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="image-guardrail")
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [{"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}}],
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.last_inputs is None
|
||||
assert self._recorded_entries(data) == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scoped_out_text_with_image_records_not_run(self):
|
||||
"""Scoping removed text too, so the skip is recorded even though an image sat beside it"""
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
guardrail = MockGuardrail(guardrail_name="image-guardrail")
|
||||
guardrail.skip_system_message_in_guardrail = True
|
||||
data = {
|
||||
"messages": [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": "Describe this picture."},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}},
|
||||
],
|
||||
},
|
||||
]
|
||||
}
|
||||
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=guardrail)
|
||||
|
||||
assert guardrail.last_inputs is None
|
||||
entries = self._recorded_entries(data)
|
||||
assert len(entries) == 1
|
||||
assert entries[0]["guardrail_status"] == "not_run"
|
||||
assert entries[0]["guardrail_response"] == "no scannable content after message scoping"
|
||||
|
||||
|
||||
class ToolDroppingTextGuardrail(CustomGuardrail):
|
||||
"""Answers one text per non-tool message it saw, the way a guardrail that
|
||||
filters tool rows out before scanning does, and hands back only texts."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__(guardrail_name="tool-dropping-redactor")
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
kept = [m for m in inputs.get("structured_messages") or [] if m.get("role") != "tool"]
|
||||
return {**inputs, "texts": [str(m.get("content")).replace("POISON", "[BLOCKED]") for m in kept]}
|
||||
|
||||
|
||||
class TestPerMessageTextWriteBack:
|
||||
"""Texts that no longer pair one-to-one with what the handler extracted must be
|
||||
rejected by name instead of sliding onto the wrong messages."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fewer_texts_than_extracted_over_a_tool_message_is_rejected(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
handler = OpenAIChatCompletionsHandler()
|
||||
original_messages = [
|
||||
{"role": "system", "content": "SYSTEM-PROMPT"},
|
||||
{"role": "user", "content": "fetch the page"},
|
||||
{"role": "assistant", "content": "fetching"},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "page says POISON here"},
|
||||
{"role": "user", "content": "and then?"},
|
||||
]
|
||||
data = {"messages": json.loads(json.dumps(original_messages))}
|
||||
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await handler.process_input_messages(data=data, guardrail_to_apply=ToolDroppingTextGuardrail())
|
||||
|
||||
assert excinfo.value.guardrail_name == "tool-dropping-redactor"
|
||||
assert data["messages"] == original_messages, "a rejected rewrite must leave the request untouched"
|
||||
|
||||
|
||||
class TestBuildBlockSseChunks:
|
||||
"""build_block_sse_chunks turns a streaming ModifyResponseException into 200 SSE chunks"""
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ with guardrail transformations.
|
|||
import copy
|
||||
from collections.abc import Callable
|
||||
from typing import Any, List, Literal, Optional, Tuple
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import logging
|
||||
|
||||
|
|
@ -31,6 +31,7 @@ from litellm.llms.openai.responses.guardrail_translation.handler import (
|
|||
OpenAIResponsesHandler,
|
||||
)
|
||||
from litellm.llms.openai.responses.guardrail_translation.tool_merge import merge_guardrailed_tools
|
||||
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
|
||||
from litellm.types.llms.openai import ChatCompletionToolCallChunk
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
|
|
@ -2338,6 +2339,135 @@ def _parallel_tool_call_input() -> list:
|
|||
]
|
||||
|
||||
|
||||
SSN = "123-45-6789"
|
||||
REDACTED_SSN = "<US_SSN>"
|
||||
|
||||
|
||||
def _redacted(value: object) -> object:
|
||||
if isinstance(value, str):
|
||||
return value.replace(SSN, REDACTED_SSN)
|
||||
if isinstance(value, list):
|
||||
return [{**part, "text": _redacted(part["text"])} if "text" in part else part for part in value]
|
||||
return value
|
||||
|
||||
|
||||
def _per_message_guardrail_server(structured_messages_in_answer: bool) -> Callable[..., MagicMock]:
|
||||
"""Answers one redacted text per chat row it was shown, the way a guardrail
|
||||
that scans per message does, and optionally the rewritten rows themselves."""
|
||||
|
||||
def post(url: str, json: dict, headers: dict) -> MagicMock:
|
||||
rows = json["structured_messages"]
|
||||
answer: dict = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": [_redacted(row["content"]) if isinstance(row.get("content"), str) else "" for row in rows],
|
||||
}
|
||||
if structured_messages_in_answer:
|
||||
answer["structured_messages"] = [{**row, "content": _redacted(row.get("content"))} for row in rows]
|
||||
response = MagicMock()
|
||||
response.json.return_value = answer
|
||||
response.raise_for_status = MagicMock()
|
||||
return response
|
||||
|
||||
return post
|
||||
|
||||
|
||||
def _per_message_redactor() -> GenericGuardrailAPI:
|
||||
return GenericGuardrailAPI(
|
||||
api_base="https://guardrail.test",
|
||||
guardrail_name="per-message-redactor",
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
|
||||
def _tool_replay_request() -> dict:
|
||||
return {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Never repeat the SSN " + SSN + " back.",
|
||||
"input": [
|
||||
{"role": "user", "content": "Look up " + SSN + " for me."},
|
||||
{"type": "function_call", "call_id": "call_1", "name": "lookup_customer", "arguments": '{"id": "42"}'},
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": '{"ssn": "' + SSN + '"}'},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _string_input_request() -> dict:
|
||||
return {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Never repeat the SSN " + SSN + " back.",
|
||||
"input": "My SSN is " + SSN + ".",
|
||||
}
|
||||
|
||||
|
||||
class TestPerMessageRewriteWriteBack:
|
||||
"""A guardrail that rewrites per chat row hands the rows back as
|
||||
structured_messages, and the handler lands them on the instructions and the
|
||||
input items they came from; the same rewrite handed back as texts alone has
|
||||
no item to land on and is rejected by name instead of sent unrewritten."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_structured_rows_land_on_instructions_and_tool_output(self):
|
||||
guardrail = _per_message_redactor()
|
||||
data = _tool_replay_request()
|
||||
function_call_item = data["input"][1]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(True)):
|
||||
result = await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back."
|
||||
assert _texts(result["input"][0]) == ["Look up " + REDACTED_SSN + " for me."]
|
||||
assert result["input"][1] == function_call_item
|
||||
assert result["input"][2] == {
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_1",
|
||||
"output": '{"ssn": "' + REDACTED_SSN + '"}',
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_texts_only_per_message_answer_is_rejected_by_name(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
guardrail = _per_message_redactor()
|
||||
data = _tool_replay_request()
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)):
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "per-message-redactor"
|
||||
assert data["input"] == original["input"]
|
||||
assert data["instructions"] == original["instructions"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_structured_rows_land_on_instructions_and_string_input(self):
|
||||
guardrail = _per_message_redactor()
|
||||
data = _string_input_request()
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(True)):
|
||||
result = await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back."
|
||||
assert [_texts(item) for item in result["input"]] == [["My SSN is " + REDACTED_SSN + "."]]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_texts_only_per_message_answer_over_a_string_input_is_rejected_by_name(self):
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
guardrail = _per_message_redactor()
|
||||
data = _string_input_request()
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)):
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "per-message-redactor"
|
||||
assert data["input"] == original["input"]
|
||||
assert data["instructions"] == original["instructions"]
|
||||
|
||||
|
||||
class TestProvenancePatching:
|
||||
"""The O(n) provenance pass must keep patching rewritten rows in place for the
|
||||
shapes real agent loops produce, and fall back safely everywhere else."""
|
||||
|
|
|
|||
|
|
@ -133,7 +133,20 @@ def test_namespace_keeps_a_non_function_member_when_a_function_member_is_edited(
|
|||
assert merged[0]["tools"][1] == custom_member
|
||||
|
||||
|
||||
def test_namespace_keeps_its_non_function_members_when_every_function_member_is_dropped():
|
||||
def test_namespace_keeps_its_custom_member_when_every_function_member_is_dropped():
|
||||
custom_member = {"type": "custom", "name": "grep", "description": "Grep", "format": {"type": "text"}}
|
||||
original = [
|
||||
{"type": "namespace", "name": "ns", "description": "NS", "tools": [_function("read"), custom_member]},
|
||||
_function("a"),
|
||||
]
|
||||
groups = _groups(original)
|
||||
|
||||
merged = merge_guardrailed_tools(original, groups, [groups[0][1], groups[1][0]])
|
||||
|
||||
assert list(merged) == [{"type": "namespace", "name": "ns", "description": "NS", "tools": [custom_member]}, _function("a")]
|
||||
|
||||
|
||||
def test_namespace_custom_member_is_dropped_when_the_guardrail_drops_its_chat_form():
|
||||
custom_member = {"type": "custom", "name": "grep", "description": "Grep", "format": {"type": "text"}}
|
||||
original = [
|
||||
{"type": "namespace", "name": "ns", "description": "NS", "tools": [_function("read"), custom_member]},
|
||||
|
|
@ -143,7 +156,37 @@ def test_namespace_keeps_its_non_function_members_when_every_function_member_is_
|
|||
|
||||
merged = merge_guardrailed_tools(original, groups, [groups[1][0]])
|
||||
|
||||
assert list(merged) == [{"type": "namespace", "name": "ns", "description": "NS", "tools": [custom_member]}, _function("a")]
|
||||
assert list(merged) == [_function("a")]
|
||||
|
||||
|
||||
def test_custom_member_description_edit_lands_without_the_namespace_prefix_or_grammar_block():
|
||||
grammar = {"type": "grammar", "syntax": "lark", "definition": "start: X"}
|
||||
custom_member = {"type": "custom", "name": "exec", "description": "Run a command", "format": grammar}
|
||||
original = [{"type": "namespace", "name": "shell", "description": "Shell", "tools": [custom_member]}]
|
||||
groups = _groups(original)
|
||||
assert groups[0][0]["function"]["description"] == "Shell\n\nRun a command\n\nFormat:\n```lark\nstart: X\n```"
|
||||
edited = copy.deepcopy(_flat(groups))
|
||||
edited[0]["function"]["description"] = "Shell\n\nRun a command (guarded)\n\nFormat:\n```lark\nstart: X\n```"
|
||||
|
||||
merged = merge_guardrailed_tools(original, groups, edited)
|
||||
|
||||
guarded_member = {**custom_member, "description": "Run a command (guarded)"}
|
||||
assert list(merged) == [{"type": "namespace", "name": "shell", "description": "Shell", "tools": [guarded_member]}]
|
||||
|
||||
|
||||
def test_text_appended_after_the_grammar_block_lands_on_the_member_without_the_block():
|
||||
grammar = {"type": "grammar", "syntax": "lark", "definition": "start: X"}
|
||||
custom_member = {"type": "custom", "name": "exec", "description": "Run a command", "format": grammar}
|
||||
original = [{"type": "namespace", "name": "shell", "description": "Shell", "tools": [custom_member]}]
|
||||
groups = _groups(original)
|
||||
edited = copy.deepcopy(_flat(groups))
|
||||
edited[0]["function"]["description"] = edited[0]["function"]["description"] + " [checked]"
|
||||
|
||||
merged = merge_guardrailed_tools(original, groups, edited)
|
||||
|
||||
assert merged[0]["tools"][0]["description"] == "Run a command [checked]"
|
||||
reflattened = _flat(_groups(merged))
|
||||
assert reflattened[0]["function"]["description"] == "Shell\n\nRun a command [checked]\n\nFormat:\n```lark\nstart: X\n```"
|
||||
|
||||
|
||||
def test_member_extras_edited_by_the_guardrail_land_on_that_member():
|
||||
|
|
|
|||
|
|
@ -252,11 +252,52 @@ class TestRender:
|
|||
def test_savings_header_and_bars_against_the_routers_baseline(self, config_dir):
|
||||
text = render("claude-sonnet-5", RECORDED, config_dir, use_color=False, bar_width=10)
|
||||
assert text.splitlines() == [
|
||||
"claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5",
|
||||
"LiteLLM ████░░░░░░ $0.14",
|
||||
"Routed to: claude-sonnet-5 -63% vs Claude Opus 5",
|
||||
"claude-auto ████░░░░░░ $0.14",
|
||||
"Claude Opus 5 ██████████ $0.38",
|
||||
]
|
||||
|
||||
def test_a_long_router_name_keeps_both_cost_bars_aligned(self, config_dir: Path) -> None:
|
||||
session: Final = RECORDED._replace(router_name="engineering-smart-router")
|
||||
text: Final = render("claude-sonnet-5", session, config_dir, use_color=False, bar_width=10)
|
||||
assert text.splitlines()[1:] == [
|
||||
"engineering-smart-router ████░░░░░░ $0.14",
|
||||
"Claude Opus 5 ██████████ $0.38",
|
||||
]
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("router_name", "baseline_name", "router_padding", "baseline_padding"),
|
||||
(
|
||||
("路由-router", "Claude Opus 5", 3, 1),
|
||||
("智能模型路由器", "Claude Opus 5", 1, 2),
|
||||
("ABC-router", "Claude Opus 5", 1, 1),
|
||||
("cafe\u0301-router", "Claude Opus 5", 3, 1),
|
||||
("a\u20dd-router", "Claude Opus 5", 6, 1),
|
||||
("カ\u3099-router", "Claude Opus 5", 5, 1),
|
||||
("auto", "基準モデル", 7, 1),
|
||||
("auto", "cafe\u0301", 1, 1),
|
||||
),
|
||||
)
|
||||
@pytest.mark.parametrize("use_color", (False, True))
|
||||
def test_unicode_labels_align_cost_bars_by_terminal_columns(
|
||||
self,
|
||||
config_dir: Path,
|
||||
router_name: str,
|
||||
baseline_name: str,
|
||||
router_padding: int,
|
||||
baseline_padding: int,
|
||||
use_color: bool,
|
||||
) -> None:
|
||||
(config_dir / "cache" / "gateway-models.json").write_text(
|
||||
json.dumps({"models": [{"id": "claude-opus-5", "display_name": baseline_name}]})
|
||||
)
|
||||
session: Final = RECORDED._replace(router_name=router_name)
|
||||
text: Final = ANSI.sub("", render("claude-sonnet-5", session, config_dir, use_color, bar_width=10))
|
||||
assert text.splitlines()[1:] == [
|
||||
f"{router_name}{' ' * router_padding}████░░░░░░ $0.14",
|
||||
f"{baseline_name}{' ' * baseline_padding}██████████ $0.38",
|
||||
]
|
||||
|
||||
def test_control_characters_in_any_externally_sourced_label_never_reach_the_terminal(self, tmp_path, config_dir):
|
||||
# The transcript, the proxy payload and Claude Code's model cache all feed labels straight into a
|
||||
# terminal, and none is under this script's control. Only the control bytes are dropped (ESC, BEL,
|
||||
|
|
@ -289,7 +330,7 @@ class TestRender:
|
|||
assert "+25% vs Claude Opus 5" in render("m", dearer, config_dir, use_color=False)
|
||||
|
||||
def test_without_a_baseline_only_the_routed_line_shows(self, config_dir):
|
||||
assert render("m", RECORDED._replace(baseline_model=None), config_dir, False) == "claude-auto · Routed to: m"
|
||||
assert render("m", RECORDED._replace(baseline_model=None), config_dir, False) == "Routed to: m"
|
||||
assert render("m", None, config_dir, False) == "Routed to: m"
|
||||
|
||||
def test_color_wraps_the_same_text(self, config_dir):
|
||||
|
|
@ -311,7 +352,8 @@ class TestClaudeCodeMode:
|
|||
return Fetched(RECORDED, definitive=True)
|
||||
|
||||
text: Final = _run(_payload(transcript), _env(tmp_path, config_dir), fetch)
|
||||
assert text.startswith("claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5\n")
|
||||
assert text.startswith("Routed to: claude-sonnet-5 -63% vs Claude Opus 5\n")
|
||||
assert text.splitlines()[1].startswith("claude-auto ")
|
||||
|
||||
def test_a_discovered_display_name_labels_the_sessions_model(
|
||||
self, tmp_path: Path, transcript: Path, config_dir: Path
|
||||
|
|
@ -322,7 +364,7 @@ class TestClaudeCodeMode:
|
|||
return Fetched(session, definitive=True)
|
||||
|
||||
text: Final = _run(_payload(transcript), _env(tmp_path, config_dir), fetch)
|
||||
assert text.startswith("claude-auto · Routed to: Claude Opus 5 -63% vs Claude Opus 5\n")
|
||||
assert text.startswith("Routed to: Claude Opus 5 -63% vs Claude Opus 5\n")
|
||||
|
||||
def test_an_unrecorded_session_degrades_to_the_routed_line(self, tmp_path, transcript, config_dir):
|
||||
assert _run(_payload(transcript), _env(tmp_path, config_dir), lambda c, s: Fetched(None, True)) == (
|
||||
|
|
@ -378,7 +420,8 @@ class TestCodexMode:
|
|||
|
||||
out = _run({"hook_event_name": "Stop", "session_id": SESSION_ID, "transcript_path": "/nope"}, env, fetch)
|
||||
message = json.loads(out)["systemMessage"]
|
||||
assert message.splitlines()[1] == "claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5"
|
||||
assert message.splitlines()[1] == "Routed to: claude-sonnet-5 -63% vs Claude Opus 5"
|
||||
assert message.splitlines()[2].startswith("claude-auto ")
|
||||
assert message.startswith("\n")
|
||||
assert seen == [Credentials("http://127.0.0.1:4000", "sk-codex")]
|
||||
|
||||
|
|
|
|||
|
|
@ -143,3 +143,18 @@ def test_a_status_carried_by_an_exception_drives_the_type_it_reports():
|
|||
exc = HTTPException(status_code=403, detail="blocked by policy")
|
||||
|
||||
assert openai_error_type(exc, error_status_code(exc, 400)) == "permission_error"
|
||||
|
||||
|
||||
def test_a_stringified_none_type_or_param_is_treated_as_absent():
|
||||
from litellm.exceptions import BadRequestError
|
||||
|
||||
carried = BadRequestError(
|
||||
message="Content blocked",
|
||||
model="claude-haiku-4-5",
|
||||
llm_provider="litellm_proxy",
|
||||
body={"message": "Content blocked", "type": "None", "param": "None", "code": "400"},
|
||||
)
|
||||
|
||||
assert carried.type == "None"
|
||||
assert openai_error_type(carried, 400) == "invalid_request_error"
|
||||
assert openai_error_param(carried) is None
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Tests for the credential management endpoints."""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -9,6 +10,7 @@ from fastapi.testclient import TestClient
|
|||
import litellm
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.credential_endpoints.endpoints import get_llm_router
|
||||
from litellm.proxy.proxy_server import app
|
||||
from litellm.types.utils import CredentialItem
|
||||
|
||||
|
|
@ -47,23 +49,27 @@ def _list_credentials():
|
|||
@pytest.fixture
|
||||
def credential_store():
|
||||
"""Stands the credential store up for one test: whether the database is reachable, what
|
||||
the proxy is already serving from memory, and what each repository call hands back."""
|
||||
the proxy is already serving from memory, which router deployments resolve against, and
|
||||
what each repository call hands back."""
|
||||
|
||||
def install(
|
||||
*,
|
||||
connected: bool = True,
|
||||
in_memory: tuple[object, ...] = (),
|
||||
llm_router: object | None = None,
|
||||
**repository_calls: AsyncMock,
|
||||
) -> None:
|
||||
patch("litellm.proxy.proxy_server.prisma_client", MagicMock() if connected else None).start()
|
||||
patch("litellm.proxy.proxy_server.master_key", "sk-test-master").start()
|
||||
patch.object(litellm, "credential_list", list(in_memory)).start()
|
||||
app.dependency_overrides[get_llm_router] = lambda: llm_router
|
||||
repository = patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository").start()
|
||||
for call_name, result in repository_calls.items():
|
||||
setattr(repository.return_value, call_name, result)
|
||||
|
||||
yield install
|
||||
patch.stopall()
|
||||
app.dependency_overrides.pop(get_llm_router, None)
|
||||
|
||||
|
||||
def test_update_credential_answers_404_when_the_credential_does_not_exist(credential_store):
|
||||
|
|
@ -122,7 +128,9 @@ def test_delete_credential_answers_404_when_the_credential_does_not_exist(creden
|
|||
|
||||
response = _delete_credential("definitely-not-there")
|
||||
|
||||
assert response.status_code == 404, f"delete of a missing credential answered {response.status_code}: {response.text}"
|
||||
assert response.status_code == 404, (
|
||||
f"delete of a missing credential answered {response.status_code}: {response.text}"
|
||||
)
|
||||
assert "definitely-not-there" in response.text
|
||||
|
||||
|
||||
|
|
@ -195,3 +203,130 @@ def test_get_credentials_answers_an_error_status_when_the_listing_fails(credenti
|
|||
|
||||
assert response.status_code == 500, f"failed listing answered {response.status_code}: {response.text}"
|
||||
assert response.json().get("success") is not True
|
||||
|
||||
|
||||
def _create_credential(body: dict):
|
||||
return _call_as_admin("POST", "/credentials", body)
|
||||
|
||||
|
||||
class _UniqueViolation(Exception):
|
||||
code = "P2002"
|
||||
|
||||
|
||||
def test_create_credential_answers_409_when_the_name_is_already_taken(credential_store):
|
||||
"""Regression: the unique index used to surface as a Prisma 500 that callers string-matched."""
|
||||
credential_store(
|
||||
create=AsyncMock(side_effect=_UniqueViolation("Unique constraint failed on the fields: (`credential_name`)")),
|
||||
)
|
||||
|
||||
response = _create_credential(
|
||||
{"credential_name": "aws_bedrock", "credential_values": {"aws_access_key_id": "new"}, "credential_info": {}},
|
||||
)
|
||||
|
||||
assert response.status_code == 409, f"name collision answered {response.status_code}: {response.text}"
|
||||
message = response.json()["error"]["message"]
|
||||
assert message == (
|
||||
"Credential 'aws_bedrock' already exists. Update it with PATCH /credentials/aws_bedrock, or delete it first."
|
||||
), f"the operator reads this message verbatim: {message}"
|
||||
assert "Unique constraint" not in response.text, f"the Prisma internals must not leak: {response.text}"
|
||||
|
||||
|
||||
def test_create_credential_still_answers_500_when_the_write_fails_for_another_reason(credential_store):
|
||||
credential_store(create=AsyncMock(side_effect=Exception("connection reset by peer")))
|
||||
|
||||
response = _create_credential(
|
||||
{"credential_name": "aws_bedrock", "credential_values": {"aws_access_key_id": "new"}, "credential_info": {}},
|
||||
)
|
||||
|
||||
assert response.status_code == 500, f"database fault answered {response.status_code}: {response.text}"
|
||||
|
||||
|
||||
def test_create_credential_still_answers_200_for_a_name_that_is_free(credential_store):
|
||||
find_by_name = AsyncMock()
|
||||
credential_store(find_by_name=find_by_name, create=AsyncMock(return_value=None))
|
||||
|
||||
response = _create_credential(
|
||||
{"credential_name": "brand_new", "credential_values": {"aws_access_key_id": "new"}, "credential_info": {}},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["success"] is True
|
||||
find_by_name.assert_not_awaited(), "the unique index is the guard; create must not add a lookup"
|
||||
|
||||
|
||||
def test_update_credential_resolves_credential_values_from_model_id_like_create(credential_store):
|
||||
"""Regression: PATCH dropped ``model_id`` from the body, so an update that named a
|
||||
deployment instead of raw values wrote whatever the caller sent, or nothing."""
|
||||
stored = CredentialItem(
|
||||
credential_name="from-deployment",
|
||||
credential_values={"api_key": "sk-old"},
|
||||
credential_info={},
|
||||
)
|
||||
update_by_name = AsyncMock(return_value=None)
|
||||
router = MagicMock()
|
||||
router.get_deployment.return_value = {"model_name": "gpt-5.2"}
|
||||
router.get_deployment_credentials.return_value = {"api_key": "sk-from-deployment"}
|
||||
credential_store(find_by_name=AsyncMock(return_value=stored), update_by_name=update_by_name, llm_router=router)
|
||||
|
||||
response = _patch_credential(
|
||||
"from-deployment",
|
||||
{"credential_name": "from-deployment", "model_id": "deployment-1", "credential_info": {}},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
router.get_deployment_credentials.assert_called_once_with("deployment-1")
|
||||
written = json.loads(update_by_name.await_args.kwargs["data"]["credential_values"])
|
||||
assert set(written) == {"api_key"}
|
||||
assert written["api_key"] != "sk-old", "the deployment's values must replace the stored ones"
|
||||
assert written["api_key"] != "sk-from-deployment", "values are encrypted before they reach the table"
|
||||
|
||||
|
||||
def test_update_credential_answers_404_when_model_id_names_no_deployment(credential_store):
|
||||
stored = CredentialItem(
|
||||
credential_name="from-deployment", credential_values={"api_key": "sk-old"}, credential_info={}
|
||||
)
|
||||
update_by_name = AsyncMock(return_value=None)
|
||||
router = MagicMock()
|
||||
router.get_deployment.return_value = None
|
||||
credential_store(find_by_name=AsyncMock(return_value=stored), update_by_name=update_by_name, llm_router=router)
|
||||
|
||||
response = _patch_credential(
|
||||
"from-deployment",
|
||||
{"credential_name": "from-deployment", "model_id": "no-such-deployment", "credential_info": {}},
|
||||
)
|
||||
|
||||
assert response.status_code == 404, response.text
|
||||
update_by_name.assert_not_awaited()
|
||||
|
||||
|
||||
def test_update_credential_answers_500_when_model_id_is_given_but_no_router_is_loaded(credential_store):
|
||||
stored = CredentialItem(
|
||||
credential_name="from-deployment", credential_values={"api_key": "sk-old"}, credential_info={}
|
||||
)
|
||||
update_by_name = AsyncMock(return_value=None)
|
||||
credential_store(find_by_name=AsyncMock(return_value=stored), update_by_name=update_by_name, llm_router=None)
|
||||
|
||||
response = _patch_credential(
|
||||
"from-deployment",
|
||||
{"credential_name": "from-deployment", "model_id": "deployment-1", "credential_info": {}},
|
||||
)
|
||||
|
||||
assert response.status_code == 500, response.text
|
||||
update_by_name.assert_not_awaited()
|
||||
|
||||
|
||||
def test_update_credential_still_accepts_a_body_without_credential_values(credential_store):
|
||||
"""Renaming or re-tagging a credential sends only ``credential_info``; that must not 422."""
|
||||
stored = CredentialItem(credential_name="existing", credential_values={"api_key": "sk-old"}, credential_info={})
|
||||
update_by_name = AsyncMock(return_value=None)
|
||||
credential_store(find_by_name=AsyncMock(return_value=stored), update_by_name=update_by_name)
|
||||
|
||||
response = _patch_credential(
|
||||
"existing",
|
||||
{"credential_name": "existing", "credential_info": {"custom_llm_provider": "openai"}},
|
||||
)
|
||||
|
||||
assert response.status_code == 200, response.text
|
||||
written = update_by_name.await_args.kwargs["data"]
|
||||
assert json.loads(written["credential_info"]) == {"custom_llm_provider": "openai"}
|
||||
assert set(json.loads(written["credential_values"])) == {"api_key"}, "stored values survive an info-only patch"
|
||||
|
|
|
|||
|
|
@ -1820,7 +1820,7 @@ async def test_unalignable_rewrite_is_rejected_never_sent_unredacted(
|
|||
Skipping the write-back would hand the model the unredacted text, so a
|
||||
guardrail could be bypassed by adding ``instructions`` or a tool call.
|
||||
"""
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
|
||||
data: dict[str, object] = {"model": "gpt-4o", "input": responses_input}
|
||||
if instructions is not None:
|
||||
|
|
|
|||
|
|
@ -582,6 +582,145 @@ class TestGuardrailActions:
|
|||
assert result_images is None
|
||||
|
||||
|
||||
class TestStructuredMessagesInResponse:
|
||||
"""A guardrail server that rewrites per chat row answers with the rewritten
|
||||
rows as structured_messages, which the endpoint handlers write back by row."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returned_rows_are_handed_back_as_structured_messages(
|
||||
self, generic_guardrail, mock_request_data_input
|
||||
):
|
||||
rewritten_rows = [
|
||||
{"role": "system", "content": "Never repeat an SSN."},
|
||||
{"role": "user", "content": "Look up [REDACTED] for me."},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'},
|
||||
]
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": ["Never repeat an SSN.", "Look up [REDACTED] for me.", '{"ssn": "[REDACTED]"}'],
|
||||
"structured_messages": rewritten_rows,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["Look up 123-45-6789 for me."]},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert guardrailed_inputs["structured_messages"] == rewritten_rows
|
||||
assert guardrailed_inputs["texts"] == mock_response.json.return_value["texts"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rows_echoed_back_as_shown_keep_their_original_keys(
|
||||
self, generic_guardrail, mock_request_data_input
|
||||
):
|
||||
tool_call_row = {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}, "index": 0}
|
||||
],
|
||||
}
|
||||
original_rows = [
|
||||
{"role": "user", "content": "Look up 123-45-6789 for me.", "name": "pat"},
|
||||
tool_call_row,
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "123-45-6789"}'},
|
||||
]
|
||||
|
||||
def echo_with_tool_output_redacted(url, json, headers):
|
||||
shown_rows = json["structured_messages"]
|
||||
assert "index" not in shown_rows[1]["tool_calls"][0]
|
||||
assert "name" not in shown_rows[0]
|
||||
answer = MagicMock()
|
||||
answer.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": ["Look up 123-45-6789 for me."],
|
||||
"structured_messages": [
|
||||
shown_rows[0],
|
||||
shown_rows[1],
|
||||
{**shown_rows[2], "content": '{"ssn": "[REDACTED]"}'},
|
||||
],
|
||||
}
|
||||
answer.raise_for_status = MagicMock()
|
||||
return answer
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", side_effect=echo_with_tool_output_redacted):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["Look up 123-45-6789 for me."], "structured_messages": original_rows},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
returned_rows = guardrailed_inputs["structured_messages"]
|
||||
assert returned_rows[0] is original_rows[0]
|
||||
assert returned_rows[1] is tool_call_row
|
||||
assert returned_rows[2] == {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rows_all_echoed_back_as_shown_leave_the_rewrite_to_texts(
|
||||
self, generic_guardrail, mock_request_data_input
|
||||
):
|
||||
"""A server written against the texts contract that echoes the request rows back
|
||||
untouched while rewriting texts still gets its texts rewrite applied."""
|
||||
original_rows = [
|
||||
{"role": "system", "content": "Never repeat an SSN."},
|
||||
{"role": "user", "content": "Look up 123-45-6789 for me."},
|
||||
]
|
||||
|
||||
def echo_rows_and_rewrite_texts(url, json, headers):
|
||||
answer = MagicMock()
|
||||
answer.json.return_value = {
|
||||
"action": "NONE",
|
||||
"texts": [text.replace("123-45-6789", "[REDACTED]") for text in json["texts"]],
|
||||
"structured_messages": json["structured_messages"],
|
||||
}
|
||||
answer.raise_for_status = MagicMock()
|
||||
return answer
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", side_effect=echo_rows_and_rewrite_texts):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={
|
||||
"texts": ["Never repeat an SSN.", "Look up 123-45-6789 for me."],
|
||||
"structured_messages": original_rows,
|
||||
},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert "structured_messages" not in guardrailed_inputs
|
||||
assert guardrailed_inputs["texts"] == ["Never repeat an SSN.", "Look up [REDACTED] for me."]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"structured_messages",
|
||||
[[], [{"content": "a row with no role"}], "not a list"],
|
||||
ids=["empty", "no_role", "not_a_list"],
|
||||
)
|
||||
async def test_rows_that_are_not_chat_messages_are_ignored(
|
||||
self, generic_guardrail, mock_request_data_input, structured_messages
|
||||
):
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": ["[REDACTED]"],
|
||||
"structured_messages": structured_messages,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["Look up 123-45-6789 for me."]},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert "structured_messages" not in guardrailed_inputs
|
||||
assert guardrailed_inputs["texts"] == ["[REDACTED]"]
|
||||
|
||||
|
||||
class TestImageSupport:
|
||||
"""Test image handling in guardrail requests"""
|
||||
|
||||
|
|
|
|||
|
|
@ -4620,46 +4620,27 @@ class TestPanwAirsLatestRoleMessageOnly:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_system_plus_multiturn_no_fallback(self):
|
||||
"""Anthropic with top-level system + multi-turn messages[]
|
||||
— latest-user works, no scan-all fallback.
|
||||
"""Anthropic with a top-level system prompt and multi-turn messages[]
|
||||
scans only the latest user turn, with no scan-all fallback.
|
||||
|
||||
Key scenario: Anthropic top-level `system` field causes
|
||||
structured_messages to have an injected system entry, but
|
||||
request_data["messages"] does NOT include it.
|
||||
The Anthropic handler hoists the top-level `system` field into both
|
||||
`texts` and `structured_messages`, so the latest-user walk has to
|
||||
count the same entries the framework flattened.
|
||||
"""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
from litellm.llms.anthropic.chat.guardrail_translation.handler import (
|
||||
AnthropicMessagesHandler,
|
||||
)
|
||||
|
||||
# Original Anthropic messages (no system in messages array)
|
||||
original_messages = [
|
||||
{"role": "user", "content": "First user turn"},
|
||||
{"role": "assistant", "content": "First assistant turn"},
|
||||
{"role": "user", "content": "Latest user turn"},
|
||||
]
|
||||
|
||||
# texts extracted from original_messages (3 text entries)
|
||||
texts = ["First user turn", "First assistant turn", "Latest user turn"]
|
||||
|
||||
# structured_messages has an INJECTED system message from translation
|
||||
structured_messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "First user turn"},
|
||||
{"role": "assistant", "content": "First assistant turn"},
|
||||
{"role": "user", "content": "Latest user turn"},
|
||||
]
|
||||
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"texts": texts,
|
||||
"structured_messages": structured_messages,
|
||||
}
|
||||
handler = make_handler()
|
||||
request_data = {
|
||||
"litellm_call_id": "test-call-id",
|
||||
"model": "anthropic/claude-sonnet-4-20250514",
|
||||
"messages": original_messages,
|
||||
"system": "You are a helpful assistant.",
|
||||
"messages": [
|
||||
{"role": "user", "content": "First user turn"},
|
||||
{"role": "assistant", "content": "First assistant turn"},
|
||||
{"role": "user", "content": "Latest user turn"},
|
||||
],
|
||||
"proxy_server_request": {
|
||||
"url": "http://localhost:4000/v1/messages",
|
||||
},
|
||||
|
|
@ -4670,13 +4651,11 @@ class TestPanwAirsLatestRoleMessageOnly:
|
|||
) as mock_api:
|
||||
mock_api.return_value = {"action": "allow", "category": "benign"}
|
||||
|
||||
await handler.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
await AnthropicMessagesHandler().process_input_messages(
|
||||
data=request_data,
|
||||
guardrail_to_apply=handler,
|
||||
)
|
||||
|
||||
# Should scan ONLY the latest user message, not fall back to scan-all
|
||||
assert mock_api.call_count == 1
|
||||
assert mock_api.call_args.kwargs["content"] == "Latest user turn"
|
||||
|
||||
|
|
|
|||
|
|
@ -16,14 +16,21 @@ Streaming: CSW.__anext__ stores args on logging_obj at stream end.
|
|||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Any, Final
|
||||
from collections.abc import Callable, Mapping
|
||||
from datetime import datetime
|
||||
from typing import Any, Final, cast
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.utils import StandardLoggingPayload
|
||||
from litellm.utils import _dispatch_success_logging
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
@ -54,6 +61,27 @@ def _attach_mock_success_dispatch(mock_logging_obj, async_success_fn):
|
|||
mock_logging_obj.async_success_handler = async_success_fn
|
||||
|
||||
|
||||
async def _wait_until(condition: Callable[[], bool]) -> None:
|
||||
"""Give the logging worker a bounded window to run what the closure enqueued."""
|
||||
for _ in range(200):
|
||||
if condition():
|
||||
return
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
"""Keeps what the async success callback was handed, the way a spend logger sees it."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.standard_logging_object: StandardLoggingPayload | None = None
|
||||
|
||||
async def async_log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
self.standard_logging_object = cast(StandardLoggingPayload, kwargs["standard_logging_object"])
|
||||
|
||||
|
||||
class PostCallGuardrail(CustomGuardrail):
|
||||
"""A post-call guardrail."""
|
||||
|
||||
|
|
@ -259,6 +287,120 @@ async def test_deferred_flag_stores_and_executes_closure():
|
|||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deferred_slot_keeps_the_innermost_wrapper_result():
|
||||
"""Nested @client wrappers exit through _dispatch_success_logging with one shared logging
|
||||
object. The deferred slot must keep the first stored result, the way the immediate path's
|
||||
has_logged dedupe keeps the first fired task, so the spend log reads usage from the
|
||||
innermost provider-shaped response and never from an outer wrapper's translation of it."""
|
||||
logging_obj: Final = MagicMock()
|
||||
logging_obj._defer_async_logging = True
|
||||
logging_obj._enqueue_deferred_logging = None
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
inner_result: Final = object()
|
||||
outer_result: Final = object()
|
||||
|
||||
for result in (inner_result, outer_result):
|
||||
_dispatch_success_logging(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
is_completion_with_fallbacks=False,
|
||||
is_litellm_internal_call=False,
|
||||
)
|
||||
|
||||
logging_obj._enqueue_deferred_logging()
|
||||
await _wait_until(lambda: logging_obj.async_success_handler.await_count > 0)
|
||||
|
||||
logging_obj.async_success_handler.assert_awaited_once()
|
||||
assert logging_obj.async_success_handler.await_args.kwargs["result"] is inner_result
|
||||
assert logging_obj.handle_sync_success_callbacks_for_async_calls.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deferred_anthropic_messages_bridged_to_the_responses_api_logs_the_provider_usage(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""/v1/messages on an Azure gpt-5.4+ deployment with function tools runs three nested
|
||||
wrappers: anthropic_messages, the chat adapter's acompletion, and the Responses bridge
|
||||
acompletion hands the call to, which retags the call as ``responses``. With logging
|
||||
deferred for a post-call guardrail the stored closure must carry the innermost provider
|
||||
response: logging the Anthropic-shaped reply under Responses semantics books this
|
||||
7,336-token prompt as 3 tokens, since Anthropic's input_tokens excludes the cache hit."""
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True")
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
respx_mock.post(url__regex=r"https://deferred-nested\.openai\.azure\.com/openai/.*responses.*").mock(
|
||||
return_value=httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"id": "resp_deferred_nested",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-5.4-nano",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_deferred_nested",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "Hello!", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"input_tokens": 7336,
|
||||
"input_tokens_details": {"cached_tokens": 7333},
|
||||
"output_tokens": 23,
|
||||
"output_tokens_details": {"reasoning_tokens": 0},
|
||||
"total_tokens": 7359,
|
||||
},
|
||||
},
|
||||
)
|
||||
)
|
||||
recorder: Final = _RecordingLogger()
|
||||
logging_obj: Final = Logging(
|
||||
model="azure/gpt-5.4-nano",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
stream=False,
|
||||
call_type="anthropic_messages",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="deferred-nested-anthropic-messages",
|
||||
function_id="deferred-nested-anthropic-messages",
|
||||
dynamic_async_success_callbacks=[recorder],
|
||||
)
|
||||
logging_obj._defer_async_logging = True
|
||||
|
||||
response: Final = await litellm.anthropic_messages(
|
||||
model="azure/gpt-5.4-nano",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
max_tokens=16,
|
||||
tools=[
|
||||
{
|
||||
"name": "lookup_volume",
|
||||
"description": "Look up a storage volume by name",
|
||||
"input_schema": {"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]},
|
||||
}
|
||||
],
|
||||
api_key="sk-deferred-nested",
|
||||
api_base="https://deferred-nested.openai.azure.com",
|
||||
api_version="2025-04-01-preview",
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
assert response["content"] == [{"type": "text", "text": "Hello!"}]
|
||||
assert response["usage"]["input_tokens"] == 3
|
||||
assert response["usage"]["cache_read_input_tokens"] == 7333
|
||||
|
||||
logging_obj._enqueue_deferred_logging()
|
||||
await _wait_until(lambda: recorder.standard_logging_object is not None)
|
||||
|
||||
assert recorder.standard_logging_object is not None
|
||||
assert recorder.standard_logging_object["prompt_tokens"] == 7336
|
||||
assert recorder.standard_logging_object["metadata"]["usage_object"]["prompt_tokens_details"]["cached_tokens"] == 7333
|
||||
assert recorder.standard_logging_object["response_cost"] == pytest.approx(3 * 2e-7 + 7333 * 2e-8 + 23 * 1.25e-6)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. Non-streaming regression: without flag, create_task fires normally
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import base64
|
||||
from collections.abc import Mapping, Sequence
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -12,6 +13,7 @@ from litellm.proxy.guardrails.guardrail_hooks.prompt_security.prompt_security im
|
|||
PromptSecurityGuardrailMissingSecrets,
|
||||
)
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch):
|
||||
|
|
@ -174,6 +176,123 @@ async def test_apply_guardrail_modify_request(monkeypatch: pytest.MonkeyPatch):
|
|||
assert result["texts"] == ["User prompt with PII: SSN [REDACTED]"]
|
||||
|
||||
|
||||
def _modify_response(modified_messages: Sequence[Mapping[str, object]]) -> Response:
|
||||
mock_response = Response(
|
||||
json={"result": {"prompt": {"action": "modify", "modified_messages": modified_messages}}},
|
||||
status_code=200,
|
||||
request=Request(method="POST", url="https://test.prompt.security/api/protect"),
|
||||
)
|
||||
mock_response.raise_for_status = lambda: None
|
||||
return mock_response
|
||||
|
||||
|
||||
def _tool_replay_messages() -> list[AllMessageValues]:
|
||||
return [
|
||||
{"role": "system", "content": "Never echo an SSN like 123-45-6789."},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Look up 123-45-6789"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/id-card.png"}},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "123-45-6789"}'},
|
||||
{"role": "user", "content": "Summarize what you found."},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_modify_returns_structured_messages_with_tool_rows_kept(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A per-message modify verdict comes back as structured_messages so the
|
||||
endpoint handler can write it back by message, with the rows Prompt Security
|
||||
never saw (tool results) and the non-text parts (images) left in place."""
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True)
|
||||
messages = _tool_replay_messages()
|
||||
inputs = {"texts": ["Look up 123-45-6789", "Summarize what you found."], "structured_messages": messages}
|
||||
modified_messages = [
|
||||
{"role": "system", "content": "Never echo an SSN like [REDACTED]."},
|
||||
{"role": "user", "content": [{"type": "text", "text": "Look up [REDACTED]"}]},
|
||||
{"role": "assistant", "content": None},
|
||||
{"role": "user", "content": "Summarize what you found."},
|
||||
]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={"messages": messages}, input_type="request"
|
||||
)
|
||||
|
||||
assert result["structured_messages"] == [
|
||||
{"role": "system", "content": "Never echo an SSN like [REDACTED]."},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "Look up [REDACTED]"},
|
||||
{"type": "image_url", "image_url": {"url": "https://example.com/id-card.png"}},
|
||||
],
|
||||
},
|
||||
messages[2],
|
||||
messages[3],
|
||||
{"role": "user", "content": "Summarize what you found."},
|
||||
]
|
||||
assert result["structured_messages"] is not messages
|
||||
assert result["texts"] == [
|
||||
"Never echo an SSN like [REDACTED].",
|
||||
"Look up [REDACTED]",
|
||||
"Summarize what you found.",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_modify_with_unexpected_message_count_keeps_texts_only(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True)
|
||||
messages = _tool_replay_messages()
|
||||
inputs = {"texts": ["Look up 123-45-6789", "Summarize what you found."], "structured_messages": messages}
|
||||
modified_messages = [{"role": "user", "content": "Look up [REDACTED]"}]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={"messages": messages}, input_type="request"
|
||||
)
|
||||
|
||||
assert result["structured_messages"] is messages
|
||||
assert result["texts"] == ["Look up [REDACTED]"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_modify_keeps_empty_text_parts_as_slots(monkeypatch: pytest.MonkeyPatch):
|
||||
"""The chat handler counts an empty text part as a slot, so a modify verdict
|
||||
that echoes the empty part still lines up with the row and its texts."""
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
guardrail = PromptSecurityGuardrail(guardrail_name="test-guard", event_hook="pre_call", default_on=True)
|
||||
messages: list[AllMessageValues] = [
|
||||
{"role": "user", "content": [{"type": "text", "text": "Look up 123-45-6789"}, {"type": "text", "text": ""}]}
|
||||
]
|
||||
inputs = {"texts": ["Look up 123-45-6789", ""], "structured_messages": messages}
|
||||
modified_messages = [
|
||||
{"role": "user", "content": [{"type": "text", "text": "Look up [REDACTED]"}, {"type": "text", "text": ""}]}
|
||||
]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)):
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={"messages": messages}, input_type="request"
|
||||
)
|
||||
|
||||
assert result["structured_messages"] == modified_messages
|
||||
assert result["texts"] == ["Look up [REDACTED]", ""]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_allow_request(monkeypatch: pytest.MonkeyPatch):
|
||||
"""Test that apply_guardrail allows safe prompts"""
|
||||
|
|
|
|||
|
|
@ -682,3 +682,67 @@ async def test_detail_prev_trend_query_is_bounded():
|
|||
prev_wheres = [w for w in wheres if "lt" in w.get("date", {})]
|
||||
assert prev_wheres
|
||||
assert all("gte" in w["date"] for w in prev_wheres)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logs_report_not_run_entries_as_not_run_not_passed():
|
||||
"""LIT-6314: a guardrail that never scanned must not be reported as a pass in the drill-down."""
|
||||
index_row = MagicMock()
|
||||
index_row.request_id = "req-nr"
|
||||
index_row.guardrail_id = "db-1"
|
||||
index_row.start_time = datetime(2026, 4, 22)
|
||||
spend_log = MagicMock()
|
||||
spend_log.request_id = "req-nr"
|
||||
spend_log.model = "gpt-4o-mini"
|
||||
spend_log.startTime = datetime(2026, 4, 22)
|
||||
spend_log.metadata = {
|
||||
"guardrail_information": [
|
||||
{"guardrail_name": "db-1", "guardrail_status": "not_run", "duration": 0.0},
|
||||
]
|
||||
}
|
||||
prisma = _prisma(find_unique=_db_row(), index_find_many=[index_row])
|
||||
prisma.db.litellm_spendlogs.find_many = AsyncMock(return_value=[spend_log])
|
||||
handler = _config_handler()
|
||||
p1, p2 = _patches(prisma, handler)
|
||||
with p1, p2:
|
||||
resp = await guardrails_usage_logs(
|
||||
guardrail_id="db-1",
|
||||
policy_id=None,
|
||||
page=1,
|
||||
page_size=50,
|
||||
action=None,
|
||||
start_date=START,
|
||||
end_date=END,
|
||||
user_api_key_dict=ADMIN,
|
||||
)
|
||||
assert [log.action for log in resp.logs] == ["not_run"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_logs_action_passed_filter_excludes_not_run_entries():
|
||||
"""LIT-6314: filtering the drill-down for passes must not return unscanned requests."""
|
||||
index_row = MagicMock()
|
||||
index_row.request_id = "req-nr"
|
||||
index_row.guardrail_id = "db-1"
|
||||
index_row.start_time = datetime(2026, 4, 22)
|
||||
spend_log = MagicMock()
|
||||
spend_log.request_id = "req-nr"
|
||||
spend_log.model = "gpt-4o-mini"
|
||||
spend_log.startTime = datetime(2026, 4, 22)
|
||||
spend_log.metadata = {"guardrail_information": [{"guardrail_name": "db-1", "guardrail_status": "not_run"}]}
|
||||
prisma = _prisma(find_unique=_db_row(), index_find_many=[index_row])
|
||||
prisma.db.litellm_spendlogs.find_many = AsyncMock(return_value=[spend_log])
|
||||
handler = _config_handler()
|
||||
p1, p2 = _patches(prisma, handler)
|
||||
with p1, p2:
|
||||
resp = await guardrails_usage_logs(
|
||||
guardrail_id="db-1",
|
||||
policy_id=None,
|
||||
page=1,
|
||||
page_size=50,
|
||||
action="passed",
|
||||
start_date=START,
|
||||
end_date=END,
|
||||
user_api_key_dict=ADMIN,
|
||||
)
|
||||
assert resp.logs == []
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue