Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_fix_v1_messages_midstream_timeout_failure_logging

This commit is contained in:
mateo-berri 2026-09-03 16:42:10 -07:00
commit edea0eeb15
70 changed files with 3389 additions and 282 deletions

View file

@ -55,7 +55,10 @@ RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/
ENV PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries
COPY docker/prod_entrypoint.sh /app/docker/prod_entrypoint.sh
RUN sed -i 's/\r$//' /app/docker/prod_entrypoint.sh && chmod +x /app/docker/prod_entrypoint.sh
EXPOSE 4000/tcp
ENTRYPOINT ["litellm"]
ENTRYPOINT ["/app/docker/prod_entrypoint.sh"]
CMD ["--port", "4000"]

View file

@ -1,8 +1,10 @@
#!/bin/sh
if [ "$USE_DDTRACE" = "true" ]; then
export DD_TRACE_OPENAI_ENABLED="False"
exec ddtrace-run "$@"
fi
case "$USE_DDTRACE" in
[Tt][Rr][Uu][Ee])
export DD_TRACE_OPENAI_ENABLED="False"
exec ddtrace-run "$@"
;;
esac
exec "$@"

View file

@ -1,8 +1,10 @@
#!/bin/sh
if [ "$USE_DDTRACE" = "true" ]; then
export DD_TRACE_OPENAI_ENABLED="False"
exec ddtrace-run litellm "$@"
else
exec litellm "$@"
fi
case "$USE_DDTRACE" in
[Tt][Rr][Uu][Ee])
export DD_TRACE_OPENAI_ENABLED="False"
exec ddtrace-run litellm "$@"
;;
esac
exec litellm "$@"

View file

@ -48,6 +48,18 @@ def _build_json_field_or_condition(json_key: str, value: str) -> dict[str, objec
}
def _build_search_condition(search: str) -> dict[str, object]:
"""Match a row whose id, changed_by, object_id, or changed_by_api_key equals the search value."""
return {
"OR": (
{"id": search},
{"changed_by": search},
{"object_id": search},
{"changed_by_api_key": search},
)
}
@router.get(
"/audit",
tags=["Audit Logging"],
@ -83,6 +95,10 @@ async def get_audit_logs(
None,
description="Filter by token (key hash) present in before_value or updated_values JSON (PostgreSQL only)",
),
search: str | None = Query(
None,
description="Match a row whose id, object_id, changed_by, or changed_by_api_key equals this value",
),
# Sorting parameters
sort_by: str | None = Query(
None,
@ -118,6 +134,11 @@ async def get_audit_logs(
*([_build_json_field_or_condition("token", object_key_hash)] if object_key_hash else []),
]
and_conditions: Final[tuple[dict[str, object], ...]] = (
*json_field_conditions,
*((_build_search_condition(search),) if search else ()),
)
# Build filter conditions
where_conditions: Final[dict[str, object]] = {
**({"changed_by": changed_by} if changed_by else {}),
@ -126,14 +147,14 @@ async def get_audit_logs(
**({"table_name": table_name} if table_name else {}),
**({"object_id": object_id} if object_id else {}),
**({"updated_at": date_filter} if start_date or end_date else {}),
**({"AND": json_field_conditions} if json_field_conditions else {}),
**({"AND": and_conditions} if and_conditions else {}),
}
order_by: Final[dict[str, str]] = (
{sort_by: sort_order} if sort_by and isinstance(sort_by, str) else {"updated_at": sort_order}
)
audit_log_table: Final[TableActions["prisma_models.LiteLLM_AuditLog"]] = AuditLogRepository(prisma_client).table
audit_log_table: Final[TableActions[prisma_models.LiteLLM_AuditLog]] = AuditLogRepository(prisma_client).table
# Get paginated results
audit_logs: Final = await audit_log_table.find_many(
@ -195,7 +216,7 @@ async def get_audit_log_by_id(
detail={"message": CommonProxyErrors.db_not_connected_error.value},
)
audit_log_table: Final[TableActions["prisma_models.LiteLLM_AuditLog"]] = AuditLogRepository(prisma_client).table
audit_log_table: Final[TableActions[prisma_models.LiteLLM_AuditLog]] = AuditLogRepository(prisma_client).table
# Get the audit log by ID
audit_log: Final = await audit_log_table.find_unique(where={"id": id})

View file

@ -50,10 +50,12 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
MAX_FILE_LIST_LIMIT,
_is_base64_encoded_unified_file_id,
apply_unified_file_ids,
decode_model_from_file_id,
ensure_batch_response_managed_file_ids,
get_batch_id_from_unified_batch_id,
get_content_type_from_file_object,
get_model_id_from_unified_batch_id,
get_original_file_id,
map_raw_file_ids_to_unified,
normalize_mime_type_for_provider,
resolve_managed_output_file_model_name,
@ -427,6 +429,103 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
detail=f"Object not found: {unified_object_id}",
)
async def enforce_batch_object_access(
self, object_id: str, user_api_key_dict: UserAPIKeyAuth
) -> None:
"""Deny access to a provider-format batch id owned by another caller.
Ids with no ownership row (batches created before ownership tracking,
or directly on the provider account) stay accessible so pass-through
reads keep working.
"""
if self.prisma_client is None:
return
managed_object = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={"OR": [{"unified_object_id": object_id}, {"model_object_id": object_id}]}
)
)
if managed_object is None:
return
if not can_access_resource(
user_api_key_dict=user_api_key_dict,
created_by=managed_object.created_by,
resource_team_id=managed_object.team_id,
):
raise HTTPException(
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to the object {object_id}",
)
async def enforce_provider_file_access(
self, file_id: str, user_api_key_dict: UserAPIKeyAuth
) -> None:
"""Deny access to a provider-format file id owned by another caller.
Ownership rows for provider-format ids are written when a managed
batch's output/error files are first synced; ids with no row stay
accessible so pass-through reads keep working.
"""
if self.prisma_client is None:
return
managed_file = (
await self.prisma_client.db.litellm_managedfiletable.find_first(
where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]}
)
)
if managed_file is None:
return
if not can_access_resource(
user_api_key_dict=user_api_key_dict,
created_by=managed_file.created_by,
resource_team_id=managed_file.team_id,
):
raise HTTPException(
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to the file {file_id}",
)
async def store_batch_output_file_ownership(
self, response: LiteLLMBatch, litellm_parent_otel_span: Optional[Span]
) -> None:
"""Record ownership rows for a batch's provider-format output/error
file ids, inherited from the owning batch row (never the caller), so
file reads can be isolation-checked."""
provider_file_ids = tuple(
file_id
for file_id in (
getattr(response, "output_file_id", None),
getattr(response, "error_file_id", None),
)
if file_id and not _is_base64_encoded_unified_file_id(file_id)
)
if not provider_file_ids:
return
if self.prisma_client is None:
return
batch_row = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={"unified_object_id": response.id}
)
)
if batch_row is None or (
batch_row.created_by is None and batch_row.team_id is None
):
return
owner_identity = UserAPIKeyAuth(
user_id=batch_row.created_by, team_id=batch_row.team_id
)
for file_id in provider_file_ids:
model_name = decode_model_from_file_id(file_id)
raw_file_id = get_original_file_id(file_id)
await self.store_unified_file_id(
file_id=file_id,
file_object=None,
litellm_parent_otel_span=litellm_parent_otel_span,
model_mappings={model_name: raw_file_id} if model_name else {},
user_api_key_dict=owner_identity,
)
async def list_user_batches(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -613,6 +712,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
status_code=403,
detail=f"User {user_api_key_dict.user_id} does not have access to the file {retrieve_file_id}",
)
if retrieve_file_id:
await self.enforce_provider_file_access(
retrieve_file_id, user_api_key_dict
)
return False
async def check_file_ids_access(self, file_ids: List[str], user_api_key_dict: UserAPIKeyAuth) -> None:
@ -765,6 +868,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
data["model"] = potential_model_id
data[accessor_key] = get_batch_id_from_unified_batch_id(potential_llm_object_id)
elif retrieve_object_id and accessor_key == "batch_id":
await self.enforce_batch_object_access(retrieve_object_id, user_api_key_dict)
elif call_type == CallTypes.acreate_fine_tuning_job.value:
input_file_id = cast(Optional[str], data.get("training_file"))
if input_file_id:
@ -1297,7 +1402,13 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
user_api_key_dict=user_api_key_dict,
request_tags=request_tags_from_metadata(request_metadata if isinstance(request_metadata, dict) else {}),
persist_attribution=is_batch_create,
create_if_missing=is_batch_create,
)
if not is_batch_create:
await self.store_batch_output_file_ownership(
response=response,
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
)
# Only record batch creation metric on actual create (not retrieve/cancel).
# unified_file_id in _hidden_params is only set by the create_batch endpoint.

View file

@ -27,6 +27,7 @@ from litellm.constants import (
REDIS_CIRCUIT_BREAKER_ENABLED,
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD,
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT,
REDIS_CIRCUIT_BREAKER_TIMEOUT_MIN_DURATION,
)
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
@ -41,6 +42,8 @@ from .base_cache import BaseCache
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from prometheus_client import Counter as _PromCounter
from prometheus_client import Gauge as _PromGauge
from redis.asyncio import Redis, RedisCluster
from redis.asyncio.client import Pipeline
from redis.asyncio.cluster import ClusterPipeline
@ -135,10 +138,20 @@ class RedisCircuitBreaker:
HALF_OPEN - recovery probe: allow one request through
Transitions:
CLOSED -> OPEN after failure_threshold consecutive failures
CLOSED -> OPEN after failure_threshold consecutive hard connectivity
failures, or after an unbroken run of timeout failures
(no success or hard failure in between) that reaches
failure_threshold and spans timeout_min_duration seconds
OPEN -> HALF_OPEN after recovery_timeout seconds
HALF_OPEN -> CLOSED on success
HALF_OPEN -> OPEN on failure (resets timer)
Timeouts are accounted separately from hard connectivity failures because the async
Redis timeout includes time waiting for the worker event loop to resume: one loop
stall makes every in-flight operation time out together, which satisfies a purely
consecutive threshold instantly even though Redis is healthy. Requiring a
timeout-only streak to also span timeout_min_duration filters such bursts while a
real outage that surfaces as timeouts still opens the breaker after that duration.
"""
CLOSED = "closed"
@ -150,13 +163,19 @@ class RedisCircuitBreaker:
failure_threshold: int,
recovery_timeout: int,
enabled: bool = True,
timeout_min_duration: float = REDIS_CIRCUIT_BREAKER_TIMEOUT_MIN_DURATION,
) -> None:
self.failure_threshold = failure_threshold
self.recovery_timeout = recovery_timeout
self.enabled = enabled
self.timeout_min_duration = timeout_min_duration
self._failure_count = 0
self._hard_failure_count = 0
self._timeout_count = 0
self._timeout_streak_started_at: float | None = None
self._opened_at: float | None = None
self._state = self.CLOSED
_breaker_metrics().record_state_change(None, self._state)
def is_open(self) -> bool:
"""Returns True if Redis calls should be skipped."""
@ -169,24 +188,45 @@ class RedisCircuitBreaker:
return True
if self._state == self.OPEN:
if time.time() - (self._opened_at or 0) > self.recovery_timeout:
self._state = self.HALF_OPEN
self._set_state(self.HALF_OPEN)
return False # this caller is the designated probe
return True
return False
def record_failure(self) -> None:
def _should_open(self, now: float) -> bool:
if self._state == self.HALF_OPEN:
return True
if self._hard_failure_count >= self.failure_threshold:
return True
if self._timeout_count < self.failure_threshold:
return False
return now - (self._timeout_streak_started_at or now) >= self.timeout_min_duration
def record_failure(self, is_timeout: bool = False) -> None:
if not self.enabled:
return
now: Final = time.time()
self._failure_count += 1
self._opened_at = time.time()
if self._failure_count >= self.failure_threshold:
if is_timeout:
self._timeout_count += 1
if self._timeout_streak_started_at is None:
self._timeout_streak_started_at = now
else:
self._hard_failure_count += 1
self._timeout_count = 0
self._timeout_streak_started_at = None
self._opened_at = now
_breaker_metrics().record_failure("timeout" if is_timeout else "connectivity")
if self._should_open(now):
if self._state != self.OPEN:
verbose_logger.warning(
"Redis circuit breaker OPENED after %d consecutive failures — fast-failing Redis calls for %ds",
"Redis circuit breaker OPENED after %d consecutive failures"
" (%d hard connectivity) — fast-failing Redis calls for %ds",
self._failure_count,
self._hard_failure_count,
self.recovery_timeout,
)
self._state = self.OPEN
self._set_state(self.OPEN)
def record_success(self) -> None:
if not self.enabled:
@ -194,7 +234,17 @@ class RedisCircuitBreaker:
if self._state == self.HALF_OPEN:
verbose_logger.info("Redis circuit breaker CLOSED — Redis recovered")
self._failure_count = 0
self._state = self.CLOSED
self._hard_failure_count = 0
self._timeout_count = 0
self._timeout_streak_started_at = None
self._set_state(self.CLOSED)
def _set_state(self, state: str) -> None:
if state == self._state:
return
_breaker_metrics().record_transition(state)
_breaker_metrics().record_state_change(self._state, state)
self._state = state
_RedisCallResult = TypeVar("_RedisCallResult")
@ -234,6 +284,78 @@ def _is_redis_health_failure(exc: BaseException) -> bool:
return True
@functools.lru_cache(maxsize=1)
def _redis_timeout_error_types() -> tuple[type, ...]:
"""Health failures that are timeouts rather than unambiguous connectivity errors.
``builtins.TimeoutError`` covers ``asyncio.TimeoutError`` and ``socket.timeout``
(aliases since py3.11 / py3.10). ``redis.exceptions.TimeoutError`` does not subclass
either, so it is listed explicitly.
"""
try:
from redis.exceptions import TimeoutError as RedisTimeoutError
except ImportError:
return (TimeoutError,)
return (RedisTimeoutError, TimeoutError)
def _is_redis_timeout_failure(exc: BaseException) -> bool:
return isinstance(exc, _redis_timeout_error_types())
class _BreakerMetrics:
"""Prometheus metrics for the Redis circuit breaker; no-ops when the client is absent.
Registered lazily on the default registry (which /metrics serves) via the module-level
``_breaker_metrics`` singleton so repeated RedisCache construction never re-registers.
"""
def __init__(self) -> None:
self._state_gauge: _PromGauge | None = None
self._transitions: _PromCounter | None = None
self._failures: _PromCounter | None = None
try:
from prometheus_client import Counter as PromCounter
from prometheus_client import Gauge
except ImportError:
return
self._state_gauge = Gauge(
"litellm_redis_circuit_breaker_state",
"Number of Redis circuit breakers currently in each state",
labelnames=("state",),
)
self._transitions = PromCounter(
"litellm_redis_circuit_breaker_transitions",
"Redis circuit breaker state transitions",
labelnames=("state",),
)
self._failures = PromCounter(
"litellm_redis_circuit_breaker_failures",
"Redis health failures counted by the circuit breaker",
labelnames=("failure_class",),
)
def record_state_change(self, old_state: str | None, new_state: str) -> None:
if self._state_gauge is None:
return
if old_state is not None:
self._state_gauge.labels(old_state).dec()
self._state_gauge.labels(new_state).inc()
def record_transition(self, state: str) -> None:
if self._transitions is not None:
self._transitions.labels(state).inc()
def record_failure(self, failure_class: str) -> None:
if self._failures is not None:
self._failures.labels(failure_class).inc()
@functools.lru_cache(maxsize=1)
def _breaker_metrics() -> _BreakerMetrics:
return _BreakerMetrics()
def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseException) -> None:
"""Record a Redis failure that the calling method is about to swallow.
@ -245,7 +367,7 @@ def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseExcep
"""
if not _is_redis_health_failure(exc):
return
breaker.record_failure()
breaker.record_failure(is_timeout=_is_redis_timeout_failure(exc))
_swallowed_redis_failures.set(_swallowed_redis_failures.get() + 1)
@ -281,7 +403,7 @@ async def _run_under_circuit_breaker(
result: Final = await call()
except Exception as e:
if _is_redis_health_failure(e):
breaker.record_failure()
breaker.record_failure(is_timeout=_is_redis_timeout_failure(e))
raise
_exit_circuit_breaker(breaker, swallowed_before)
return result

View file

@ -432,6 +432,9 @@ REDIS_CONNECTION_POOL_TIMEOUT: Final = int(os.getenv("REDIS_CONNECTION_POOL_TIME
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD: Final = int(os.getenv("REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD", 5))
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT: Final = int(os.getenv("REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT", 60))
REDIS_CIRCUIT_BREAKER_ENABLED: Final = os.getenv("REDIS_CIRCUIT_BREAKER_ENABLED", "true").lower() == "true"
# minimum seconds a timeout-only failure streak must span before it can open the breaker,
# so one event-loop stall timing out many queued calls at once does not trip it
REDIS_CIRCUIT_BREAKER_TIMEOUT_MIN_DURATION: Final = float(os.getenv("REDIS_CIRCUIT_BREAKER_TIMEOUT_MIN_DURATION", 5.0))
# Seconds of idle before a Redis cluster connection is validated with a PING and
# reconnected if dead, so a connection silently dropped by a cluster restart
# (e.g. ElastiCache Serverless maintenance) is not reused while broken

View file

@ -70,7 +70,7 @@ def to_basic_auth(auth_value: str) -> str:
def strip_auth_scheme(auth_value: str, scheme: str) -> str:
"""Return ``auth_value`` with a leading ``<scheme> `` removed, or unchanged when absent.
"""Return ``auth_value`` with a leading ``<scheme>`` and separator removed, or unchanged when absent.
Callers supply both a bare credential and a complete header value, so prefixing
unconditionally yields ``Bearer Bearer <jwt>``. Scheme names are case-insensitive per
@ -78,10 +78,9 @@ def strip_auth_scheme(auth_value: str, scheme: str) -> str:
with the scheme text and a scheme with nothing behind it are returned untouched.
Surrounding whitespace is left to ``_strip_header_whitespace`` at header-build time.
"""
scheme_name, _, remainder = auth_value.lstrip().partition(" ")
credential: Final = remainder.lstrip()
if credential and scheme_name.lower() == scheme.lower():
return credential
parts: Final = auth_value.split(None, 1)
if len(parts) == 2 and parts[0].lower() == scheme.lower():
return parts[1]
return auth_value

View file

@ -415,40 +415,64 @@ def _is_off_peak(off_peak: Mapping[str, object], current_time: datetime | None =
return False
def _coerce_off_peak_rate(value: object, default: float) -> float:
@dataclass(frozen=True, slots=True)
class TokenRates:
input_rate: float
output_rate: float
cache_read_rate: float
cache_creation_rate: float
reasoning_rate: float | None
@property
def billed_reasoning_rate(self) -> float:
return self.output_rate if self.reasoning_rate is None else self.reasoning_rate
def _parse_off_peak_rate(value: object) -> float | None:
if isinstance(value, bool):
return default
return None
if isinstance(value, (int, float)):
return float(value)
if isinstance(value, str):
try:
return float(value)
except ValueError:
return default
return default
return None
return None
def apply_off_peak_pricing(
model_info: ModelInfo,
current_time: datetime | None,
prompt_base_cost: float,
completion_base_cost: float,
cache_read_cost: float,
) -> tuple[float, float, float]:
def _off_peak_rate(off_peak: Mapping[str, object], key: str, standard_rate: float) -> float:
parsed: Final = _parse_off_peak_rate(off_peak.get(key))
return standard_rate if parsed is None else parsed
def _open_off_peak_block(model_info: ModelInfo, current_time: datetime | None) -> Mapping[str, object] | None:
off_peak: Final = model_info.get("off_peak_pricing")
if not isinstance(off_peak, Mapping) or not _is_off_peak(off_peak, current_time):
return None
return off_peak
def apply_off_peak_pricing(model_info: ModelInfo, current_time: datetime | None, rates: TokenRates) -> TokenRates:
"""Swap in off-peak per-token rates when the current UTC time is inside one of the model's
off_peak_pricing rules, the every-day hours_utc windows or a day-of-week-qualified entry in
windows. An off-peak rate replaces the rate that would otherwise apply rather than
discounting it, so a model that also has tiered or above-threshold pricing bills the flat
off-peak rate for the whole request while the window is open. Any rate left unset in
off_peak_pricing falls back to the standard rate.
off_peak_pricing falls back to the standard rate, so a block without
output_cost_per_reasoning_token keeps the model's own reasoning rate, or its off-peak output
rate when reasoning has no dedicated rate at all.
"""
off_peak: Final = model_info.get("off_peak_pricing")
if not isinstance(off_peak, Mapping) or not _is_off_peak(off_peak, current_time):
return prompt_base_cost, completion_base_cost, cache_read_cost
return (
_coerce_off_peak_rate(off_peak.get("input_cost_per_token"), prompt_base_cost),
_coerce_off_peak_rate(off_peak.get("output_cost_per_token"), completion_base_cost),
_coerce_off_peak_rate(off_peak.get("cache_read_input_token_cost"), cache_read_cost),
off_peak: Final = _open_off_peak_block(model_info, current_time)
if off_peak is None:
return rates
off_peak_reasoning_rate: Final = _parse_off_peak_rate(off_peak.get("output_cost_per_reasoning_token"))
return TokenRates(
input_rate=_off_peak_rate(off_peak, "input_cost_per_token", rates.input_rate),
output_rate=_off_peak_rate(off_peak, "output_cost_per_token", rates.output_rate),
cache_read_rate=_off_peak_rate(off_peak, "cache_read_input_token_cost", rates.cache_read_rate),
cache_creation_rate=_off_peak_rate(off_peak, "cache_creation_input_token_cost", rates.cache_creation_rate),
reasoning_rate=rates.reasoning_rate if off_peak_reasoning_rate is None else off_peak_reasoning_rate,
)
@ -458,14 +482,28 @@ def _apply_off_peak_to_base_costs(
base_costs: tuple[float, float, float, float, float],
) -> tuple[float, float, float, float, float]:
"""Apply off-peak rates to an already-resolved set of base costs, whichever pricing path
produced them. Cache-creation rates are passed through untouched, since off_peak_pricing
has no field for them.
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.
"""
prompt, completion, cache_creation, cache_creation_above_1hr, cache_read = base_costs
off_peak_prompt, off_peak_completion, off_peak_cache_read = apply_off_peak_pricing(
model_info, current_time, prompt, completion, cache_read
rates: Final = apply_off_peak_pricing(
model_info,
current_time,
TokenRates(
input_rate=prompt,
output_rate=completion,
cache_read_rate=cache_read,
cache_creation_rate=cache_creation,
reasoning_rate=None,
),
)
return (
rates.input_rate,
rates.output_rate,
rates.cache_creation_rate,
cache_creation_above_1hr,
rates.cache_read_rate,
)
return (off_peak_prompt, off_peak_completion, cache_creation, cache_creation_above_1hr, off_peak_cache_read)
def _get_token_base_cost(
@ -1029,6 +1067,29 @@ def _resolve_reasoning_token_cost(
return standard_reasoning_cost if standard_reasoning_cost is not None else completion_base_cost
def _resolve_billed_reasoning_rate(
model_info: ModelInfo,
usage: Usage,
service_tier: str | None,
completion_base_cost: float,
current_time: datetime | None,
) -> float:
off_peak: Final = _open_off_peak_block(model_info, current_time)
off_peak_reasoning_rate: Final = (
None if off_peak is None else _parse_off_peak_rate(off_peak.get("output_cost_per_reasoning_token"))
)
if off_peak_reasoning_rate is not None:
return off_peak_reasoning_rate
tiered_reasoning_rate: Final = _get_tiered_reasoning_rate(model_info=model_info, usage=usage)
if tiered_reasoning_rate is not None:
return tiered_reasoning_rate
return _resolve_reasoning_token_cost(
model_info=model_info,
service_tier=service_tier,
completion_base_cost=completion_base_cost,
)
def generic_cost_per_token(
model: str,
usage: Usage,
@ -1037,6 +1098,7 @@ def generic_cost_per_token(
data_residency: str | None = None,
model_info: ModelInfo | None = None,
vertex_location: str | None = None,
current_time: datetime | None = None,
) -> tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -1051,6 +1113,7 @@ def generic_cost_per_token(
- vertex_location: optional Vertex AI location the request was served from
(e.g. "us-east5", "global"), used to apply the per-model
regional-endpoint uplift multiplier when non-global.
- current_time: the moment the request is billed at, for off_peak_pricing; defaults to now, UTC
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
@ -1117,6 +1180,7 @@ def generic_cost_per_token(
usage.prompt_tokens - cache_hit - audio_tokens - cache_creation - image_tokens - video_tokens, 0
)
billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
(
prompt_base_cost,
completion_base_cost,
@ -1127,6 +1191,7 @@ def generic_cost_per_token(
model_info=model_info,
usage=usage,
service_tier=service_tier,
current_time=billing_time,
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
)
@ -1185,17 +1250,13 @@ def generic_cost_per_token(
## REASONING COST
if not is_text_tokens_total and reasoning_tokens and reasoning_tokens > 0:
tiered_reasoning_rate: Final = _get_tiered_reasoning_rate(model_info=model_info, usage=usage)
_output_cost_per_reasoning_token = (
tiered_reasoning_rate
if tiered_reasoning_rate is not None
else _resolve_reasoning_token_cost(
model_info=model_info,
service_tier=service_tier,
completion_base_cost=completion_base_cost,
)
completion_cost += float(reasoning_tokens) * _resolve_billed_reasoning_rate(
model_info=model_info,
usage=usage,
service_tier=service_tier,
completion_base_cost=completion_base_cost,
current_time=billing_time,
)
completion_cost += float(reasoning_tokens) * _output_cost_per_reasoning_token
## IMAGE COST
if not is_text_tokens_total and image_tokens and image_tokens > 0:
@ -1247,6 +1308,7 @@ def get_token_type_cost_breakdown(
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
current_time: datetime | None = None,
) -> TokenTypeCostBreakdown:
"""
Provider-agnostic cost of reasoning and cache tokens, derived from the usage
@ -1265,6 +1327,7 @@ def get_token_type_cost_breakdown(
except Exception:
return TokenTypeCostBreakdown(0.0, 0.0, 0.0)
billing_time: Final = current_time if current_time is not None else datetime.now(timezone.utc)
(
_prompt_base_cost,
completion_base_cost,
@ -1275,6 +1338,7 @@ def get_token_type_cost_breakdown(
model_info=model_info,
usage=usage,
service_tier=service_tier,
current_time=billing_time,
threshold_is_inclusive=_uses_inclusive_token_thresholds(custom_llm_provider),
)
@ -1284,18 +1348,12 @@ def get_token_type_cost_breakdown(
if not reasoning_tokens:
reasoning_tokens = _coerce_token_count(getattr(usage, "reasoning_tokens", 0))
# Reasoning is billed at the selected tier's reasoning rate for tiered models,
# else at the service-tier-aware per-reasoning-token rate - this mirrors how the
# total completion cost is computed, so the breakdown can never diverge from it.
tiered_reasoning_rate: Final = _get_tiered_reasoning_rate(model_info=model_info, usage=usage)
reasoning_rate: Final = (
tiered_reasoning_rate
if tiered_reasoning_rate is not None
else _resolve_reasoning_token_cost(
model_info=model_info,
service_tier=service_tier,
completion_base_cost=completion_base_cost,
)
reasoning_rate: Final = _resolve_billed_reasoning_rate(
model_info=model_info,
usage=usage,
service_tier=service_tier,
completion_base_cost=completion_base_cost,
current_time=billing_time,
)
reasoning_cost = float(reasoning_tokens) * reasoning_rate

View file

@ -6,6 +6,7 @@ from litellm.types.llms.anthropic import AnthropicMessagesRequestOptionalParams
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
from litellm.types.llms.openai import ChatCompletionSystemMessage
if TYPE_CHECKING:
from litellm.exceptions import ContentPolicyViolationError
@ -36,6 +37,16 @@ def safeguard_refusal_error(model: str, stop_details: Mapping[str, object]) -> "
)
def anthropic_system_to_openai_message(system: object) -> ChatCompletionSystemMessage | None:
"""
Return the Anthropic Messages top-level ``system`` (a string or a list of text
blocks) as an OpenAI-style system message, or None when the request has none.
"""
if not isinstance(system, (str, list)) or not system:
return None
return ChatCompletionSystemMessage(role="system", content=system)
@lru_cache(maxsize=1)
def _anthropic_messages_optional_param_keys() -> frozenset[str]:
"""

View file

@ -7,12 +7,13 @@ cached, cache-creation, output, reasoning) is billed at that one tier's rate.
See https://help.aliyun.com/zh/model-studio/billing-for-model-studio
"""
from dataclasses import dataclass, replace
from dataclasses import dataclass
from datetime import datetime
from typing import Final
from litellm.litellm_core_utils.llm_cost_calc.tiered_pricing import select_tier_for_input, tier_rate
from litellm.litellm_core_utils.llm_cost_calc.utils import (
TokenRates,
apply_off_peak_pricing,
parse_completion_tokens_details,
parse_prompt_tokens_details,
@ -34,19 +35,6 @@ class TokenBreakdown:
return self.text_tokens + self.cached_tokens + self.cache_creation_tokens
@dataclass(frozen=True, slots=True)
class TokenRates:
input_rate: float
cache_read_rate: float
cache_creation_rate: float
output_rate: float
reasoning_rate: float | None
@property
def billed_reasoning_rate(self) -> float:
return self.output_rate if self.reasoning_rate is None else self.reasoning_rate
def _extract_token_breakdown(usage: Usage) -> TokenBreakdown:
prompt_details: Final = parse_prompt_tokens_details(usage)
cached_tokens: Final = prompt_details["cache_hit_tokens"]
@ -105,13 +93,6 @@ def _tier_rates(model_info: ModelInfo, tier: dict) -> TokenRates:
)
def _off_peak_rates(model_info: ModelInfo, current_time: datetime | None, rates: TokenRates) -> TokenRates:
input_rate, output_rate, cache_read_rate = apply_off_peak_pricing(
model_info, current_time, rates.input_rate, rates.output_rate, rates.cache_read_rate
)
return replace(rates, input_rate=input_rate, output_rate=output_rate, cache_read_rate=cache_read_rate)
def _bill(breakdown: TokenBreakdown, rates: TokenRates) -> tuple[float, float]:
prompt_cost: Final = (
(breakdown.text_tokens * rates.input_rate)
@ -155,6 +136,6 @@ def cost_per_token(
else None
)
standard_rates: Final = _flat_rates(model_info) if tier is None else _tier_rates(model_info, tier)
rates: Final = _off_peak_rates(model_info, current_time, standard_rates)
rates: Final = apply_off_peak_pricing(model_info, current_time, standard_rates)
return _bill(breakdown, rates)

View file

@ -8,7 +8,7 @@ from urllib.parse import urlparse
import litellm
from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str
from .common_utils import OpenAIError
from .common_utils import OpenAIError, is_openai_backed_api_base
if TYPE_CHECKING:
from collections.abc import Callable
@ -16,7 +16,6 @@ if TYPE_CHECKING:
from openai.auth import SubjectTokenProvider, WorkloadIdentity, WorkloadIdentityAuth
OPENAI_WIF_CLIENT_ID: Final = "litellm"
_OPENAI_API_HOST: Final = "api.openai.com"
_SDK_UPGRADE_MESSAGE: Final = (
"OpenAI workload identity federation requires openai>=2.32.0. "
"Upgrade the installed openai package to use OPENAI_IDENTITY_PROVIDER_ID / "
@ -75,7 +74,7 @@ def _targets_openai_api(api_base: str | None) -> bool:
if api_base is None:
return True
parsed: Final = urlparse(api_base)
return parsed.scheme == "https" and parsed.hostname == _OPENAI_API_HOST
return parsed.scheme == "https" and is_openai_backed_api_base(api_base)
@lru_cache(maxsize=16)

View file

@ -47,7 +47,7 @@ from litellm.constants import (
MCP_TOOL_LISTING_TIMEOUT,
)
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme, to_basic_credentials
from litellm.integrations.custom_guardrail import (
_sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic
)
@ -2299,13 +2299,13 @@ class MCPServerManager:
from litellm.types.mcp import MCPAuth
if server.auth_type == MCPAuth.bearer_token:
headers["Authorization"] = f"Bearer {server.authentication_token}"
headers["Authorization"] = f"Bearer {strip_auth_scheme(server.authentication_token, 'Bearer')}"
elif server.auth_type == MCPAuth.api_key:
headers["Authorization"] = f"ApiKey {server.authentication_token}"
headers["Authorization"] = f"ApiKey {strip_auth_scheme(server.authentication_token, 'ApiKey')}"
elif server.auth_type == MCPAuth.basic:
headers["Authorization"] = f"Basic {server.authentication_token}"
headers["Authorization"] = f"Basic {to_basic_credentials(server.authentication_token)}"
elif server.auth_type == MCPAuth.token:
headers["Authorization"] = f"token {server.authentication_token}"
headers["Authorization"] = f"token {strip_auth_scheme(server.authentication_token, 'token')}"
# Add any static headers from server config.
#
@ -3346,9 +3346,7 @@ class MCPServerManager:
normalized: Final = {k.lower(): v for k, v in raw_headers.items()}
auth_value = normalized.get("authorization")
if auth_value:
if auth_value.startswith("Bearer "):
return auth_value[len("Bearer ") :]
return auth_value
return strip_auth_scheme(auth_value, "Bearer")
return None
@staticmethod

View file

@ -18,6 +18,7 @@ from fastapi import HTTPException
from pydantic import SecretStr
from typing_extensions import assert_never
from litellm.experimental_mcp_client.client import strip_auth_scheme, to_basic_credentials
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
DEFAULT_CREDENTIAL_HEADER,
@ -213,7 +214,9 @@ def _shared_key_spec(
token: Final = server.authentication_token
if not token:
return None # no key configured -> defer to v1 (parity-safe)
value: Final = base64.b64encode(token.encode("utf-8")).decode() if encode else token
value: Final = (
to_basic_credentials(token) if encode else strip_auth_scheme(token, value_prefix) if value_prefix else token
)
return ServerSpec(
server_id=server.server_id,
resource=resource,

View file

@ -5788,7 +5788,7 @@ async def vector_store_access_check(
vector_store_ids_to_run: Final = litellm.vector_store_registry.get_vector_store_ids_to_run(
non_default_params=request_body, tools=request_body.get("tools", None)
)
if vector_store_ids_to_run is None:
if not vector_store_ids_to_run:
verbose_proxy_logger.debug("Vector store to run not found, skipping vector store access check")
return True

View file

@ -1,4 +1,5 @@
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import Final, Protocol
from fastapi import APIRouter, Depends, HTTPException, status
@ -20,7 +21,7 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.management_helpers.access_group_team_sync import invalidate_access_group_cache
from litellm.proxy.utils import get_prisma_client_or_throw
from litellm.repositories.table_repositories import AccessGroupRepository
from litellm.repositories.table_repositories import AccessGroupRepository, TeamRepository
from litellm.types.access_group import (
AccessGroupCreateRequest,
AccessGroupResponse,
@ -74,11 +75,11 @@ class _AccessGroupTable(Protocol):
class _TeamTable(Protocol):
async def find_unique(self, where: Mapping[str, object]) -> _TeamRecord | None: ...
async def find_unique(self, *, where: Mapping[str, object]) -> _TeamRecord | None: ...
async def find_many(self, where: Mapping[str, object]) -> Sequence[_TeamRecord]: ...
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_TeamRecord]: ...
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> object: ...
async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> object: ...
class _KeyTable(Protocol):
@ -119,8 +120,56 @@ def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None:
)
def _record_to_response(record: _AccessGroupRecord) -> AccessGroupResponse:
return AccessGroupResponse.model_validate(record.dict())
def _record_to_response(
record: _AccessGroupRecord, *, assigned_team_ids: Sequence[str] | None = None
) -> AccessGroupResponse:
stored: Final = record.dict()
payload: Final = (
stored if assigned_team_ids is None else MappingProxyType({**stored, "assigned_team_ids": assigned_team_ids})
)
return AccessGroupResponse.model_validate(payload)
def _attached_team_ids_by_group(
records: Sequence[_AccessGroupRecord], teams: Sequence[_TeamRecord]
) -> Mapping[str, tuple[str, ...]]:
"""Teams really attached to each group: the stored column minus ghosts, plus teams the mirror missed."""
real_team_ids: Final = frozenset(team.team_id for team in teams)
def attached(record: _AccessGroupRecord) -> tuple[str, ...]:
stored: Final = (team_id for team_id in (record.assigned_team_ids or ()) if team_id in real_team_ids)
carrying: Final = (team.team_id for team in teams if record.access_group_id in (team.access_group_ids or ()))
return tuple(dict.fromkeys((*stored, *carrying)))
return MappingProxyType({record.access_group_id: attached(record) for record in records})
async def _attached_team_ids_for(
team_table: _TeamTable, records: Sequence[_AccessGroupRecord]
) -> Mapping[str, tuple[str, ...]]:
if not records:
return MappingProxyType({})
group_ids: Final = tuple(record.access_group_id for record in records)
stored_team_ids: Final = tuple(
dict.fromkeys(team_id for record in records for team_id in (record.assigned_team_ids or ()))
)
carrying: Final = {"access_group_ids": {"hasSome": group_ids}} # mutable-ok: prisma where is a dict
listed: Final = {"team_id": {"in": stored_team_ids}} # mutable-ok: prisma where is a dict
teams: Final = await team_table.find_many(where={"OR": (carrying, listed)}) # mutable-ok: prisma where is a dict
return _attached_team_ids_by_group(records, teams)
async def _require_teams_exist(tx: _AccessGroupTx, team_ids: Sequence[str]) -> None:
if not team_ids:
return
where: Final = {"team_id": {"in": team_ids}} # mutable-ok: prisma where is a dict
found: Final = await tx.litellm_teamtable.find_many(where=where)
missing: Final = frozenset(team_ids) - frozenset(team.team_id for team in found)
if missing:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unknown team ids: {', '.join(sorted(missing))}",
)
def _record_to_access_group_table(record: _AccessGroupRecord) -> LiteLLM_AccessGroupTable:
@ -330,6 +379,7 @@ async def create_access_group(
status_code=status.HTTP_409_CONFLICT,
detail=f"Access group '{data.access_group_name}' already exists",
)
await _require_teams_exist(tx, data.assigned_team_ids or ())
record: Final = await tx.litellm_accessgrouptable.create(
data={
@ -390,7 +440,8 @@ async def list_access_groups(
table: Final = AccessGroupRepository(prisma_client).table
records: Final = await table.find_many(order={"created_at": "desc"})
return [_record_to_response(r) for r in records]
attached: Final = await _attached_team_ids_for(TeamRepository(prisma_client).table, records)
return [_record_to_response(r, assigned_team_ids=attached[r.access_group_id]) for r in records]
@router.get(
@ -411,7 +462,8 @@ async def get_access_group(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Access group '{access_group_id}' not found",
)
return _record_to_response(record)
attached: Final = await _attached_team_ids_for(TeamRepository(prisma_client).table, (record,))
return _record_to_response(record, assigned_team_ids=attached[record.access_group_id])
@router.put(
@ -461,8 +513,10 @@ async def update_access_group(
status_code=status.HTTP_404_NOT_FOUND,
detail=f"Access group '{access_group_id}' not found",
)
await _require_teams_exist(tx, data.assigned_team_ids or ())
old_team_ids: Final[set[str]] = set(existing.assigned_team_ids or [])
attached: Final = await _attached_team_ids_for(tx.litellm_teamtable, (existing,))
old_team_ids: Final[set[str]] = set(attached[access_group_id])
old_key_ids: Final[set[str]] = set(existing.assigned_key_ids or [])
new_team_ids: Final[set[str]] = (
set(update_fields["assigned_team_ids"] or []) if "assigned_team_ids" in update_fields else old_team_ids

View file

@ -147,6 +147,7 @@ from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyResponse,
BulkUpdateTeamKeysRequest,
FailedKeyUpdate,
KeySearchWhere,
SuccessfulKeyUpdate,
)
from litellm.types.router import Deployment
@ -5800,6 +5801,10 @@ async def list_keys(
None,
description="Filter keys by key alias. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching.",
),
search: str | None = Query(
None,
description="Combined search: matches keys whose token (key hash) equals the value OR whose key_alias contains it (case-insensitive).",
),
return_full_object: bool = Query(False, description="Return full key object"),
include_team_keys: bool = Query(False, description="Include all keys for teams that user is an admin of."),
include_created_by_keys: bool = Query(False, description="Include keys created by the user"),
@ -5943,6 +5948,7 @@ async def list_keys(
agent_id=agent_id,
use_substring_matching=use_substring_matching,
expires_filter=expires if isinstance(expires, str) else None,
search=search,
)
verbose_proxy_logger.debug("Successfully prepared response")
@ -6162,6 +6168,16 @@ def _build_expires_where_clause(expires_filter: str, now: datetime) -> dict[str,
return {"OR": [{"expires": None}, {"expires": {"gte": now}}]}
def _build_key_search_where(search: str) -> KeySearchWhere:
search_where: Final[KeySearchWhere] = {
"OR": (
{"token": search},
{"key_alias": {"contains": search, "mode": "insensitive"}},
)
}
return search_where
def _build_key_filter_conditions(
user_id: str | None,
team_id: str | None,
@ -6177,6 +6193,7 @@ def _build_key_filter_conditions(
agent_id: str | None = None,
use_substring_matching: bool = False,
expires_filter: str | None = None,
search: str | None = None,
) -> Mapping[str, object]:
"""Build filter conditions for key listing.
@ -6266,7 +6283,7 @@ def _build_key_filter_conditions(
# Apply team_id, project_id and access_group_id as global AND filters so they
# narrow results across all visibility conditions (own keys, team keys, etc.)
global_filters: Final[tuple[dict[str, object], ...]] = (
global_filters: Final[tuple[Mapping[str, object], ...]] = (
*(
(
{"key_alias": {"contains": key_alias, "mode": "insensitive"}}
@ -6277,6 +6294,7 @@ def _build_key_filter_conditions(
else ()
),
*(({"token": key_hash},) if key_hash and isinstance(key_hash, str) else ()),
*((_build_key_search_where(search),) if isinstance(search, str) and search else ()),
*(({"team_id": team_id},) if team_id and isinstance(team_id, str) else ()),
*(({"project_id": project_id},) if project_id else ()),
*(({"access_group_ids": {"hasSome": [access_group_id]}},) if access_group_id else ()),
@ -6316,6 +6334,7 @@ async def _list_key_helper(
agent_id: str | None = None,
use_substring_matching: bool = False,
expires_filter: str | None = None,
search: str | None = None,
) -> KeyListResponseObject:
"""
Helper function to list keys
@ -6354,6 +6373,7 @@ async def _list_key_helper(
agent_id=agent_id,
use_substring_matching=use_substring_matching,
expires_filter=expires_filter,
search=search,
)
# Calculate skip for pagination

View file

@ -1123,13 +1123,15 @@ async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_grou
user_id=user_id,
user_email=user_id, # We don't have email from group membership
user_alias=None,
teams=[], # Teams will be added separately
metadata={"created_via": created_via},
auto_create_key=False,
user_role=default_role,
)
created_user: Final = await new_user(data=new_user_request)
created_user: Final = await new_user(
data=new_user_request,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
verbose_proxy_logger.info("Created user %s via %s", user_id, created_via)
return created_user
@ -1699,7 +1701,7 @@ async def create_user(
user_id=user_id,
user_email=user_data["user_email"],
user_alias=user_data["user_alias"],
teams=user_data["teams"],
teams=user_data["teams"] or None,
metadata=metadata,
auto_create_key=False,
user_role=resolved_role if admin_group is not None else default_role,
@ -1717,6 +1719,7 @@ async def create_user(
created_user: Final = await new_user(
data=new_user_request,
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
)
scim_user: Final = await ScimTransformations.transform_litellm_user_to_scim_user(user=created_user)
@ -1771,22 +1774,25 @@ async def update_user(
roles=user_data["roles"],
)
# SCIM User.groups is readOnly (RFC 7643 4.1.2): IdPs sync membership via /Groups and send
# no groups or `groups: []` on profile PUTs, so empty means unspecified, not "remove from every team"
target_teams: Final = user_data["teams"] or existing_user.teams
await _handle_team_membership_changes(
user_id=user_id,
existing_teams=existing_user.teams or [],
new_teams=user_data["teams"],
existing_teams=existing_user.teams,
new_teams=target_teams,
)
update_data: Final = {
"user_email": user_data["user_email"],
"user_alias": user_data["user_alias"],
"sso_user_id": user_data["sso_user_id"],
"teams": user_data["teams"],
"teams": target_teams,
"metadata": safe_dumps(metadata),
}
admin_group: Final = await _get_scim_admin_group()
if admin_group is not None:
if admin_group is not None and user_data["teams"]:
update_data["user_role"] = _resolve_scim_user_role(
user.groups or [], admin_group, _default_scim_user_role()
)

View file

@ -22,6 +22,7 @@ from collections.abc import Mapping
from typing import TYPE_CHECKING, Final
from fastapi import APIRouter, Depends, HTTPException, Query
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
@ -91,6 +92,36 @@ def _visibility_filter(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, object
return {"OR": ors}
class _StartsWith(TypedDict):
startsWith: ReadOnly[str]
class _MemoryKeyWhere(TypedDict):
key: ReadOnly[str | _StartsWith]
class _MemoryIdWhere(TypedDict):
memory_id: ReadOnly[str]
class _MemorySearchWhere(TypedDict):
OR: ReadOnly[tuple[_MemoryKeyWhere, _MemoryIdWhere]]
def _key_filter(search: str | None, key_prefix: str | None, key: str | None) -> Mapping[str, object] | None:
"""`search` matches a key prefix or an exact memory_id; otherwise `key_prefix` wins over `key`."""
if search is not None:
search_where: Final[_MemorySearchWhere] = {"OR": ({"key": {"startsWith": search}}, {"memory_id": search})}
return search_where
if key_prefix is not None:
prefix_where: Final[_MemoryKeyWhere] = {"key": {"startsWith": key_prefix}}
return prefix_where
if key is not None:
exact_where: Final[_MemoryKeyWhere] = {"key": key}
return exact_where
return None
def _row_to_model(row: "prisma_models.LiteLLM_MemoryTable") -> LiteLLM_MemoryRow:
return LiteLLM_MemoryRow(
memory_id=row.memory_id,
@ -326,6 +357,13 @@ async def list_memory(
"Mutually exclusive with `key`; if both are provided, `key_prefix` wins."
),
),
search: str | None = Query(
None,
description=(
"Match entries whose key starts with this value or whose memory_id equals it. "
"Takes precedence over `key_prefix` and `key` when provided."
),
),
page: int = Query(1, ge=1),
page_size: int = Query(50, ge=1, le=500),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
@ -333,22 +371,16 @@ async def list_memory(
"""List memory entries visible to the caller."""
prisma_client: Final = _require_prisma()
# Build the key filter first (prefix wins if both `key` and `key_prefix`
# are passed). Then AND it with the visibility filter via an explicit
# top-level "AND" — safer than `dict.update` since future visibility
# filters could grow an "OR" key that would clobber this one if merged
# by key.
key_filter: Final[dict[str, object]] = {}
if key_prefix is not None:
key_filter["key"] = {"startsWith": key_prefix}
elif key is not None:
key_filter["key"] = key
# AND the key filter with the visibility filter via an explicit top-level
# "AND": both sides can carry an "OR" key (`search`, non-admin visibility),
# so merging them by key would let one clobber the other and leak rows.
key_filter: Final = _key_filter(search=search, key_prefix=key_prefix, key=key)
vis: Final = _visibility_filter(user_api_key_dict)
where: Mapping[str, object]
where: Mapping[str, object] | None
if vis is None:
where = key_filter
elif not key_filter:
elif key_filter is None:
where = vis
else:
where = {"AND": [key_filter, vis]}

View file

@ -2229,6 +2229,31 @@ async def calculate_spend(request: SpendCalculateRequest):
)
class _SpendLogSearchCondition(NamedTuple):
sql: str
params: tuple[object, ...]
def _build_spend_log_search_condition(
search: str,
start_date: datetime,
end_date: datetime,
next_param_index: int,
) -> _SpendLogSearchCondition:
"""request_id (indexed) matches across all time; the unindexed id columns only inside the window."""
raw: Final = f"${next_param_index}"
window_start: Final = f"${next_param_index + 1}"
window_end: Final = f"${next_param_index + 2}"
sql: Final = (
f"(request_id = {raw} OR ("
f"\"startTime\" >= ({window_start}::timestamptz AT TIME ZONE 'UTC') "
f"AND \"startTime\" <= ({window_end}::timestamptz AT TIME ZONE 'UTC') "
f'AND (api_key = {raw} OR team_id = {raw} OR "user" = {raw} OR end_user = {raw} '
f"OR session_id = {raw} OR model_id = {raw})))"
)
return _SpendLogSearchCondition(sql=sql, params=(search, start_date, end_date))
@router.get(
"/spend/logs/v2",
tags=["Budget & Spend Tracking"],
@ -2329,6 +2354,14 @@ async def ui_view_spend_logs(
"UI route only, honored when sorting by startTime"
),
),
search: str | None = fastapi.Query(
default=None,
description=(
"Match a log whose request_id, api_key (hash), team_id, user, end_user, "
"session_id, or model_id equals this value. request_id matches across all time; the other columns "
"match inside start_date/end_date, which stay required"
),
),
):
"""
View spend logs with pagination support.
@ -2392,8 +2425,10 @@ async def ui_view_spend_logs(
try:
is_admin_view: Final = _is_admin_view_safe(user_api_key_dict=user_api_key_dict)
is_request_id_lookup: Final = request_id is not None and not is_v2
is_search_lookup: Final = search is not None
search_owns_window: Final = is_search_lookup and not is_v2
if is_request_id_lookup:
if is_request_id_lookup and not is_search_lookup:
# request_id is the @id primary key: it identifies a single row, so a
# time window is meaningless. The dashboard always sends a default 24h
# window, which hid ids copied from an older page (LIT-3981). Drop the
@ -2576,7 +2611,7 @@ async def ui_view_spend_logs(
# Date range. Wrap the param side with `AT TIME ZONE 'UTC'` so comparison
# against the plain `timestamp` column does not depend on the DB session
# timezone (see #22529). Absent for a request_id-only lookup (see above).
if start_date_obj is not None and end_date_obj is not None:
if start_date_obj is not None and end_date_obj is not None and not search_owns_window:
sql_conditions.append(f"\"startTime\" >= (${p}::timestamptz AT TIME ZONE 'UTC')")
sql_params.append(start_date_obj)
p += 1
@ -2584,6 +2619,17 @@ async def ui_view_spend_logs(
sql_params.append(end_date_obj)
p += 1
if search is not None and start_date_obj is not None and end_date_obj is not None:
search_condition: Final = _build_spend_log_search_condition(
search=search,
start_date=start_date_obj,
end_date=end_date_obj,
next_param_index=p,
)
sql_conditions.append(search_condition.sql)
sql_params.extend(search_condition.params)
p += len(search_condition.params) # rebind-ok: advances the file's shared $N placeholder counter
# Equality filters - read effective values from where_conditions (post-authorization)
for sql_col, wc_key in [
("team_id", "team_id"),
@ -2662,7 +2708,13 @@ async def ui_view_spend_logs(
sql_params.append(f"%{error_message}%")
p += 1
if group_by_session is True and not is_v2 and not is_request_id_lookup and sort_by == "startTime":
if (
group_by_session is True
and not is_v2
and not is_request_id_lookup
and not is_search_lookup
and sort_by == "startTime"
):
return await _ui_session_grouped_spend_logs(
prisma_client=prisma_client,
sql_conditions=sql_conditions,
@ -2696,7 +2748,7 @@ async def ui_view_spend_logs(
_order_expr = order_column
joined_conditions: Final = " AND ".join(sql_conditions)
session_grouping: Final = group_by_session is True
session_grouping: Final = group_by_session is True and not is_search_lookup
count_group_clause: Final = f"GROUP BY {_SESSION_GROUP_KEY_SQL}" if session_grouping else ""
count_query: Final = f"""
SELECT COUNT(*) AS total_count

View file

@ -176,6 +176,10 @@ class PolicyAttachmentRepository(PrismaTableRepository["prisma_models.LiteLLM_Po
table_name = "litellm_policyattachmenttable"
class TeamRepository(PrismaTableRepository["prisma_models.LiteLLM_TeamTable"]):
table_name = "litellm_teamtable"
class DeletedTeamRepository(PrismaTableRepository["prisma_models.LiteLLM_DeletedTeamTable"]):
table_name = "litellm_deletedteamtable"

View file

@ -197,6 +197,7 @@ from litellm.router_utils.router_callbacks.track_deployment_metrics import (
from litellm.scheduler import FlowItem, Scheduler
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolParam,
FileTypes,
OpenAIFileObject,
OpenAIFilesPurpose,
@ -11762,7 +11763,7 @@ class Router:
self,
messages: list[dict[str, str]] | None,
input: str | list | None,
instructions: str | None = None,
request_kwargs: Mapping[str, object] | None = None,
) -> int:
"""
Count input tokens for context-window pre-call checks.
@ -11772,9 +11773,28 @@ class Router:
The Responses payload is normalized to chat messages via the shared
LiteLLMCompletionResponsesConfig transform so the same token_counter path covers
both API surfaces and `instructions` tokens are included in the count.
Prompt content the message list never carries is read from `request_kwargs`:
`tools` (Chat Completions, Responses and Anthropic Messages shapes) and the
Anthropic Messages top-level `system` block.
"""
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
anthropic_system_to_openai_message,
)
extras: Final = request_kwargs if request_kwargs is not None else MappingProxyType({})
raw_instructions: Final = extras.get("instructions")
instructions: Final = raw_instructions if isinstance(raw_instructions, str) else None
raw_tools: Final = extras.get("tools")
tools: Final = (
cast(list[ChatCompletionToolParam], raw_tools) # cast-ok: token_counter formats any tool dict shape
if isinstance(raw_tools, list) and raw_tools
else None
)
system_message: Final = anthropic_system_to_openai_message(extras.get("system"))
if messages is not None:
return litellm.token_counter(messages=messages)
counted_messages: Final = (system_message, *messages) if system_message is not None else messages
return litellm.token_counter(messages=counted_messages, tools=tools)
if input is not None:
from openai.types.responses.response_create_params import ResponseInputParam
@ -11787,7 +11807,10 @@ class Router:
input=typed_input,
responses_api_request={"instructions": instructions} if instructions is not None else {},
)
return litellm.token_counter(messages=cast(list, input_messages)) # cast-ok: transformed chat messages
return litellm.token_counter(
messages=cast(list, input_messages), # cast-ok: transformed chat messages
tools=tools,
)
raise ValueError("Either messages or input must be provided to count tokens")
def _deployment_max_input_tokens(self, model: str, deployment: Mapping[str, object]) -> int | None:
@ -11833,14 +11856,13 @@ class Router:
"""
if messages is None and input is None:
return None
raw_instructions: Final = request_kwargs.get("instructions") if request_kwargs else None
try:
if not self._pre_call_checks_need_token_count(model, healthy_deployments):
return None
return await asyncify(self._count_pre_call_check_tokens)(
messages=cast(list[dict[str, str]] | None, messages), # cast-ok: forwarded to the sync counter
input=cast(str | list | None, input), # cast-ok: forwarded to the sync counter
instructions=raw_instructions if isinstance(raw_instructions, str) else None,
request_kwargs=request_kwargs,
)
except Exception as e: # noqa: BLE001 # best-effort: an uncountable prompt must not fail the request
verbose_router_logger.error(
@ -11887,8 +11909,6 @@ class Router:
_rate_limit_error = False
parent_otel_span: Final = _get_parent_otel_span_from_kwargs(request_kwargs)
raw_instructions: Final = request_kwargs.get("instructions") if request_kwargs else None
instructions: Final = raw_instructions if isinstance(raw_instructions, str) else None
has_countable_input: Final = messages is not None or input is not None
## get model group RPM ##
@ -11919,7 +11939,7 @@ class Router:
return _returned_deployments
try:
input_tokens = self._count_pre_call_check_tokens(
messages=messages, input=input, instructions=instructions
messages=messages, input=input, request_kwargs=request_kwargs
)
except Exception as e:
verbose_router_logger.error(

View file

@ -2,6 +2,23 @@ from datetime import datetime
from typing import Any, Final, Literal
from pydantic import BaseModel, ConfigDict, model_validator
from typing_extensions import ReadOnly, TypedDict
from litellm.types.proxy.management_endpoints.internal_user_endpoints import InsensitiveContains
class KeyTokenWhere(TypedDict):
token: ReadOnly[str]
class KeyAliasContainsWhere(TypedDict):
key_alias: ReadOnly[InsensitiveContains]
class KeySearchWhere(TypedDict):
"""Prisma filter behind `/key/list?search=`: exact token or case-insensitive alias substring."""
OR: ReadOnly[tuple[KeyTokenWhere, KeyAliasContainsWhere]]
class BulkUpdateKeyRequestItem(BaseModel):

View file

@ -222,7 +222,9 @@ class OffPeakPricing(TypedDict, total=False):
weekday_timezone: ReadOnly[str]
input_cost_per_token: ReadOnly[float]
output_cost_per_token: ReadOnly[float]
output_cost_per_reasoning_token: ReadOnly[float]
cache_read_input_token_cost: ReadOnly[float]
cache_creation_input_token_cost: ReadOnly[float]
class ModelInfoBase(ProviderSpecificModelInfo, total=False):

View file

@ -214,8 +214,8 @@ locals {
gateway_uvicorn_args = "--host 0.0.0.0 --port 4000 --workers ${var.gateway_num_workers}"
backend_uvicorn_args = "--host 0.0.0.0 --port 4001"
gateway_launch_cmd = "if [ \"$USE_DDTRACE\" = \"true\" ]; then export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args}; else exec uvicorn gateway.main:app ${local.gateway_uvicorn_args}; fi"
backend_launch_cmd = "if [ \"$USE_DDTRACE\" = \"true\" ]; then export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args}; else exec uvicorn backend.main:app ${local.backend_uvicorn_args}; fi"
gateway_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args};; *) exec uvicorn gateway.main:app ${local.gateway_uvicorn_args};; esac"
backend_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args};; *) exec uvicorn backend.main:app ${local.backend_uvicorn_args};; esac"
gateway_proxy_overrides = local.proxy_config_enabled ? {
entryPoint = ["sh", "-c"]

View file

@ -138,8 +138,8 @@ locals {
gateway_uvicorn_args = "--host 0.0.0.0 --port 4000 --workers ${var.gateway_num_workers}"
backend_uvicorn_args = "--host 0.0.0.0 --port 4001"
gateway_launch_cmd = "if [ \"$USE_DDTRACE\" = \"true\" ]; then export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args}; else exec uvicorn gateway.main:app ${local.gateway_uvicorn_args}; fi"
backend_launch_cmd = "if [ \"$USE_DDTRACE\" = \"true\" ]; then export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args}; else exec uvicorn backend.main:app ${local.backend_uvicorn_args}; fi"
gateway_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args};; *) exec uvicorn gateway.main:app ${local.gateway_uvicorn_args};; esac"
backend_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args};; *) exec uvicorn backend.main:app ${local.backend_uvicorn_args};; esac"
gateway_args = join(" && ", concat(
local.redis_ca_fragment,

View file

@ -11,6 +11,7 @@ from litellm.caching import DualCache
from litellm.proxy._types import CallTypes
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
encode_file_id_with_model,
)
@ -3153,6 +3154,421 @@ async def test_same_user_different_keys_can_access_batch():
assert result1["batch_id"] == result2["batch_id"]
MODEL_ENCODED_BATCH_ID = encode_file_id_with_model(
"batch_provider123", "gpt-4o-team-alias", id_type="batch"
)
MODEL_ENCODED_OUTPUT_FILE_ID = encode_file_id_with_model(
"file-output456", "gpt-4o-team-alias", id_type="file"
)
RAW_PROVIDER_BATCH_ID = "batch_provider123"
RAW_PROVIDER_FILE_ID = "file-output456"
def _owned_record(created_by, team_id):
record = MagicMock()
record.created_by = created_by
record.team_id = team_id
return record
def _batch_response(batch_id, output_file_id=None, is_create=False):
from litellm.types.utils import LiteLLMBatch
batch = LiteLLMBatch(
id=batch_id,
completion_window="24h",
created_at=1700000000,
endpoint="/v1/chat/completions",
input_file_id="file-input789",
object="batch",
status="completed",
output_file_id=output_file_id,
)
if is_create:
batch._hidden_params["unified_file_id"] = "unified-input-file-id"
return batch
@pytest.mark.asyncio
@pytest.mark.parametrize("call_type", ["aretrieve_batch", "acancel_batch"])
@pytest.mark.parametrize(
"batch_id", [MODEL_ENCODED_BATCH_ID, RAW_PROVIDER_BATCH_ID]
)
async def test_team_b_cannot_access_team_a_provider_format_batch(
call_type, batch_id
):
"""
Cross-team retrieve/cancel of a model-encoded or raw provider batch id
must 403 when an ownership row exists for another team.
Regression test: before this check only unified (litellm_proxy-prefixed)
batch ids were enforced, so any key could read any model-encoded or raw
provider batch.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
with pytest.raises(HTTPException) as exc_info:
await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"batch_id": batch_id},
call_type=call_type,
)
assert exc_info.value.status_code == 403
prisma_client.db.litellm_managedobjecttable.find_first.assert_awaited_once_with(
where={"OR": [{"unified_object_id": batch_id}, {"model_object_id": batch_id}]}
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"caller_kwargs",
[
{"user_id": "user_a", "team_id": "team_a"},
{"user_id": "teammate_of_a", "team_id": "team_a"},
{"user_id": "admin_user", "user_role": "proxy_admin"},
],
)
async def test_authorized_callers_can_access_provider_format_batch(caller_kwargs):
"""
The creator, a same-team member, and a proxy admin can all retrieve a
model-encoded batch owned by team_a. Data must pass through unmodified so
the endpoint's own routing still applies.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
parent_otel_span=MagicMock(), **caller_kwargs
),
cache=MagicMock(),
data={"batch_id": MODEL_ENCODED_BATCH_ID},
call_type="aretrieve_batch",
)
assert result["batch_id"] == MODEL_ENCODED_BATCH_ID
assert "model" not in result
@pytest.mark.asyncio
async def test_provider_format_batch_without_ownership_row_stays_accessible():
"""
A provider-format batch id with no ownership row (created before ownership
tracking, or directly on the provider account) must stay retrievable.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first.return_value = None
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"batch_id": RAW_PROVIDER_BATCH_ID},
call_type="aretrieve_batch",
)
assert result["batch_id"] == RAW_PROVIDER_BATCH_ID
@pytest.mark.asyncio
async def test_fine_tuning_provider_format_id_not_enforced():
"""
Provider-format fine-tuning job ids are deliberately out of scope for
ownership enforcement; only unified fine-tuning ids are checked.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
DualCache(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"fine_tuning_job_id": "ftjob-abc123"},
call_type="aretrieve_fine_tuning_job",
)
assert result["fine_tuning_job_id"] == "ftjob-abc123"
prisma_client.db.litellm_managedobjecttable.find_first.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
"call_type", ["afile_content", "afile_retrieve", "afile_delete"]
)
@pytest.mark.parametrize(
"file_id", [MODEL_ENCODED_OUTPUT_FILE_ID, RAW_PROVIDER_FILE_ID]
)
async def test_team_b_cannot_access_team_a_provider_format_file(
call_type, file_id
):
"""
Cross-team content/retrieve/delete of a model-encoded or raw provider
file id must 403 when an ownership row exists for another team.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedfiletable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(), prisma_client=prisma_client
)
with pytest.raises(HTTPException) as exc_info:
await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"file_id": file_id},
call_type=call_type,
)
assert exc_info.value.status_code == 403
prisma_client.db.litellm_managedfiletable.find_first.assert_awaited_once_with(
where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]}
)
@pytest.mark.asyncio
async def test_same_team_can_access_provider_format_file():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedfiletable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="teammate_of_a", team_id="team_a", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"file_id": MODEL_ENCODED_OUTPUT_FILE_ID},
call_type="afile_content",
)
assert result["file_id"] == MODEL_ENCODED_OUTPUT_FILE_ID
@pytest.mark.asyncio
async def test_provider_format_file_without_ownership_row_stays_accessible():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedfiletable.find_first.return_value = None
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(), prisma_client=prisma_client
)
result = await proxy_managed_files.async_pre_call_hook(
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
cache=MagicMock(),
data={"file_id": RAW_PROVIDER_FILE_ID},
call_type="afile_content",
)
assert result["file_id"] == RAW_PROVIDER_FILE_ID
@pytest.mark.asyncio
async def test_post_call_batch_create_stores_ownership_row():
"""
Batch creation (response hidden params carry the unified input file id)
must write an ownership row attributed to the creating key.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={
"input_file_id": "file-input789",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
user_api_key_dict=UserAPIKeyAuth(
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
),
response=_batch_response(MODEL_ENCODED_BATCH_ID, is_create=True),
)
upsert_call = prisma_client.db.litellm_managedobjecttable.upsert.await_args
assert upsert_call.kwargs["where"] == {
"unified_object_id": MODEL_ENCODED_BATCH_ID
}
create_data = upsert_call.kwargs["data"]["create"]
assert create_data["created_by"] == "user_a"
assert create_data["team_id"] == "team_a"
prisma_client.db.litellm_managedobjecttable.update_many.assert_not_awaited()
@pytest.mark.asyncio
async def test_post_call_batch_sync_does_not_claim_ownership():
"""
Retrieve/cancel of a batch with no ownership row must NOT create one:
otherwise the first foreign key to touch a legacy batch would become its
owner and lock out the real creator once enforcement is on.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 0
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={"batch_id": MODEL_ENCODED_BATCH_ID},
user_api_key_dict=UserAPIKeyAuth(
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
),
response=_batch_response(MODEL_ENCODED_BATCH_ID),
)
prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited()
prisma_client.db.litellm_managedobjecttable.update_many.assert_awaited_once()
@pytest.mark.asyncio
async def test_post_call_batch_sync_updates_existing_row():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={"batch_id": MODEL_ENCODED_BATCH_ID},
user_api_key_dict=UserAPIKeyAuth(
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
),
response=_batch_response(MODEL_ENCODED_BATCH_ID),
)
update_call = prisma_client.db.litellm_managedobjecttable.update_many.await_args
assert update_call.kwargs["where"] == {
"unified_object_id": MODEL_ENCODED_BATCH_ID
}
assert update_call.kwargs["data"]["status"] == "completed"
prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited()
@pytest.mark.asyncio
async def test_post_call_batch_sync_stores_output_file_ownership_from_batch_row():
"""
When a synced batch reports a provider-format output file id, an
ownership row for that file must be written with the BATCH row's
created_by/team_id, not the caller's identity.
"""
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
_owned_record(created_by="user_a", team_id="team_a")
)
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={"batch_id": MODEL_ENCODED_BATCH_ID},
user_api_key_dict=UserAPIKeyAuth(
user_id="admin_user",
user_role="proxy_admin",
parent_otel_span=MagicMock(),
),
response=_batch_response(
MODEL_ENCODED_BATCH_ID, output_file_id=MODEL_ENCODED_OUTPUT_FILE_ID
),
)
file_upsert = prisma_client.db.litellm_managedfiletable.upsert.await_args
assert file_upsert.kwargs["where"] == {
"unified_file_id": MODEL_ENCODED_OUTPUT_FILE_ID
}
create_data = file_upsert.kwargs["data"]["create"]
assert create_data["created_by"] == "user_a"
assert create_data["team_id"] == "team_a"
assert create_data["flat_model_file_ids"] == ["file-output456"]
@pytest.mark.asyncio
async def test_post_call_batch_create_does_not_store_output_file_ownership():
from litellm.proxy._types import UserAPIKeyAuth
prisma_client = AsyncMock()
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
)
await proxy_managed_files.async_post_call_success_hook(
data={
"input_file_id": "file-input789",
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
user_api_key_dict=UserAPIKeyAuth(
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
),
response=_batch_response(
MODEL_ENCODED_BATCH_ID,
output_file_id=MODEL_ENCODED_OUTPUT_FILE_ID,
is_create=True,
),
)
prisma_client.db.litellm_managedfiletable.upsert.assert_not_awaited()
@pytest.mark.asyncio
async def test_file_list_cursors_are_scoped_to_the_caller():
"""A non-owner must not learn other callers' file ids through the page cursors."""

View file

@ -1,5 +1,6 @@
from datetime import datetime, timedelta
from unittest.mock import AsyncMock, patch
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import FastAPI
@ -8,10 +9,12 @@ from litellm_enterprise.proxy.audit_logging_endpoints import router as audit_rou
from litellm_enterprise.types.proxy.audit_logging_endpoints import AuditLogResponse
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
# Create an app with just the audit router for testing
app = FastAPI()
app.include_router(audit_router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role="proxy_admin")
client = TestClient(app)
# Mock data for testing
@ -130,3 +133,45 @@ async def test_get_audit_log_by_id_not_found(mock_prisma_client):
data = response.json()
assert "message" in data["detail"]
assert "not found" in data["detail"]["message"].lower()
def _list_audit_logs_where(mock_prisma_client: MagicMock, query: str) -> dict[str, object]:
mock_prisma_client.db.litellm_auditlog.find_many.return_value = []
mock_prisma_client.db.litellm_auditlog.count.return_value = 0
response: Final = client.get(f"/audit?{query}")
assert response.status_code == 200, response.text
find_many_where: Final = mock_prisma_client.db.litellm_auditlog.find_many.call_args.kwargs["where"]
assert mock_prisma_client.db.litellm_auditlog.count.call_args.kwargs["where"] == find_many_where
return find_many_where
def test_search_matches_any_id_column_alongside_the_other_filters(mock_prisma_client):
where: Final = _list_audit_logs_where(mock_prisma_client, "search=abc-123&action=create&object_team_id=team-1")
assert where == {
"action": "create",
"AND": (
{
"OR": [
{"before_value": {"path": ["team_id"], "string_contains": "team-1"}},
{"updated_values": {"path": ["team_id"], "string_contains": "team-1"}},
]
},
{
"OR": (
{"id": "abc-123"},
{"changed_by": "abc-123"},
{"object_id": "abc-123"},
{"changed_by_api_key": "abc-123"},
)
},
),
}
def test_an_empty_search_leaves_the_where_clause_unchanged(mock_prisma_client):
where: Final = _list_audit_logs_where(mock_prisma_client, "action=create&search=")
assert where == {"action": "create"}

View file

@ -763,7 +763,7 @@ class _MigrateDeployHarness:
"_resolve_specific_migration",
staticmethod(self.resolved.append),
)
monkeypatch.setattr(utils_module.subprocess, "run", self._fake_run)
monkeypatch.setattr(utils_module.prisma_toolchain, "run_prisma", self._fake_run)
monkeypatch.setattr(utils_module.time, "sleep", lambda seconds: None)
self.baseline_succeeds = True

View file

@ -779,7 +779,7 @@ async def test_concurrent_success_is_not_cancelled_by_another_calls_failure():
"error, opens_breaker",
[
pytest.param("ConnectionError", True, id="connection_refused_is_unhealthy"),
pytest.param("TimeoutError", True, id="timeout_is_unhealthy"),
pytest.param("TimeoutError", False, id="timeout_burst_is_ambiguous"),
pytest.param("BusyLoadingError", True, id="loading_is_unhealthy"),
pytest.param("ResponseError", False, id="wrong_type_command_is_not"),
pytest.param("DataError", False, id="bad_data_is_not"),
@ -791,6 +791,10 @@ async def test_only_connectivity_failures_open_the_breaker(error, opens_breaker)
They say nothing about connectivity, and a caller able to provoke them (an INCR against
a non-numeric value, say) could otherwise trip the shared breaker on demand and drop
rate limiting to per-process counters, which spreading traffic across replicas outruns.
A rapid burst of timeouts is ambiguous too: the async timeout includes event-loop
scheduling delay, so a loop stall times out every queued call at once against a
healthy Redis. It must not open the breaker until the streak spans a minimum duration.
"""
import redis.exceptions
@ -810,3 +814,147 @@ async def test_only_connectivity_failures_open_the_breaker(error, opens_breaker)
await _run_under_circuit_breaker(breaker, "op", failing_call)
assert breaker.is_open() is opens_breaker
@pytest.mark.asyncio
async def test_event_loop_stall_timeout_burst_keeps_breaker_closed():
"""One blocking stall of the worker event loop must not trip the breaker.
Every operation already waiting on the loop times out together when the loop resumes,
so a purely consecutive threshold is satisfied instantly even though the Redis on the
other end (here an in-process fake that answers immediately) is healthy.
"""
import time as time_mod
from litellm.caching.redis_cache import RedisCircuitBreaker, _run_under_circuit_breaker
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, timeout_min_duration=5.0)
async def healthy_redis_call_with_client_timeout():
return await asyncio.wait_for(asyncio.sleep(0.001, result="ok"), timeout=0.05)
async def stall_the_loop():
await asyncio.sleep(0)
time_mod.sleep(0.2)
results = await asyncio.gather(
*(_run_under_circuit_breaker(breaker, "op", healthy_redis_call_with_client_timeout) for _ in range(8)),
stall_the_loop(),
return_exceptions=True,
)
timeouts = [r for r in results if isinstance(r, asyncio.TimeoutError)]
assert len(timeouts) >= breaker.failure_threshold, "the stall must time out a full burst"
assert breaker.is_open() is False, "a healthy Redis behind one loop stall must stay in the pool"
assert await _run_under_circuit_breaker(breaker, "op", healthy_redis_call_with_client_timeout) == "ok"
@pytest.mark.asyncio
async def test_persistent_timeouts_still_open_the_breaker():
"""A real outage that surfaces only as timeouts must still open the breaker.
Once the timeout-only streak spans the minimum duration with no success in between,
Redis is genuinely unusable from this worker and protection has to kick in.
"""
from redis.exceptions import TimeoutError as RedisTimeoutError
from litellm.caching.redis_cache import RedisCircuitBreaker, _run_under_circuit_breaker
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, timeout_min_duration=0.1)
async def timing_out_call():
raise RedisTimeoutError("read timed out")
for _ in range(breaker.failure_threshold):
with pytest.raises(RedisTimeoutError):
await _run_under_circuit_breaker(breaker, "op", timing_out_call)
assert breaker.is_open() is False, "the burst has not spanned the minimum duration yet"
await asyncio.sleep(0.12)
with pytest.raises(RedisTimeoutError):
await _run_under_circuit_breaker(breaker, "op", timing_out_call)
assert breaker.is_open() is True
@pytest.mark.asyncio
async def test_stale_timeout_does_not_let_sub_threshold_hard_failures_open_the_breaker():
"""Hard connectivity failures below the threshold must not open the breaker just
because an old timeout already started the streak and the duration has elapsed.
Each class has to earn the open on its own terms: hard failures by reaching the
threshold, timeouts by reaching the threshold and spanning the minimum duration.
"""
from redis.exceptions import ConnectionError as RedisConnectionError
from redis.exceptions import TimeoutError as RedisTimeoutError
from litellm.caching.redis_cache import RedisCircuitBreaker, _is_redis_timeout_failure
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, timeout_min_duration=0.05)
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("read timed out")))
await asyncio.sleep(0.06)
for _ in range(breaker.failure_threshold - 1):
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
assert breaker.is_open() is False, "2 hard failures and 1 stale timeout are below both thresholds"
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
assert breaker.is_open() is True, "the threshold-th hard failure must still open it"
@pytest.mark.asyncio
async def test_hard_failure_resets_timeout_streak_so_a_later_burst_must_earn_its_own_duration():
"""A stale timeout followed by hard failures must not pre-age the duration gate:
a later short timeout burst has to span timeout_min_duration on its own.
"""
from redis.exceptions import ConnectionError as RedisConnectionError
from redis.exceptions import TimeoutError as RedisTimeoutError
from litellm.caching.redis_cache import RedisCircuitBreaker, _is_redis_timeout_failure
breaker = RedisCircuitBreaker(failure_threshold=3, recovery_timeout=60, timeout_min_duration=0.05)
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("read timed out")))
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
await asyncio.sleep(0.06)
for _ in range(breaker.failure_threshold):
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("read timed out")))
assert breaker.is_open() is False, "the burst is instantaneous, so the duration gate must hold it closed"
await asyncio.sleep(0.06)
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("read timed out")))
assert breaker.is_open() is True, "the same run of timeouts persisting past the duration must open it"
@pytest.mark.asyncio
async def test_breaker_metrics_track_state_and_failure_class():
"""Breaker accounting must be observable: failure class, transitions, and state."""
from prometheus_client import REGISTRY
from redis.exceptions import ConnectionError as RedisConnectionError
from redis.exceptions import TimeoutError as RedisTimeoutError
from litellm.caching.redis_cache import RedisCircuitBreaker, _is_redis_timeout_failure
def sample(name, labels=None):
return REGISTRY.get_sample_value(name, labels) or 0.0
timeout_before = sample("litellm_redis_circuit_breaker_failures_total", {"failure_class": "timeout"})
hard_before = sample("litellm_redis_circuit_breaker_failures_total", {"failure_class": "connectivity"})
opened_before = sample("litellm_redis_circuit_breaker_transitions_total", {"state": "open"})
open_gauge_before = sample("litellm_redis_circuit_breaker_state", {"state": "open"})
closed_gauge_before = sample("litellm_redis_circuit_breaker_state", {"state": "closed"})
breaker = RedisCircuitBreaker(failure_threshold=2, recovery_timeout=60, timeout_min_duration=5.0)
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisTimeoutError("t")))
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
breaker.record_failure(is_timeout=_is_redis_timeout_failure(RedisConnectionError("refused")))
assert sample("litellm_redis_circuit_breaker_failures_total", {"failure_class": "timeout"}) == timeout_before + 1
assert sample("litellm_redis_circuit_breaker_failures_total", {"failure_class": "connectivity"}) == hard_before + 2
assert sample("litellm_redis_circuit_breaker_transitions_total", {"state": "open"}) == opened_before + 1
assert sample("litellm_redis_circuit_breaker_state", {"state": "open"}) == open_gauge_before + 1
assert sample("litellm_redis_circuit_breaker_state", {"state": "closed"}) == closed_gauge_before
breaker.record_success()
assert sample("litellm_redis_circuit_breaker_state", {"state": "open"}) == open_gauge_before
assert sample("litellm_redis_circuit_breaker_state", {"state": "closed"}) == closed_gauge_before + 1

View file

@ -1014,6 +1014,8 @@ class TestAuthSchemeNormalization:
[
("Bearer abc", "Bearer", "abc"),
("bearer abc", "Bearer", "abc"),
("Bearer\tabc", "Bearer", "abc"),
("bearer\t\tabc", "Bearer", "abc"),
(" Bearer abc ", "Bearer", "abc "),
("abc", "Bearer", "abc"),
("Bearerabc", "Bearer", "Bearerabc"),

View file

@ -29,11 +29,13 @@ from litellm.types.utils import (
from litellm.litellm_core_utils.llm_cost_calc.utils import (
CostCalculatorUtils,
PromptTokensDetailsResult,
TokenRates,
TokenTypeCostBreakdown,
_calculate_input_cost,
_get_token_base_cost,
_is_off_peak,
_is_within_off_peak_window,
apply_off_peak_pricing,
calculate_cache_writing_cost,
generic_cost_per_token,
get_token_type_cost_breakdown,
@ -782,6 +784,258 @@ def test_get_token_base_cost_off_peak_wins_over_tiered_pricing():
assert outside[:2] == (3e-6, 6e-6)
def _register_off_peak_reasoning_model(
model_name: str, off_peak_pricing: dict, reasoning_rate: float | None = 4e-6, **service_tier_rates: float
) -> None:
reasoning_entry = {} if reasoning_rate is None else {"output_cost_per_reasoning_token": reasoning_rate}
litellm.register_model(
{
model_name: {
"litellm_provider": "openai",
"mode": "chat",
"input_cost_per_token": 1e-6,
"output_cost_per_token": 2e-6,
"cache_read_input_token_cost": 1e-7,
"cache_creation_input_token_cost": 1.25e-6,
"off_peak_pricing": off_peak_pricing,
**reasoning_entry,
**service_tier_rates,
}
}
)
def _off_peak_reasoning_usage() -> Usage:
return Usage(
prompt_tokens=100,
completion_tokens=80,
total_tokens=180,
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=30, text_tokens=50),
)
def test_generic_cost_per_token_off_peak_reasoning_rate():
from datetime import datetime, timezone
model_name = "litellm-test-off-peak-reasoning"
_register_off_peak_reasoning_model(
model_name,
{"hours_utc": "16:30-00:30", "output_cost_per_token": 1e-6, "output_cost_per_reasoning_token": 5e-7},
)
_, inside = generic_cost_per_token(
model=model_name,
usage=_off_peak_reasoning_usage(),
custom_llm_provider="openai",
current_time=datetime(2026, 1, 1, 18, 0, tzinfo=timezone.utc),
)
assert inside == pytest.approx(50 * 1e-6 + 30 * 5e-7)
_, outside = generic_cost_per_token(
model=model_name,
usage=_off_peak_reasoning_usage(),
custom_llm_provider="openai",
current_time=datetime(2026, 1, 1, 12, 0, tzinfo=timezone.utc),
)
assert outside == pytest.approx(50 * 2e-6 + 30 * 4e-6)
def test_generic_cost_per_token_off_peak_block_without_reasoning_rate():
from datetime import datetime, timezone
inside_window = datetime(2026, 1, 1, 18, 0, tzinfo=timezone.utc)
block = {"hours_utc": "16:30-00:30", "output_cost_per_token": 1e-6}
_register_off_peak_reasoning_model("litellm-test-off-peak-model-reasoning-rate", block)
_, with_model_rate = generic_cost_per_token(
model="litellm-test-off-peak-model-reasoning-rate",
usage=_off_peak_reasoning_usage(),
custom_llm_provider="openai",
current_time=inside_window,
)
assert with_model_rate == pytest.approx(50 * 1e-6 + 30 * 4e-6)
_register_off_peak_reasoning_model("litellm-test-off-peak-no-reasoning-rate", block, reasoning_rate=None)
_, without_model_rate = generic_cost_per_token(
model="litellm-test-off-peak-no-reasoning-rate",
usage=_off_peak_reasoning_usage(),
custom_llm_provider="openai",
current_time=inside_window,
)
assert without_model_rate == pytest.approx(80 * 1e-6)
def test_generic_cost_per_token_off_peak_reasoning_rate_wins_over_the_tier():
from datetime import datetime, timezone
model_name = "litellm-test-off-peak-tiered-reasoning"
litellm.register_model(
{
model_name: {
"litellm_provider": "openai",
"mode": "chat",
"tiered_pricing": [
{
"range": [0, 128000],
"input_cost_per_token": 3e-6,
"output_cost_per_token": 6e-6,
"output_cost_per_reasoning_token": 8e-6,
},
],
"off_peak_pricing": {
"hours_utc": "16:30-00:30",
"output_cost_per_token": 1e-6,
"output_cost_per_reasoning_token": 5e-7,
},
}
}
)
_, inside = generic_cost_per_token(
model=model_name,
usage=_off_peak_reasoning_usage(),
custom_llm_provider="openai",
current_time=datetime(2026, 1, 1, 18, 0, tzinfo=timezone.utc),
)
assert inside == pytest.approx(50 * 1e-6 + 30 * 5e-7)
_, outside = generic_cost_per_token(
model=model_name,
usage=_off_peak_reasoning_usage(),
custom_llm_provider="openai",
current_time=datetime(2026, 1, 1, 12, 0, tzinfo=timezone.utc),
)
assert outside == pytest.approx(50 * 6e-6 + 30 * 8e-6)
def test_generic_cost_per_token_off_peak_reasoning_rate_wins_over_the_service_tier():
from datetime import datetime, timezone
model_name = "litellm-test-off-peak-reasoning-service-tier"
_register_off_peak_reasoning_model(
model_name,
{"hours_utc": "16:30-00:30", "output_cost_per_token": 1e-6, "output_cost_per_reasoning_token": 5e-7},
output_cost_per_token_priority=3e-6,
output_cost_per_reasoning_token_priority=6e-6,
)
_, inside = generic_cost_per_token(
model=model_name,
usage=_off_peak_reasoning_usage(),
custom_llm_provider="openai",
service_tier="priority",
current_time=datetime(2026, 1, 1, 18, 0, tzinfo=timezone.utc),
)
assert inside == pytest.approx(50 * 1e-6 + 30 * 5e-7)
_, outside = generic_cost_per_token(
model=model_name,
usage=_off_peak_reasoning_usage(),
custom_llm_provider="openai",
service_tier="priority",
current_time=datetime(2026, 1, 1, 12, 0, tzinfo=timezone.utc),
)
assert outside == pytest.approx(50 * 3e-6 + 30 * 6e-6)
def test_apply_off_peak_pricing_treats_bool_as_unset_and_parses_strings():
from datetime import datetime, timezone
model_name = "litellm-test-off-peak-odd-values"
_register_off_peak_reasoning_model(
model_name,
{
"hours_utc": "16:30-00:30",
"cache_creation_input_token_cost": True,
"output_cost_per_reasoning_token": "5e-7",
},
)
standard = TokenRates(
input_rate=1e-6, output_rate=2e-6, cache_read_rate=1e-7, cache_creation_rate=1.25e-6, reasoning_rate=4e-6
)
rates = apply_off_peak_pricing(
litellm.get_model_info(model_name, custom_llm_provider="openai"),
datetime(2026, 1, 1, 18, 0, tzinfo=timezone.utc),
standard,
)
assert rates.cache_creation_rate == 1.25e-6
assert rates.reasoning_rate == 5e-7
def test_get_token_base_cost_off_peak_cache_creation_rate():
from datetime import datetime, timezone
from typing import cast
from litellm.types.utils import ModelInfo
model_info = cast(
ModelInfo,
{
"input_cost_per_token": 1e-6,
"output_cost_per_token": 2e-6,
"cache_creation_input_token_cost": 1.25e-6,
"cache_creation_input_token_cost_above_1hr": 2e-6,
"off_peak_pricing": {"hours_utc": "16:30-00:30", "cache_creation_input_token_cost": 5e-7},
},
)
usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150)
inside_window = datetime(2026, 1, 1, 18, 0, tzinfo=timezone.utc)
inside = _get_token_base_cost(model_info, usage, current_time=inside_window)
assert inside[2] == 5e-7
assert inside[3] == 2e-6
outside = _get_token_base_cost(model_info, usage, current_time=datetime(2026, 1, 1, 12, 0, tzinfo=timezone.utc))
assert outside[2] == 1.25e-6
without_key = cast(
ModelInfo,
{**model_info, "off_peak_pricing": {"hours_utc": "16:30-00:30", "input_cost_per_token": 5e-7}},
)
assert _get_token_base_cost(without_key, usage, current_time=inside_window)[2] == 1.25e-6
def test_get_token_type_cost_breakdown_reflects_off_peak_reasoning_and_cache_creation_rates():
from datetime import datetime, timezone
model_name = "litellm-test-off-peak-breakdown"
_register_off_peak_reasoning_model(
model_name,
{
"hours_utc": "16:30-00:30",
"output_cost_per_reasoning_token": 5e-7,
"cache_creation_input_token_cost": 5e-7,
},
)
usage = Usage(
prompt_tokens=1000,
completion_tokens=80,
total_tokens=1080,
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=30, text_tokens=50),
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=100, cache_creation_tokens=400, text_tokens=500),
)
inside = get_token_type_cost_breakdown(
model=model_name,
custom_llm_provider="openai",
usage=usage,
current_time=datetime(2026, 1, 1, 18, 0, tzinfo=timezone.utc),
)
assert inside.reasoning_cost == pytest.approx(30 * 5e-7)
assert inside.cache_creation_cost == pytest.approx(400 * 5e-7)
assert inside.cache_read_cost == pytest.approx(100 * 1e-7)
outside = get_token_type_cost_breakdown(
model=model_name,
custom_llm_provider="openai",
usage=usage,
current_time=datetime(2026, 1, 1, 12, 0, tzinfo=timezone.utc),
)
assert outside.reasoning_cost == pytest.approx(30 * 4e-6)
assert outside.cache_creation_cost == pytest.approx(400 * 1.25e-6)
def test_generic_cost_per_token_gpt54_above_272k_tokens(_local_model_cost_map):
"""GPT-5.4/5.4-pro: prompts >272K input tokens priced at 2x input, 1.5x output."""
model = "gpt-5.4"

View file

@ -649,6 +649,90 @@ class TestDashscopeCostCalculator:
assert math.isclose(completion_cost, 200 * 2.4e-06, rel_tol=1e-10)
def test_dashscope_off_peak_reasoning_rate_replaces_the_dedicated_reasoning_rate(self):
self._register_off_peak_flat_model(
"dashscope/qwen-reasoning-rate-off-peak-test",
{
"hours_utc": self.OFF_PEAK_WINDOW,
"output_cost_per_token": 2.4e-06,
"output_cost_per_reasoning_token": 4.5e-06,
},
)
litellm.model_cost["dashscope/qwen-reasoning-rate-off-peak-test"]["output_cost_per_reasoning_token"] = 9e-06
usage = Usage(
prompt_tokens=100,
completion_tokens=200,
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=50),
)
_, completion_cost = dashscope_cost_per_token(
model="qwen-reasoning-rate-off-peak-test", usage=usage, current_time=self.INSIDE_WINDOW
)
assert math.isclose(completion_cost, (150 * 2.4e-06) + (50 * 4.5e-06), rel_tol=1e-10)
_, peak_completion_cost = dashscope_cost_per_token(
model="qwen-reasoning-rate-off-peak-test", usage=usage, current_time=self.OUTSIDE_WINDOW
)
assert math.isclose(peak_completion_cost, (150 * 4.8e-06) + (50 * 9e-06), rel_tol=1e-10)
def test_dashscope_off_peak_cache_creation_rate_replaces_the_standard_rate(self):
self._register_off_peak_flat_model(
"dashscope/qwen-cache-creation-off-peak-test",
{"hours_utc": self.OFF_PEAK_WINDOW, "cache_creation_input_token_cost": 1.5e-06},
)
usage = Usage(
prompt_tokens=1000,
completion_tokens=10,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=300, cache_creation_tokens=100),
)
prompt_cost, _ = dashscope_cost_per_token(
model="qwen-cache-creation-off-peak-test", usage=usage, current_time=self.INSIDE_WINDOW
)
assert math.isclose(prompt_cost, (600 * 2.4e-06) + (300 * 2e-07) + (100 * 1.5e-06), rel_tol=1e-10)
peak_prompt_cost, _ = dashscope_cost_per_token(
model="qwen-cache-creation-off-peak-test", usage=usage, current_time=self.OUTSIDE_WINDOW
)
assert math.isclose(peak_prompt_cost, (600 * 2.4e-06) + (300 * 2e-07) + (100 * 3e-06), rel_tol=1e-10)
def test_dashscope_off_peak_reasoning_and_cache_creation_rates_override_the_selected_tier(self):
self._register_tiered_model(
"dashscope/qwen-tiered-reasoning-off-peak-test",
[
{
"range": [0, 1000],
"input_cost_per_token": 4e-07,
"cache_creation_input_token_cost": 3e-07,
"output_cost_per_token": 1.6e-06,
"output_cost_per_reasoning_token": 3.2e-06,
},
],
)
litellm.model_cost["dashscope/qwen-tiered-reasoning-off-peak-test"]["off_peak_pricing"] = {
"hours_utc": self.OFF_PEAK_WINDOW,
"cache_creation_input_token_cost": 1e-07,
"output_cost_per_reasoning_token": 8e-07,
}
usage = Usage(
prompt_tokens=500,
completion_tokens=100,
prompt_tokens_details=PromptTokensDetailsWrapper(cache_creation_tokens=200),
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=40),
)
prompt_cost, completion_cost = dashscope_cost_per_token(
model="qwen-tiered-reasoning-off-peak-test", usage=usage, current_time=self.INSIDE_WINDOW
)
assert math.isclose(prompt_cost, (300 * 4e-07) + (200 * 1e-07), rel_tol=1e-10)
assert math.isclose(completion_cost, (60 * 1.6e-06) + (40 * 8e-07), rel_tol=1e-10)
peak_prompt_cost, peak_completion_cost = dashscope_cost_per_token(
model="qwen-tiered-reasoning-off-peak-test", usage=usage, current_time=self.OUTSIDE_WINDOW
)
assert math.isclose(peak_prompt_cost, (300 * 4e-07) + (200 * 3e-07), rel_tol=1e-10)
assert math.isclose(peak_completion_cost, (60 * 1.6e-06) + (40 * 3.2e-06), rel_tol=1e-10)
def test_dashscope_off_peak_defaults_to_the_current_time(self):
"""The proxy's cost dispatch passes no clock, so an all-day window has to apply on the
default current time."""

View file

@ -85,6 +85,31 @@ class TestResolveConfig:
def test_plaintext_http_api_base_disables(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
assert resolve_openai_workload_identity_config(api_key=None, api_base="http://api.openai.com/v1") is None
@pytest.mark.parametrize(
"api_base",
(
"https://southcentralus.privatelink.api.openai.com/v1",
"https://eu.api.openai.com/v1",
"https://us.api.openai.com/v1",
),
)
def test_openai_backed_api_base_allows(self, wif_env: OpenAIWorkloadIdentityConfig, api_base: str) -> None:
assert resolve_openai_workload_identity_config(api_key=None, api_base=api_base) == wif_env
@pytest.mark.parametrize(
"api_base",
(
"https://api.openai.com.evil.example/v1",
"https://openai.com/v1",
"https://euapi.openai.com/v1",
"http://southcentralus.privatelink.api.openai.com/v1",
),
)
def test_lookalike_or_plaintext_api_base_disables(
self, wif_env: OpenAIWorkloadIdentityConfig, api_base: str
) -> None:
assert resolve_openai_workload_identity_config(api_key=None, api_base=api_base) is None
def test_foreign_env_base_url_disables(
self, wif_env: OpenAIWorkloadIdentityConfig, monkeypatch: pytest.MonkeyPatch
) -> None:
@ -158,6 +183,14 @@ class TestClientConstruction:
assert client.api_key == "workload-identity-auth"
assert client._workload_identity_auth is not None
def test_privatelink_client_uses_workload_identity(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
client: Final = OpenAIChatCompletion()._get_openai_client(
is_async=False, api_key=None, api_base="https://southcentralus.privatelink.api.openai.com/v1"
)
assert isinstance(client, OpenAI)
assert client.api_key == "workload-identity-auth"
assert client._workload_identity_auth is not None
def test_static_key_client_unaffected(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
client: Final = OpenAIChatCompletion()._get_openai_client(is_async=False, api_key="sk-static", api_base=None)
assert isinstance(client, OpenAI)
@ -231,6 +264,16 @@ class TestResponsesValidateEnvironment:
)
assert headers["Authorization"] == "Bearer None"
@respx.mock
def test_privatelink_api_base_mints_bearer(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
mock_token_exchange()
headers: Final = OpenAIResponsesAPIConfig().validate_environment(
headers={},
model="gpt-4o-mini",
litellm_params=GenericLiteLLMParams(api_base="https://southcentralus.privatelink.api.openai.com/v1"),
)
assert headers["Authorization"] == "Bearer exchanged-bearer-token"
def test_litellm_proxy_subclass_never_mints_wif(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
headers: Final = LiteLLMProxyResponsesAPIConfig().validate_environment(
headers={}, model="gpt-4o-mini", litellm_params=GenericLiteLLMParams()

View file

@ -12,6 +12,7 @@ import pytest
from fastapi import HTTPException
from pydantic import ValidationError
from litellm.experimental_mcp_client.client import MCPClient
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
oauth_protected_resource_path,
raise_public,
@ -94,6 +95,49 @@ def test_basic_scheme_base64_encodes_the_token():
assert spec.config.key_source.value.get_secret_value() == expected
@pytest.mark.parametrize(
"auth_type, authentication_token, expected_value, expected_header",
[
(MCPAuth.bearer_token, "Bearer abc", "abc", ("Authorization", "Bearer abc")),
(MCPAuth.token, "token abc", "abc", ("Authorization", "token abc")),
(MCPAuth.basic, "user:pass", "dXNlcjpwYXNz", ("Authorization", "Basic dXNlcjpwYXNz")),
(MCPAuth.basic, "Basic dXNlcjpwYXNz", "dXNlcjpwYXNz", ("Authorization", "Basic dXNlcjpwYXNz")),
(MCPAuth.basic, "Basic user:pass", "dXNlcjpwYXNz", ("Authorization", "Basic dXNlcjpwYXNz")),
],
)
def test_shared_key_normalizes_schemed_authentication_token(
auth_type, authentication_token, expected_value, expected_header
):
spec = to_server_spec(_server(auth_type=auth_type, authentication_token=authentication_token))
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
assert spec.config.key_source.value.get_secret_value() == expected_value
assert spec.config.header(expected_value) == expected_header
@pytest.mark.parametrize(
"auth_type, authentication_token",
[
(MCPAuth.bearer_token, "Bearer abc"),
(MCPAuth.bearer_token, "abc"),
(MCPAuth.token, "token abc"),
(MCPAuth.token, "abc"),
(MCPAuth.basic, "user:pass"),
(MCPAuth.basic, "Basic dXNlcjpwYXNz"),
(MCPAuth.basic, "Basic user:pass"),
],
)
def test_shared_key_authorization_matches_v1(auth_type, authentication_token):
spec = to_server_spec(_server(auth_type=auth_type, authentication_token=authentication_token))
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
client = MCPClient(server_url="https://x", auth_type=auth_type)
client.update_auth_value(authentication_token)
assert spec.config.header(spec.config.key_source.value.get_secret_value())[1] == client._get_auth_headers()[
"Authorization"
]
@pytest.mark.parametrize(
"oauth2_flow",
[None, "authorization_code"],

View file

@ -2616,6 +2616,69 @@ class TestMCPServerManager:
await manager.preflight_token_exchange(server=server, oauth2_headers=None, user_api_key_auth=None)
assert resolved == ["good-subject"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"authorization",
[
"Bearer subj-jwt",
"bearer subj-jwt",
"BEARER subj-jwt",
"Bearer\tsubj-jwt",
"Bearer subj-jwt",
],
)
async def test_preflight_token_exchange_strips_inbound_authorization_scheme(self, authorization):
"""The resolver posts inbound_token verbatim as subject_token."""
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError, ServerSpec, Subject
resolved: Final[list[str | None]] = []
class _FakeProvider:
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]:
resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None)
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization"))
manager = MCPServerManager(cred_provider=_FakeProvider())
server = self._token_exchange_server(f"te-preflight-auth-{authorization!r}")
await manager.preflight_token_exchange(
server=server,
oauth2_headers={"Authorization": authorization},
user_api_key_auth=None,
)
assert resolved == ["subj-jwt"]
@pytest.mark.asyncio
async def test_preflight_token_exchange_preserves_authorization_without_separator(self):
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
StaticHeaderAuth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError, ServerSpec, Subject
resolved: Final[list[str | None]] = []
class _FakeProvider:
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]:
resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None)
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization"))
manager = MCPServerManager(cred_provider=_FakeProvider())
server = self._token_exchange_server("te-preflight-auth-no-separator")
await manager.preflight_token_exchange(
server=server,
oauth2_headers={"Authorization": "Bearersubj-jwt"},
user_api_key_auth=None,
)
assert resolved == ["Bearersubj-jwt"]
@pytest.mark.asyncio
async def test_preflight_token_exchange_skips_discovery_for_other_auth_modes(self):
"""Preflight must not make unrelated auth modes depend on OAuth discovery."""
@ -4571,6 +4634,81 @@ class TestMCPServerManager:
# auth_type is none here, so a 401 from this upstream must not be dressed up as a re-auth signal
assert captured["relays_upstream_auth"] is False
@pytest.mark.asyncio
@pytest.mark.parametrize(
"auth_type, authentication_token, expected_authorization",
[
(MCPAuth.bearer_token, "Bearer abc", "Bearer abc"),
(MCPAuth.bearer_token, "abc", "Bearer abc"),
(MCPAuth.api_key, "ApiKey abc", "ApiKey abc"),
(MCPAuth.token, "token abc", "token abc"),
(MCPAuth.basic, "user:pass", "Basic dXNlcjpwYXNz"),
(MCPAuth.basic, "Basic dXNlcjpwYXNz", "Basic dXNlcjpwYXNz"),
(MCPAuth.basic, "Basic user:pass", "Basic dXNlcjpwYXNz"),
],
)
async def test_register_openapi_tools_normalizes_authentication_token(
self, tmp_path, monkeypatch, auth_type, authentication_token, expected_authorization
):
manager = MCPServerManager()
spec_path = tmp_path / "openapi.json"
spec_path.write_text(
json.dumps(
{
"openapi": "3.0.0",
"info": {"title": "Demo", "version": "1.0.0"},
"paths": {
"/health": {
"get": {
"operationId": "health_check",
"summary": "health",
}
}
},
}
)
)
server = MCPServer(
server_id="openapi-server",
name="openapi-server",
server_name="openapi-server",
url="https://example.com",
transport=MCPTransport.http,
auth_type=auth_type,
authentication_token=authentication_token,
)
captured: dict = {}
def fake_create_tool_function(
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False
):
captured["headers"] = headers
async def tool_func(**kwargs):
return "ok"
return tool_func
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.create_tool_function",
fake_create_tool_function,
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.build_input_schema",
lambda *args, **kwargs: {"type": "object", "properties": {}, "required": []},
)
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.tool_registry.global_mcp_tool_registry.register_tool",
lambda *args, **kwargs: None,
)
await manager._register_openapi_tools(
spec_path=str(spec_path),
server=server,
base_url="https://example.com",
)
assert captured["headers"]["Authorization"] == expected_authorization
@pytest.mark.asyncio
async def test_pre_call_tool_check_allowed_tools_list_allows_tool(self):
"""Test pre_call_tool_check allows tool when it's in allowed_tools list"""

View file

@ -1349,6 +1349,38 @@ async def test_vector_store_access_check_early_returns(
assert result == expected_result
@pytest.mark.asyncio
async def test_vector_store_access_check_skips_db_lookup_when_no_vector_stores_requested():
"""Registry returns [] (not None) for plain requests; no object permission DB lookup should happen."""
valid_token = UserAPIKeyAuth(token="test-token", object_permission_id="perm-123")
team_object = MagicMock()
team_object.object_permission_id = "team-permission"
mock_prisma_client = MagicMock()
find_unique = AsyncMock()
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = find_unique
mock_vector_store_registry = MagicMock()
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = []
with (
patch( # test-quality-ok: production auth reads these module globals; no dependency injection seam exists
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
),
patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists
"litellm.vector_store_registry", mock_vector_store_registry
),
):
result = await vector_store_access_check(
request_body={"messages": [{"role": "user", "content": "test"}]},
team_object=team_object,
valid_token=valid_token,
)
assert result is True
find_unique.assert_not_awaited()
@pytest.mark.parametrize(
"object_permissions,vector_store_ids,should_raise,error_type",
[

View file

@ -19,11 +19,13 @@ from litellm.proxy._types import (
NewUserResponse,
ProxyErrorTypes,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.management_endpoints.scim.scim_v2 import (
SCIMRosterSyncError,
UserProvisionerHelpers,
_apply_group_patch_updates,
_create_user_if_not_exists,
_extract_group_member_ids,
_extract_ids_from_path_filter,
_handle_group_membership_changes,
@ -37,8 +39,8 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import (
delete_group,
delete_user,
get_groups,
get_users,
get_service_provider_config,
get_users,
merge_placeholder,
patch_group,
patch_team_membership,
@ -304,6 +306,85 @@ async def test_create_user_uses_default_internal_user_params_role(mocker, monkey
assert called_args.user_role == LitellmUserRoles.PROXY_ADMIN
def _mock_scim_create_user_deps(mocker: MockerFixture, scim_user: SCIMUser) -> AsyncMock:
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_usertable.find_many = AsyncMock(return_value=())
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable into create_user
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)
mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable into create_user
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
AsyncMock(return_value=scim_user),
)
return mocker.patch( # test-quality-ok: endpoint collaborators are module-level, not injectable into create_user
"litellm.proxy.management_endpoints.scim.scim_v2.new_user",
AsyncMock(return_value=NewUserRequest(user_id=scim_user.userName)),
)
@pytest.mark.asyncio
async def test_create_user_without_groups_defers_to_default_team(mocker: MockerFixture, monkeypatch):
"""IdPs omit groups on POST /Users; teams must stay unset so new_user applies default_internal_user_params.teams"""
scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
userName="new-user",
emails=[SCIMUserEmail(value="new@example.com")],
)
monkeypatch.setattr(
"litellm.default_internal_user_params",
{"teams": [{"team_id": "default-team", "max_budget_in_team": 25}]},
raising=False,
)
new_user_mock = _mock_scim_create_user_deps(mocker, scim_user)
await create_user(user=scim_user)
assert new_user_mock.call_args.kwargs["data"].teams is None
assert new_user_mock.call_args.kwargs["user_api_key_dict"] == UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
@pytest.mark.asyncio
async def test_create_user_with_groups_keeps_idp_teams(mocker: MockerFixture, monkeypatch):
scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
userName="new-user",
emails=[SCIMUserEmail(value="new@example.com")],
groups=[SCIMUserGroup(value="idp-team")],
)
monkeypatch.setattr(
"litellm.default_internal_user_params",
{"teams": [{"team_id": "default-team", "max_budget_in_team": 25}]},
raising=False,
)
new_user_mock = _mock_scim_create_user_deps(mocker, scim_user)
await create_user(user=scim_user)
assert new_user_mock.call_args.kwargs["data"].teams == ["idp-team"]
@pytest.mark.asyncio
async def test_create_user_if_not_exists_defers_to_default_team(mocker: MockerFixture, monkeypatch):
monkeypatch.setattr(
"litellm.default_internal_user_params",
{"teams": [{"team_id": "default-team", "max_budget_in_team": 25}]},
raising=False,
)
new_user_mock = mocker.patch( # test-quality-ok: new_user is imported inside the helper, not injectable
"litellm.proxy.management_endpoints.internal_user_endpoints.new_user",
AsyncMock(return_value=NewUserResponse(user_id="group-user", key="k")),
)
created = await _create_user_if_not_exists(user_id="group-user")
assert created is not None
assert new_user_mock.call_args.kwargs["data"].teams is None
assert new_user_mock.call_args.kwargs["user_api_key_dict"] == UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
@pytest.mark.asyncio
async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeypatch):
"""
@ -1176,6 +1257,67 @@ async def test_update_user_success(mocker):
assert call_args[1]["data"]["teams"] == ["new-team"]
@pytest.mark.asyncio
@pytest.mark.parametrize("groups", [None, []], ids=["groups-omitted", "groups-empty"])
async def test_update_user_without_groups_preserves_memberships_and_role(mocker, monkeypatch, groups):
"""Okta profile PUTs carry no `groups` or `groups: []`; neither may drop teams (and their keys) or recompute role"""
from litellm.proxy.proxy_server import proxy_config
async def mock_get_config():
return {"litellm_settings": {"scim_admin_group": "litellm-admins"}}
monkeypatch.setattr(proxy_config, "get_config", mock_get_config)
monkeypatch.setattr("litellm.default_internal_user_params", None, raising=False)
existing_user = mocker.MagicMock()
existing_user.teams = ["litellm-admins", "engineering"]
existing_user.metadata = {}
scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
userName="okta-user",
name=SCIMUserName(familyName="Renamed", givenName="Okta"),
emails=[SCIMUserEmail(value="okta@example.com")],
**({} if groups is None else {"groups": groups}),
)
response_scim_user = SCIMUser(
schemas=["urn:ietf:params:scim:schemas:core:2.0:User"],
id="okta-user",
userName="okta-user",
emails=[SCIMUserEmail(value="okta@example.com")],
)
mock_prisma_client = mocker.MagicMock()
mock_prisma_client.db = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={"user_id": "okta-user"})
mocker.patch( # test-quality-ok: update_user's collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
AsyncMock(return_value=mock_prisma_client),
)
mocker.patch( # test-quality-ok: update_user's collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists",
AsyncMock(return_value=existing_user),
)
mocker.patch( # test-quality-ok: update_user's collaborators are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user",
AsyncMock(return_value=response_scim_user),
)
patch_membership = mocker.patch( # test-quality-ok: roster writes are module-level, not injectable
"litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership",
AsyncMock(),
)
result = await update_user(user_id="okta-user", user=scim_user)
assert result == response_scim_user
patch_membership.assert_not_awaited()
update_data = mock_prisma_client.db.litellm_usertable.update.call_args.kwargs["data"]
assert update_data["teams"] == ["litellm-admins", "engineering"]
assert "user_role" not in update_data
@pytest.mark.asyncio
async def test_update_user_not_found(mocker):
"""Should raise 404 when user doesn't exist"""

View file

@ -12,13 +12,12 @@ from fastapi.testclient import TestClient
from prisma.errors import PrismaError
import litellm.proxy.proxy_server as ps
from litellm.proxy.proxy_server import app
from litellm.proxy._types import (
CommonProxyErrors,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.proxy_server import app
def _make_access_group_record(
@ -58,6 +57,10 @@ def _make_access_group_record(
return record
def _make_team_record(team_id: str, access_group_ids: list[str] | None = None):
return types.SimpleNamespace(team_id=team_id, access_group_ids=access_group_ids or [])
@pytest.fixture
def client_and_mocks(monkeypatch):
"""Setup mock prisma and admin auth for access group endpoints."""
@ -185,7 +188,8 @@ ACCESS_GROUP_PATHS = ["/v1/access_group", "/v1/unified_access_group"]
)
def test_create_access_group_success(client_and_mocks, base_path, payload):
"""Create access group with various payloads returns 201."""
client, _, mock_table, *_ = client_and_mocks
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[_make_team_record("team-1")])
resp = client.post(base_path, json=payload)
assert resp.status_code == 201
@ -277,13 +281,45 @@ def test_create_access_group_500_on_non_constraint_prisma_error(client_and_mocks
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_list_access_groups_success_empty(client_and_mocks, base_path):
"""List access groups returns empty list when none exist."""
client, _, mock_table, *_ = client_and_mocks
"""List access groups returns empty list when none exist, without querying teams."""
client, mock_prisma, mock_table, *_ = client_and_mocks
resp = client.get(base_path)
assert resp.status_code == 200
assert resp.json() == []
mock_table.find_many.assert_awaited_once()
mock_prisma.db.litellm_teamtable.find_many.assert_not_awaited()
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_list_access_groups_attributes_teams_per_group_with_one_query(client_and_mocks, base_path):
"""List derives each group's teams from the team table in a single query, attributed per group."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
records = [
_make_access_group_record(access_group_id="ag-1", access_group_name="group-1"),
_make_access_group_record(access_group_id="ag-2", access_group_name="group-2"),
]
mock_table.find_many = AsyncMock(return_value=records)
mock_team_table.find_many = AsyncMock(
return_value=[
_make_team_record("team-x", ["ag-1"]),
_make_team_record("team-y", ["ag-2"]),
_make_team_record("team-z", ["ag-1", "ag-2"]),
]
)
resp = client.get(base_path)
assert resp.status_code == 200
body = resp.json()
assert body[0]["assigned_team_ids"] == ["team-x", "team-z"]
assert body[1]["assigned_team_ids"] == ["team-y", "team-z"]
mock_team_table.find_many.assert_awaited_once()
carrying, listed = mock_team_table.find_many.call_args.kwargs["where"]["OR"]
assert list(carrying["access_group_ids"]["hasSome"]) == ["ag-1", "ag-2"]
assert list(listed["team_id"]["in"]) == []
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
@ -373,6 +409,43 @@ def test_get_access_group_success(client_and_mocks, base_path, access_group_id):
assert resp.json()["access_group_id"] == access_group_id
@pytest.mark.parametrize("base_path", ACCESS_GROUP_PATHS)
def test_get_access_group_derives_assigned_teams_from_team_table(client_and_mocks, base_path):
"""Get drops ghost ids from the stored column and adds teams that carry the group but were never mirrored."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
record = _make_access_group_record(access_group_id="ag-123", assigned_team_ids=["team-a", "ghost-team"])
mock_table.find_unique = AsyncMock(return_value=record)
mock_team_table.find_many = AsyncMock(
return_value=[
_make_team_record("team-a", ["ag-123"]),
_make_team_record("team-b", ["ag-123"]),
_make_team_record("team-c", ["ag-123"]),
]
)
resp = client.get(f"{base_path}/ag-123")
assert resp.status_code == 200
assert resp.json()["assigned_team_ids"] == ["team-a", "team-b", "team-c"]
carrying, listed = mock_team_table.find_many.call_args.kwargs["where"]["OR"]
assert list(carrying["access_group_ids"]["hasSome"]) == ["ag-123"]
assert list(listed["team_id"]["in"]) == ["team-a", "ghost-team"]
def test_get_access_group_empty_column_and_no_teams_returns_empty(client_and_mocks):
"""Get returns [] when the column is empty and no team carries the group."""
client, mock_prisma, mock_table, *_ = client_and_mocks
mock_table.find_unique = AsyncMock(return_value=_make_access_group_record(access_group_id="ag-123"))
mock_prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
resp = client.get("/v1/access_group/ag-123")
assert resp.status_code == 200
assert resp.json()["assigned_team_ids"] == []
def test_get_access_group_not_found(client_and_mocks):
"""Get access group returns 404 when not found."""
client, _, mock_table, *_ = client_and_mocks
@ -985,6 +1058,28 @@ def test_record_to_access_group_table():
assert result.access_agent_ids == ["agent-1"]
def test_attached_team_ids_by_group_keeps_column_order_then_appends_unmirrored_teams():
"""Stored ids that resolve keep their order, ghosts drop, carriers the mirror missed append once, per group."""
from litellm.proxy.management_endpoints.access_group_endpoints import (
_attached_team_ids_by_group,
)
records = [
_make_access_group_record(access_group_id="ag-1", assigned_team_ids=["team-b", "ghost", "team-a"]),
_make_access_group_record(access_group_id="ag-2", assigned_team_ids=[]),
]
teams = [
_make_team_record("team-a", ["ag-1"]),
_make_team_record("team-b", []),
_make_team_record("team-c", ["ag-1"]),
_make_team_record("team-d", ["ag-2"]),
]
result = _attached_team_ids_by_group(records, teams)
assert dict(result) == {"ag-1": ("team-b", "team-a", "team-c"), "ag-2": ("team-d",)}
# ---------------------------------------------------------------------------
# Sync tests: CREATE
# ---------------------------------------------------------------------------
@ -997,9 +1092,8 @@ def test_create_access_group_syncs_assigned_teams(client_and_mocks):
)
mock_team_table = mock_prisma.db.litellm_teamtable
team_record = MagicMock()
team_record.team_id = "team-1"
team_record.access_group_ids = []
team_record = _make_team_record("team-1")
mock_team_table.find_many = AsyncMock(return_value=[team_record])
mock_team_table.find_unique = AsyncMock(return_value=team_record)
resp = client.post(
@ -1043,20 +1137,22 @@ def test_create_access_group_syncs_assigned_keys(client_and_mocks):
assert "ag-new" in call_kwargs["data"]["access_group_ids"]
def test_create_access_group_skips_sync_for_nonexistent_team(client_and_mocks):
"""Create skips updating a team that doesn't exist in DB."""
client, mock_prisma, _, mock_cache, mock_proxy_logging = client_and_mocks
def test_create_access_group_rejects_nonexistent_team(client_and_mocks):
"""Create refuses to store a team id that does not resolve to a team row."""
client, mock_prisma, mock_access_group_table, *_ = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
mock_team_table.find_unique = AsyncMock(return_value=None)
mock_team_table.find_many = AsyncMock(return_value=[_make_team_record("team-real")])
resp = client.post(
"/v1/access_group",
json={
"access_group_name": "new-group",
"assigned_team_ids": ["nonexistent-team"],
"assigned_team_ids": ["team-real", "nonexistent-team", "also-missing"],
},
)
assert resp.status_code == 201
assert resp.status_code == 400
assert resp.json()["detail"] == "Unknown team ids: also-missing, nonexistent-team"
mock_access_group_table.create.assert_not_awaited()
mock_team_table.update.assert_not_awaited()
@ -1065,9 +1161,8 @@ def test_create_access_group_idempotent_team_sync(client_and_mocks):
client, mock_prisma, _, mock_cache, mock_proxy_logging = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
team_record = MagicMock()
team_record.team_id = "team-1"
team_record.access_group_ids = ["ag-new"] # already synced
team_record = _make_team_record("team-1", ["ag-new"])
mock_team_table.find_many = AsyncMock(return_value=[team_record])
mock_team_table.find_unique = AsyncMock(return_value=team_record)
resp = client.post(
@ -1095,9 +1190,8 @@ def test_update_access_group_syncs_added_teams(client_and_mocks):
)
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
team_record = MagicMock()
team_record.team_id = "team-new"
team_record.access_group_ids = []
team_record = _make_team_record("team-new")
mock_team_table.find_many = AsyncMock(return_value=[_make_team_record("team-existing", ["ag-update"]), team_record])
mock_team_table.find_unique = AsyncMock(return_value=team_record)
resp = client.put(
@ -1113,6 +1207,25 @@ def test_update_access_group_syncs_added_teams(client_and_mocks):
assert "ag-update" in call_kwargs["data"]["access_group_ids"]
def test_update_access_group_rejects_nonexistent_team(client_and_mocks):
"""Update refuses to store a team id that does not resolve to a team row and leaves the group untouched."""
client, mock_prisma, mock_access_group_table, *_ = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
existing = _make_access_group_record(access_group_id="ag-update", assigned_team_ids=["team-existing"])
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
mock_team_table.find_many = AsyncMock(return_value=[_make_team_record("team-existing", ["ag-update"])])
resp = client.put(
"/v1/access_group/ag-update",
json={"assigned_team_ids": ["team-existing", "team-ghost"]},
)
assert resp.status_code == 400
assert resp.json()["detail"] == "Unknown team ids: team-ghost"
mock_access_group_table.update.assert_not_awaited()
mock_team_table.update.assert_not_awaited()
def test_update_access_group_syncs_removed_teams(client_and_mocks):
"""Update removes access_group_id from de-assigned teams."""
client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = (
@ -1125,9 +1238,8 @@ def test_update_access_group_syncs_removed_teams(client_and_mocks):
)
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
team_to_remove = MagicMock()
team_to_remove.team_id = "team-remove"
team_to_remove.access_group_ids = ["ag-update"]
team_to_remove = _make_team_record("team-remove", ["ag-update"])
mock_team_table.find_many = AsyncMock(return_value=[_make_team_record("team-keep", ["ag-update"]), team_to_remove])
mock_team_table.find_unique = AsyncMock(return_value=team_to_remove)
resp = client.put(
@ -1145,6 +1257,28 @@ def test_update_access_group_syncs_removed_teams(client_and_mocks):
assert "ag-update" not in call_kwargs["data"]["access_group_ids"]
def test_update_access_group_detaches_team_the_mirror_missed(client_and_mocks):
"""Update removes the group from a team that carries it but was never written to the stored column."""
client, mock_prisma, mock_access_group_table, *_ = client_and_mocks
mock_team_table = mock_prisma.db.litellm_teamtable
existing = _make_access_group_record(access_group_id="ag-update", assigned_team_ids=["team-keep"])
mock_access_group_table.find_unique = AsyncMock(return_value=existing)
unmirrored = _make_team_record("team-unmirrored", ["ag-update", "ag-other"])
mock_team_table.find_many = AsyncMock(return_value=[_make_team_record("team-keep", ["ag-update"]), unmirrored])
mock_team_table.find_unique = AsyncMock(return_value=unmirrored)
resp = client.put("/v1/access_group/ag-update", json={"assigned_team_ids": ["team-keep"]})
assert resp.status_code == 200
mock_team_table.find_unique.assert_awaited_once_with(where={"team_id": "team-unmirrored"})
mock_team_table.update.assert_awaited_once()
call_kwargs = mock_team_table.update.call_args.kwargs
assert call_kwargs["where"] == {"team_id": "team-unmirrored"}
assert call_kwargs["data"]["access_group_ids"] == ["ag-other"]
def test_update_access_group_no_team_sync_when_ids_not_in_payload(client_and_mocks):
"""Update does not sync teams when assigned_team_ids is absent from the payload."""
client, mock_prisma, mock_access_group_table, mock_cache, mock_proxy_logging = (

View file

@ -6347,6 +6347,93 @@ def test_build_key_filter_conditions_key_hash_narrows_team_admin_visibility():
assert {"token": "hashed-token-123"} in where["AND"], f"key_hash not ANDed: {where}"
def _search_clause(search: str, token: str) -> dict:
return {"OR": [{"token": token}, {"key_alias": {"contains": search, "mode": "insensitive"}}]}
def test_build_key_filter_conditions_search_ors_token_and_alias_contains():
"""
LIT-4741: `search` matches a key by its alias (case-insensitive contains) OR by
its ID (the token column), with the pasted value used verbatim.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_key_filter_conditions,
)
hashed_where = json.loads(
json.dumps(
_build_key_filter_conditions(
user_id=None,
team_id=None,
organization_id=None,
key_alias=None,
key_hash=None,
exclude_team_id=None,
admin_team_ids=None,
search="already-hashed-token",
)
)
)
assert _search_clause("already-hashed-token", "already-hashed-token") in hashed_where["AND"], (
f"hashed search not used verbatim: {hashed_where}"
)
def test_build_key_filter_conditions_search_narrows_team_admin_visibility():
"""
LIT-4741, same class as LIT-3243: `search` must be a top-level AND so it
narrows a team admin's admin-team branch instead of being bypassed by it.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
_build_key_filter_conditions,
)
where = json.loads(
json.dumps(
_build_key_filter_conditions(
user_id="team-admin-user",
team_id=None,
organization_id=None,
key_alias=None,
key_hash=None,
exclude_team_id=None,
admin_team_ids=["team-a"],
member_team_ids=["team-a"],
include_created_by_keys=False,
search="member-key-id",
)
)
)
assert where.get("AND"), f"expected top-level AND, got: {where}"
assert _search_clause("member-key-id", "member-key-id") in where["AND"], f"search not ANDed: {where}"
assert json.dumps({"team_id": {"in": ["team-a"]}}) in json.dumps(where)
@pytest.mark.asyncio
async def test_list_key_helper_applies_search_to_prisma_where():
"""LIT-4741: `search` given to _list_key_helper must reach the Prisma where clause."""
mock_prisma_client = AsyncMock()
mock_find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many
mock_prisma_client.db.litellm_verificationtoken.count = AsyncMock(return_value=0)
await _list_key_helper(
prisma_client=mock_prisma_client,
page=1,
size=50,
user_id=None,
team_id=None,
organization_id=None,
key_alias=None,
key_hash=None,
search="key-id-123",
)
where = json.loads(json.dumps(mock_find_many.call_args.kwargs["where"]))
assert _search_clause("key-id-123", "key-id-123") in where["AND"], f"search not in Prisma where: {where}"
@pytest.mark.asyncio
async def test_generate_key_negative_max_budget():
"""
@ -14870,6 +14957,16 @@ async def test_list_keys_non_admin_cannot_opt_into_substring():
assert kwargs["user_id"] == "alice"
@pytest.mark.asyncio
async def test_list_keys_search_is_honored_for_non_admin():
"""LIT-4741: unlike substring_matching, `search` is not admin-gated. A non-admin's
search reaches the helper while their own-user scoping stays in place."""
user = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="alice")
kwargs = await _list_keys_capture_helper_kwargs(user, user_id=None, search="key-id-123")
assert kwargs["search"] == "key-id-123"
assert kwargs["user_id"] == "alice"
@pytest.mark.asyncio
async def test_cli_session_token_delegation_ceiling_blocked_by_team_budget():
team = LiteLLM_TeamTableCachedObj(team_id="team-1", max_budget=50.0)

View file

@ -615,6 +615,96 @@ class TestMemoryEndpoints:
assert keys == {"user:profile"}
assert body["total"] == 1
def test_list_memory_search_matches_key_prefix_or_memory_id_within_scope(self):
"""
`search` matches a key prefix OR an exact memory_id, and stays ANDed
with the visibility filter so a pasted foreign id cannot leak a row.
"""
table = self.prisma.db.litellm_memorytable
table.rows.extend(
[
_make_row(memory_id="mem-own", key="user:profile", user_id="user-a", team_id=None),
_make_row(memory_id="mem-target", key="project:context", user_id="user-a", team_id=None),
_make_row(memory_id="mem-foreign", key="user:secret", user_id="user-b", team_id=None),
]
)
client = _make_client(_user_auth("user-a", "team-a"))
with _patch_prisma(self.prisma):
by_id = client.get("/v1/memory?search=mem-target")
by_prefix = client.get("/v1/memory?search=user:")
foreign_id = client.get("/v1/memory?search=mem-foreign")
assert by_id.status_code == 200, by_id.text
assert [m["memory_id"] for m in by_id.json()["memories"]] == ["mem-target"]
assert by_id.json()["total"] == 1
assert by_prefix.status_code == 200, by_prefix.text
assert {m["key"] for m in by_prefix.json()["memories"]} == {"user:profile"}
assert by_prefix.json()["total"] == 1
assert foreign_id.status_code == 200, foreign_id.text
assert foreign_id.json()["memories"] == []
assert foreign_id.json()["total"] == 0
def test_list_memory_search_by_memory_id_for_admin_sees_any_scope(self):
"""Admins have no visibility filter, so an id search returns the row whoever owns it."""
table = self.prisma.db.litellm_memorytable
table.rows.extend(
[
_make_row(memory_id="mem-a", key="a", user_id="user-a", team_id=None),
_make_row(memory_id="mem-b", key="b", user_id="user-b", team_id=None),
]
)
client = _make_client(_admin_auth())
with _patch_prisma(self.prisma):
resp = client.get("/v1/memory?search=mem-b")
assert resp.status_code == 200, resp.text
assert [m["memory_id"] for m in resp.json()["memories"]] == ["mem-b"]
assert resp.json()["total"] == 1
def test_list_memory_search_wins_over_key_prefix(self):
"""When both are sent, `search` decides the match and `key_prefix` is ignored."""
table = self.prisma.db.litellm_memorytable
table.rows.extend(
[
_make_row(memory_id="mem-own", key="user:profile", user_id="user-a", team_id=None),
_make_row(memory_id="mem-target", key="project:context", user_id="user-a", team_id=None),
]
)
client = _make_client(_user_auth("user-a", "team-a"))
with _patch_prisma(self.prisma):
resp = client.get("/v1/memory?search=mem-target&key_prefix=user:")
assert resp.status_code == 200, resp.text
assert [m["memory_id"] for m in resp.json()["memories"]] == ["mem-target"]
assert resp.json()["total"] == 1
def test_list_memory_key_prefix_never_matches_memory_id(self):
"""`key_prefix` stays a pure key-prefix match; only `search` consults memory_id."""
table = self.prisma.db.litellm_memorytable
table.rows.append(_make_row(memory_id="mem-target", key="project:context", user_id="user-a", team_id=None))
client = _make_client(_user_auth("user-a", "team-a"))
with _patch_prisma(self.prisma):
resp = client.get("/v1/memory?key_prefix=mem-target")
assert resp.status_code == 200, resp.text
assert resp.json()["memories"] == []
assert resp.json()["total"] == 0
def test_list_memory_key_exact_filter(self):
"""`key` is an exact match, never a prefix."""
table = self.prisma.db.litellm_memorytable
table.rows.extend(
[
_make_row(memory_id="m1", key="user:profile", user_id="user-a", team_id=None),
_make_row(memory_id="m2", key="user:profile:archived", user_id="user-a", team_id=None),
]
)
client = _make_client(_user_auth("user-a", "team-a"))
with _patch_prisma(self.prisma):
resp = client.get("/v1/memory?key=user:profile")
assert resp.status_code == 200, resp.text
assert [m["memory_id"] for m in resp.json()["memories"]] == ["m1"]
assert resp.json()["total"] == 1
def test_list_memory_admin_sees_all(self):
table = self.prisma.db.litellm_memorytable
table.rows.extend(

View file

@ -58,6 +58,24 @@ def _filter_logs_by_date_range(logs, where):
return filtered
_SEARCH_CLAUSE_RE = re.compile(
r'\(request_id = \$(\d+) OR \("startTime" >= \(\$(\d+)::timestamptz AT TIME ZONE \'UTC\'\) '
r'AND "startTime" <= \(\$(\d+)::timestamptz AT TIME ZONE \'UTC\'\) '
r'AND \(api_key = \$\1 OR team_id = \$\1 OR "user" = \$\1 OR end_user = \$\1 '
r"OR session_id = \$\1 OR model_id = \$\1\)\)\)"
)
def _matches_spend_log_search(log, search):
"""Mirror the search clause: request_id across all time, the other id columns inside the window."""
if log.get("request_id") == search["value"]:
return True
if not _filter_logs_by_date_range([log], {"startTime": {"gte": search["gte"], "lte": search["lte"]}}):
return False
columns = ("api_key", "team_id", "user", "end_user", "session_id", "model_id")
return any(log.get(col) == search["value"] for col in columns)
def _reconstruct_ui_where_from_sql(sql_query, params):
"""
Rebuild the Prisma-style ``where`` dict the filter_fns below expect from the
@ -77,6 +95,16 @@ def _reconstruct_ui_where_from_sql(sql_query, params):
def _iso(value):
return value.isoformat() if hasattr(value, "isoformat") else str(value)
search_clause = _SEARCH_CLAUSE_RE.search(clause.group(1))
if search_clause:
raw_index, start_index, end_index = (int(g) for g in search_clause.groups())
where["search"] = {
"value": params[raw_index - 1],
"gte": _iso(params[start_index - 1]),
"lte": _iso(params[end_index - 1]),
}
remaining = clause.group(1) if search_clause is None else clause.group(1).replace(search_clause.group(0), "")
eq_cols = {
"team_id": "team_id",
'"user"': "user",
@ -89,7 +117,7 @@ def _reconstruct_ui_where_from_sql(sql_query, params):
}
date_bounds: dict = {}
metadata_conds: list = []
for cond in (c.strip() for c in clause.group(1).split(" AND ")):
for cond in (c.strip() for c in remaining.split(" AND ")):
gte = re.search(r'"startTime" >= \(\$(\d+)', cond)
lte = re.search(r'"startTime" <= \(\$(\d+)', cond)
alias = re.search(r"user_api_key_alias' LIKE \$(\d+)", cond)
@ -2352,6 +2380,208 @@ async def test_ui_view_spend_logs_request_id_owner_scoped_by_id_only(
app.dependency_overrides.pop(ps.user_api_key_auth, None)
def test_build_spend_log_search_condition_windows_every_branch_except_request_id():
"""LIT-4741: request_id matches across all time; the six other id columns only inside the window,
all comparing the pasted value verbatim."""
start = datetime.datetime(2026, 8, 1, tzinfo=timezone.utc)
end = datetime.datetime(2026, 8, 2, tzinfo=timezone.utc)
condition = spend_management_endpoints._build_spend_log_search_condition(
search="key-hash-7", start_date=start, end_date=end, next_param_index=3
)
assert condition.sql == (
"(request_id = $3 OR (\"startTime\" >= ($4::timestamptz AT TIME ZONE 'UTC') "
"AND \"startTime\" <= ($5::timestamptz AT TIME ZONE 'UTC') "
'AND (api_key = $3 OR team_id = $3 OR "user" = $3 OR end_user = $3 OR session_id = $3 OR model_id = $3)))'
)
assert condition.params == ("key-hash-7", start, end)
def _search_fixture_logs(today):
recent = (today - datetime.timedelta(days=1)).isoformat()
old = (today - datetime.timedelta(days=90)).isoformat()
base = {
"api_key": "hashed-other",
"user": "user-x",
"team_id": "team-x",
"end_user": "cust-x",
"session_id": "sess-x",
"model_id": "mdl-x",
"spend": 0.01,
"model": "gpt-4",
}
return [
{**base, "request_id": "req-session", "session_id": "sess-42", "startTime": recent},
{**base, "request_id": "req-session-old", "session_id": "sess-42", "startTime": old},
{**base, "request_id": "req-key", "api_key": "hashed-7", "startTime": recent},
{**base, "request_id": "req-team", "team_id": "team-7", "startTime": recent},
{**base, "request_id": "req-user", "user": "user-7", "startTime": recent},
{**base, "request_id": "req-end-user", "end_user": "cust-7", "startTime": recent},
{**base, "request_id": "req-model", "model_id": "mdl-7", "startTime": recent},
]
def _search_filter_fn(logs, captured):
def filter_fn(where):
captured["where"] = where
rows = _filter_logs_by_date_range(logs, where)
if "user" in where:
rows = [row for row in rows if row["user"] == where["user"]]
if "search" in where:
rows = [row for row in rows if _matches_spend_log_search(row, where["search"])]
return rows
return filter_fn
def _five_day_window(today):
return {
"start_date": (today - datetime.timedelta(days=5)).strftime("%Y-%m-%d %H:%M:%S"),
"end_date": today.strftime("%Y-%m-%d %H:%M:%S"),
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"search,expected_request_ids",
[
("req-session-old", {"req-session-old"}),
("sess-42", {"req-session"}),
("hashed-7", {"req-key"}),
("team-7", {"req-team"}),
("user-7", {"req-user"}),
("cust-7", {"req-end-user"}),
("mdl-7", {"req-model"}),
("no-such-id", set()),
],
)
async def test_ui_view_spend_logs_search_matches_any_id(client, monkeypatch, search, expected_request_ids):
"""LIT-4741: one box matches any id column. A request_id is found across all time (the 5-day
window excludes the 90-day-old row), every other column only inside the window, and a raw
sk- key is hashed before it is compared with api_key. The window is not applied globally."""
today = datetime.datetime.now(timezone.utc)
logs = _search_fixture_logs(today)
captured = {}
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
make_ui_spend_logs_mock_prisma(logs, _search_filter_fn(logs, captured)),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
try:
response = client.get(
"/spend/logs/ui",
params={"search": search, **_five_day_window(today)},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200, response.text
data = response.json()
assert {row["request_id"] for row in data["data"]} == expected_request_ids
assert data["total"] == len(expected_request_ids)
assert "startTime" not in captured["where"]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_spend_logs_v2_search_keeps_global_window(client, monkeypatch):
"""The public route keeps the caller's window on the whole query, so a search only finds rows
inside it even by request_id; the windowless request_id branch is a dashboard-only relaxation."""
today = datetime.datetime.now(timezone.utc)
logs = _search_fixture_logs(today)
captured = {}
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
make_ui_spend_logs_mock_prisma(logs, _search_filter_fn(logs, captured)),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
try:
response = client.get(
"/spend/logs/v2",
params={"search": "req-session-old", **_five_day_window(today)},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200, response.text
data = response.json()
assert data["data"] == []
assert data["total"] == 0
assert "startTime" in captured["where"]
assert captured["where"]["search"]["value"] == "req-session-old"
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"params",
[
{"search": "req-old"},
{"search": "req-old", "request_id": "req-old"},
],
)
async def test_ui_view_spend_logs_search_requires_dates(client, monkeypatch, params):
"""A search needs the window for its non-request_id branches, so it stays required even
alongside a request_id, which on its own may drop the window."""
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
make_ui_spend_logs_mock_prisma([], lambda where: []),
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
try:
response = client.get("/spend/logs/ui", params=params, headers={"Authorization": "Bearer sk-test"})
assert response.status_code == 400
assert "date" in response.text.lower()
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"search,expected_request_ids",
[("sess-9", {"req-own"}), ("req-foreign", set())],
)
async def test_ui_view_spend_logs_search_keeps_non_admin_scope(client, monkeypatch, search, expected_request_ids):
"""A search is scoped like any other listing: an internal user only sees their own rows even
when the id is on someone else's row, and the request_id ownership shortcut is not used."""
yesterday = (datetime.datetime.now(timezone.utc) - datetime.timedelta(days=1)).isoformat()
base = {"api_key": "hashed-key", "team_id": None, "spend": 0.01, "startTime": yesterday, "model": "gpt-4"}
logs = [
{**base, "request_id": "req-own", "user": "internal_user_1", "session_id": "sess-9"},
{**base, "request_id": "req-own-other", "user": "internal_user_1", "session_id": "sess-other"},
{**base, "request_id": "req-foreign", "user": "internal_user_2", "session_id": "sess-9"},
]
captured = {}
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client",
make_ui_spend_logs_mock_prisma(logs, _search_filter_fn(logs, captured)),
)
monkeypatch.setattr(
"litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs",
AsyncMock(return_value=[]),
)
ownership_check = AsyncMock()
monkeypatch.setattr(
"litellm.proxy.spend_tracking.spend_management_endpoints._assert_user_can_view_request_id",
ownership_check,
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER, user_id="internal_user_1"
)
try:
start_date, end_date = _default_date_range()
response = client.get(
"/spend/logs/ui",
params={"search": search, "start_date": start_date, "end_date": end_date},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200, response.text
assert {row["request_id"] for row in response.json()["data"]} == expected_request_ids
assert captured["where"]["user"] == "internal_user_1"
ownership_check.assert_not_awaited()
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_spend_logs_unauthorized(client):
# Test without authorization header
@ -6351,3 +6581,46 @@ async def test_ui_view_spend_logs_group_by_session_offset_for_non_starttime_sort
assert "OFFSET" in emitted_sql[1]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_spend_logs_search_returns_flat_rows_when_grouping_by_session(client, monkeypatch):
"""The dashboard lists sessions by default; a search for an id lists every matching row instead,
so both calls of a session show up rather than one representative, and no session cursor is returned."""
rows = [_session_representative_row("req-1", "sess-1"), _session_representative_row("req-2", "sess-1")]
async def mock_query_raw(sql_query, *params):
if "mcp_tool_call_count" in sql_query:
return []
grouped = "DISTINCT ON" in sql_query or "GROUP BY" in sql_query
visible = rows[:1] if grouped else rows
if "COUNT(*)" in sql_query:
return [{"total_count": len(visible)}]
return visible
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = AsyncMock(side_effect=mock_query_raw)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user"
)
try:
start_date, end_date = _default_date_range()
response = client.get(
"/spend/logs/ui",
params={
"search": "sess-1",
"group_by_session": "true",
"start_date": start_date,
"end_date": end_date,
},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200, response.text
data = response.json()
assert [row["request_id"] for row in data["data"]] == ["req-1", "req-2"]
assert data["total"] == 2
assert "next_session_cursor" not in data
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)

View file

@ -184,6 +184,7 @@ async def test_spend_logs_ui_wraps_params_in_at_time_zone_utc(monkeypatch):
api_key=None,
user_id=None,
request_id=None,
search=None,
start_date="2026-02-16 00:00:00",
end_date="2026-02-16 23:59:59",
page=1,
@ -247,6 +248,7 @@ async def test_spend_logs_ui_uses_bounded_count_not_full_scan(monkeypatch):
api_key=None,
user_id=None,
request_id=None,
search=None,
start_date="2026-02-16 00:00:00",
end_date="2026-02-16 23:59:59",
page=1,
@ -314,6 +316,7 @@ async def test_spend_logs_ui_caps_total_for_large_result_sets(monkeypatch):
api_key=None,
user_id=None,
request_id=None,
search=None,
start_date="2026-02-16 00:00:00",
end_date="2026-02-16 23:59:59",
page=1,
@ -359,6 +362,7 @@ async def test_spend_logs_ui_empty_page_reports_zero_total(monkeypatch):
api_key=None,
user_id=None,
request_id=None,
search=None,
start_date="2026-02-16 00:00:00",
end_date="2026-02-16 23:59:59",
page=1,
@ -406,6 +410,7 @@ async def test_spend_logs_ui_out_of_range_page_keeps_total(monkeypatch):
api_key=None,
user_id=None,
request_id=None,
search=None,
start_date="2026-02-16 00:00:00",
end_date="2026-02-16 23:59:59",
page=99,
@ -552,6 +557,7 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch):
api_key=None,
user_id=None,
request_id=None,
search=None,
start_date="2026-02-16 00:00:00",
end_date="2026-02-16 23:59:59",
page=1,
@ -616,6 +622,7 @@ async def test_spend_logs_ui_group_by_session_offset_pages_for_other_sorts(monke
api_key=None,
user_id=None,
request_id=None,
search=None,
start_date="2026-02-16 00:00:00",
end_date="2026-02-16 23:59:59",
page=2,
@ -664,6 +671,7 @@ async def test_spend_logs_ui_request_id_lookup_with_grouping_returns_exact_row(m
api_key=None,
user_id=None,
request_id="req-deep-link",
search=None,
start_date=None,
end_date=None,
page=1,

View file

@ -7,7 +7,6 @@ import re
import stat
import subprocess
from pathlib import Path
from typing import Optional
import pytest
@ -16,11 +15,15 @@ COMPONENT_ENTRYPOINT = REPO_ROOT / "docker" / "component_entrypoint.sh"
PROD_ENTRYPOINT = REPO_ROOT / "docker" / "prod_entrypoint.sh"
GATEWAY_DOCKERFILE = REPO_ROOT / "gateway" / "Dockerfile"
BACKEND_DOCKERFILE = REPO_ROOT / "backend" / "Dockerfile"
BUILD_FROM_PIP_DOCKERFILE = REPO_ROOT / "docker" / "build_from_pip" / "Dockerfile.build_from_pip"
TERRAFORM_ECS = REPO_ROOT / "terraform" / "litellm" / "aws" / "ecs.tf"
TERRAFORM_CLOUDRUN = REPO_ROOT / "terraform" / "litellm" / "gcp" / "cloudrun.tf"
IMAGE_ENTRYPOINT_PATH = "/app/docker/component_entrypoint.sh"
TRUTHY_USE_DDTRACE = ("true", "True", "TRUE", "tRuE")
FALSY_USE_DDTRACE = (None, "", "false", "False", "1", "yes", "on", "truex")
PYTHONPATH_SENTINEL = "/lit-entrypoint-sentinel:/app"
_STUB_TEMPLATE = """#!/bin/sh
@ -34,6 +37,7 @@ _STUB_TEMPLATE = """#!/bin/sh
_ENTRYPOINT_RE = re.compile(r"^ENTRYPOINT\s+(\[.*\])\s*$", re.MULTILINE)
_CMD_RE = re.compile(r"^CMD\s+(\[.*\])\s*$", re.MULTILINE)
_COPY_RE = re.compile(r"^COPY\s+(?!--from)(\S+)\s+(\S+)\s*$", re.MULTILINE)
_APP_TARGET_RE = re.compile(r"(?:gateway|backend)\.main:app")
_TF_STRING_LOCAL_RE = re.compile(r'^\s*(\w+)\s*=\s*"((?:[^"\\]|\\.)*)"\s*$', re.MULTILINE)
_TF_INTERPOLATION_RE = re.compile(r"\$\{(local|var)\.(\w+)\}")
@ -53,7 +57,7 @@ def _write_stubs(bin_dir: Path, names: tuple[str, ...]) -> None:
def _run_entrypoint(
script: Path,
argv: tuple[str, ...],
use_ddtrace: Optional[str],
use_ddtrace: str | None,
tmp_path: Path,
) -> tuple[str, ...]:
"""Run `script` with stubbed executables on PATH and return the recorded lines."""
@ -84,7 +88,7 @@ def _run_entrypoint(
return tuple(record.read_text().splitlines()) if record.exists() else ()
def _run_shell_command(command: str, bin_dir: Path, record: Path, use_ddtrace: Optional[str]) -> tuple[str, ...]:
def _run_shell_command(command: str, bin_dir: Path, record: Path, use_ddtrace: str | None) -> tuple[str, ...]:
"""Run a resolved Terraform launch command through `sh -c` and return the recorded lines."""
env = {
**os.environ,
@ -184,9 +188,19 @@ def test_ddtrace_disabled_execs_the_command_directly(tmp_path: Path) -> None:
)
@pytest.mark.parametrize("use_ddtrace", [None, "", "false", "True", "TRUE", "1", "yes"])
def test_gating_matches_the_monolithic_entrypoint(use_ddtrace: Optional[str], tmp_path: Path) -> None:
"""The componentized images must honor `USE_DDTRACE` exactly as the monolith does."""
@pytest.mark.parametrize(
"use_ddtrace, traced",
[*((v, True) for v in TRUTHY_USE_DDTRACE), *((v, False) for v in FALSY_USE_DDTRACE)],
)
def test_gating_matches_the_monolithic_entrypoint_and_get_secret_bool(
use_ddtrace: str | None, traced: bool, tmp_path: Path
) -> None:
"""Both entrypoints must accept exactly the spellings `get_secret_bool` accepts.
`ProxyStartupEvent._init_dd_tracer` reads `USE_DDTRACE` through `get_secret_bool`, which
matches `true` case-insensitively. If the shell gate were stricter, `USE_DDTRACE=True` would
give in-process LLM spans without `ddtrace-run` HTTP spans, a half-enabled state.
"""
component = _run_entrypoint(
COMPONENT_ENTRYPOINT,
("uvicorn", "gateway.main:app"),
@ -200,8 +214,56 @@ def test_gating_matches_the_monolithic_entrypoint(use_ddtrace: Optional[str], tm
tmp_path=tmp_path / "monolith",
)
assert component[0].startswith("exec=") and monolith[0].startswith("exec=")
assert (component[0] == "exec=ddtrace-run") == (monolith[0] == "exec=ddtrace-run")
expected_exec = "exec=ddtrace-run" if traced else "exec=uvicorn"
expected_openai = "DD_TRACE_OPENAI_ENABLED=False" if traced else "DD_TRACE_OPENAI_ENABLED=<unset>"
assert component[0] == expected_exec
assert component[2] == expected_openai
assert monolith[0] == ("exec=ddtrace-run" if traced else "exec=litellm")
assert monolith[2] == expected_openai
assert monolith[1] == ("args=litellm --port 4000" if traced else "args=--port 4000")
def _copied_script(dockerfile: Path, image_path: str) -> Path:
"""Resolve the repo file a Dockerfile `COPY`s to `image_path`, so tests run what the image ships."""
matches = _COPY_RE.findall(dockerfile.read_text())
sources = tuple(src for src, dst in matches if dst == image_path)
assert sources, f"{dockerfile} never COPYs anything to {image_path}"
source = REPO_ROOT / sources[-1]
assert source.is_file(), f"{dockerfile} COPYs {sources[-1]}, which does not exist in the build context"
return source
@pytest.mark.parametrize(
"use_ddtrace, traced",
[*((v, True) for v in TRUTHY_USE_DDTRACE), *((v, False) for v in FALSY_USE_DDTRACE)],
)
def test_build_from_pip_image_launches_litellm_through_the_prod_entrypoint(
use_ddtrace: str | None, traced: bool, tmp_path: Path
) -> None:
"""Run the build_from_pip image's ENTRYPOINT + CMD through the script it actually COPYs.
The image used to `ENTRYPOINT ["litellm"]`, so `USE_DDTRACE` was inert there at any spelling.
Resolving the ENTRYPOINT path back to its COPY source and executing it with the Dockerfile's
CMD checks the launch the container performs, not just that the Dockerfile mentions the script.
"""
entrypoint = _entrypoint_argv(BUILD_FROM_PIP_DOCKERFILE)
assert len(entrypoint) == 1, (
f"{BUILD_FROM_PIP_DOCKERFILE} ENTRYPOINT must be the bare script so CMD reaches litellm"
)
script = _copied_script(BUILD_FROM_PIP_DOCKERFILE, entrypoint[0])
assert script == PROD_ENTRYPOINT, f"{BUILD_FROM_PIP_DOCKERFILE} bypasses the ddtrace-aware entrypoint"
assert f"chmod +x {entrypoint[0]}" in BUILD_FROM_PIP_DOCKERFILE.read_text()
cmd = _cmd_argv(BUILD_FROM_PIP_DOCKERFILE)
recorded = _run_entrypoint(script, cmd, use_ddtrace=use_ddtrace, tmp_path=tmp_path)
cmd_str = " ".join(cmd)
assert recorded == (
"exec=ddtrace-run" if traced else "exec=litellm",
f"args=litellm {cmd_str}" if traced else f"args={cmd_str}",
"DD_TRACE_OPENAI_ENABLED=False" if traced else "DD_TRACE_OPENAI_ENABLED=<unset>",
f"PYTHONPATH={PYTHONPATH_SENTINEL}",
)
def test_entrypoint_script_is_executable() -> None:
@ -239,9 +301,9 @@ def test_component_images_make_the_entrypoint_executable(dockerfile: Path) -> No
@pytest.mark.parametrize("terraform_file", TERRAFORM_LAUNCH_SITES, ids=lambda p: p.parent.name)
@pytest.mark.parametrize("component", ["gateway", "backend"])
@pytest.mark.parametrize("use_ddtrace", [None, "true", "false", "True"])
@pytest.mark.parametrize("use_ddtrace", [*TRUTHY_USE_DDTRACE, *FALSY_USE_DDTRACE])
def test_terraform_launch_command_matches_the_script_contract(
terraform_file: Path, component: str, use_ddtrace: Optional[str], tmp_path: Path
terraform_file: Path, component: str, use_ddtrace: str | None, tmp_path: Path
) -> None:
"""The Terraform command and `docker/component_entrypoint.sh` must decide identically.
@ -273,7 +335,7 @@ def test_terraform_launch_command_matches_the_script_contract(
assert from_terraform[2] == from_script[2], f"{terraform_file} disagrees with the script on the openai integration"
assert app_target in from_terraform[1]
if use_ddtrace == "true":
if use_ddtrace in TRUTHY_USE_DDTRACE:
assert from_terraform[0] == "exec=ddtrace-run"
assert from_terraform[2] == "DD_TRACE_OPENAI_ENABLED=False"
else:

View file

@ -3855,7 +3855,7 @@ def test_pre_call_checks_counts_responses_instructions_tokens(monkeypatch):
input_only_tokens = router._count_pre_call_check_tokens(messages=None, input=short_input)
with_instructions_tokens = router._count_pre_call_check_tokens(
messages=None, input=short_input, instructions=long_instructions
messages=None, input=short_input, request_kwargs={"instructions": long_instructions}
)
assert with_instructions_tokens > input_only_tokens
@ -3871,6 +3871,164 @@ def test_pre_call_checks_counts_responses_instructions_tokens(monkeypatch):
)
_OVERSIZED_TOOL_DESCRIPTION = "look up the answer in the knowledge base. " * 40
@pytest.mark.parametrize(
"prompt_kwargs, tool",
[
pytest.param(
{"messages": [{"role": "user", "content": "hi"}]},
{
"type": "function",
"function": {
"name": "lookup",
"description": _OVERSIZED_TOOL_DESCRIPTION,
"parameters": {"type": "object", "properties": {"q": {"type": "string"}}},
},
},
id="chat_completions_tool",
),
pytest.param(
{"input": "hi"},
{
"type": "function",
"name": "lookup",
"description": _OVERSIZED_TOOL_DESCRIPTION,
"parameters": {"type": "object", "properties": {"q": {"type": "string"}}},
},
id="responses_tool",
),
pytest.param(
{"messages": [{"role": "user", "content": "hi"}]},
{
"name": "lookup",
"description": _OVERSIZED_TOOL_DESCRIPTION,
"input_schema": {"type": "object", "properties": {"q": {"type": "string"}}},
},
id="anthropic_messages_tool",
),
],
)
def test_pre_call_checks_counts_tool_definition_tokens(monkeypatch, prompt_kwargs, tool):
"""
Tool definitions are sent to the model as prompt tokens but never appear in
`messages` or `input`. A request whose prompt alone fits the context window but
whose prompt plus `tools` exceeds it must be rejected before dispatch, for the
Chat Completions, Responses and Anthropic Messages tool shapes alike.
"""
router = litellm.Router(
model_list=[
{"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
],
enable_pre_call_checks=True,
)
deployments = [
{"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
]
prompt_only_tokens = router._count_pre_call_check_tokens(
messages=prompt_kwargs.get("messages"), input=prompt_kwargs.get("input")
)
monkeypatch.setattr(
router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": prompt_only_tokens}
)
assert len(router._pre_call_checks(model="m", healthy_deployments=deployments, **prompt_kwargs)) == 1
with pytest.raises(litellm.ContextWindowExceededError):
router._pre_call_checks(
model="m",
healthy_deployments=deployments,
request_kwargs={"tools": [tool]},
**prompt_kwargs,
)
@pytest.mark.parametrize(
"system",
[
pytest.param("You are a meticulous assistant. " * 40, id="system_string"),
pytest.param(
[{"type": "text", "text": "You are a meticulous assistant. " * 40}],
id="system_blocks",
),
],
)
def test_pre_call_checks_counts_anthropic_system_tokens(monkeypatch, system):
"""
The Anthropic Messages API carries the system prompt as a top-level `system` field,
not as a message. Its tokens reach the model, so a request whose `messages` fit but
whose `messages` plus `system` exceed the context window must be rejected.
"""
router = litellm.Router(
model_list=[
{"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
],
enable_pre_call_checks=True,
)
deployments = [
{"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
]
messages = [{"role": "user", "content": "hi"}]
messages_only_tokens = router._count_pre_call_check_tokens(messages=messages, input=None)
monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": messages_only_tokens})
assert len(router._pre_call_checks(model="m", healthy_deployments=deployments, messages=messages)) == 1
with pytest.raises(litellm.ContextWindowExceededError):
router._pre_call_checks(
model="m",
healthy_deployments=deployments,
messages=messages,
request_kwargs={"system": system},
)
@pytest.mark.asyncio
async def test_aanthropic_messages_enforces_context_window_with_system_and_tools():
"""
End-to-end router regression for /v1/messages: a request whose only oversized
content lives in the top-level `system` field or in `tools` must trip the pre-call
context-window check instead of being dispatched (the deployment uses mock_response,
so reaching the provider handler would return a response rather than raise).
"""
router = litellm.Router(
model_list=[
{
"model_name": "small-ctx",
"litellm_params": {"model": "anthropic/claude-3-5-haiku-20241022", "mock_response": "hi"},
"model_info": {"max_input_tokens": 20},
}
],
enable_pre_call_checks=True,
)
messages = [{"role": "user", "content": "hi"}]
response = await router.aanthropic_messages(model="small-ctx", messages=messages, max_tokens=5)
assert response is not None
with pytest.raises(litellm.ContextWindowExceededError):
await router.aanthropic_messages(
model="small-ctx",
messages=messages,
max_tokens=5,
system="You are a meticulous assistant. " * 40,
)
with pytest.raises(litellm.ContextWindowExceededError):
await router.aanthropic_messages(
model="small-ctx",
messages=messages,
max_tokens=5,
tools=[
{
"name": "lookup",
"description": _OVERSIZED_TOOL_DESCRIPTION,
"input_schema": {"type": "object", "properties": {"q": {"type": "string"}}},
}
],
)
def test_count_pre_call_check_tokens_across_api_surfaces():
"""
_count_pre_call_check_tokens must count tokens from chat `messages`, a Responses
@ -7362,14 +7520,14 @@ def test_get_configured_mode_reads_deployment_model_info():
router = litellm.Router(
model_list=[
{
"model_name": "chat-model",
"litellm_params": {"model": "openai/some-unmapped-model"},
"model_info": {"mode": "chat"},
"model_name": "tts-model",
"litellm_params": {"model": "openai/some-unmapped-tts-model"},
"model_info": {"mode": "audio_speech"},
}
]
)
assert router.get_configured_mode("chat-model") == "chat"
assert router.get_configured_mode("tts-model") == "audio_speech"
def test_get_configured_mode_returns_none_for_unset_or_unknown():
@ -12701,33 +12859,3 @@ async def test_prompt_management_factory_marks_injection_for_every_deployment(mo
bucket = captured.get("litellm_metadata") or captured["metadata"]
assert captured["model_info"]["id"] == "provisional-dep"
assert bucket["litellm_gateway_injected_cache"] == ""
def test_get_configured_mode_reads_deployment_model_info():
router = Router(
model_list=[
{
"model_name": "my-tts",
"litellm_params": {"model": "openai/some-unmapped-mode-model"},
"model_info": {"mode": "audio_speech"},
}
]
)
assert router.get_configured_mode("my-tts") == "audio_speech"
@pytest.mark.parametrize("model_info", [{}, {"mode": ""}, {"mode": " "}, {"mode": 123}])
def test_get_configured_mode_returns_none_for_unset_blank_or_unknown(model_info):
router = Router(
model_list=[
{
"model_name": "plain-model",
"litellm_params": {"model": "openai/some-unmapped-mode-model"},
"model_info": model_info,
}
]
)
assert router.get_configured_mode("plain-model") is None
assert router.get_configured_mode("unknown-model") is None

View file

@ -3,7 +3,7 @@
"limit": 22328
},
"LIT002": {
"limit": 26760
"limit": 26758
},
"LIT003": {
"limit": 261

View file

@ -90,7 +90,7 @@ describe("AgentsTable", () => {
/>,
);
const search = screen.getByPlaceholderText("Search agent names or descriptions...");
const search = screen.getByPlaceholderText("Search agents by name, ID, or description...");
await user.type(search, "billing");
expect(screen.getByText("Billing Router")).toBeInTheDocument();
expect(screen.queryByText("Second Agent")).not.toBeInTheDocument();
@ -101,11 +101,36 @@ describe("AgentsTable", () => {
expect(screen.queryByText("Billing Router")).not.toBeInTheDocument();
});
it("filters agents by a pasted agent_id so only that agent's row survives", async () => {
const user = userEvent.setup();
render(
<AgentsTable
agents={[
makeAgent({ agent_id: "5f3c2a1b-9d8e-4f7a-b6c5-d4e3f2a1b0c9", agent_name: "Billing Router" }),
makeAgent({ agent_id: "0a9b8c7d-6e5f-4a3b-8c2d-1e0f9a8b7c6d", agent_name: "Second Agent" }),
]}
{...baseProps}
/>,
);
const search = screen.getByPlaceholderText("Search agents by name, ID, or description...");
await user.click(search);
await user.paste("5f3c2a1b-9d8e-4f7a-b6c5-d4e3f2a1b0c9");
expect(screen.getByText("Billing Router")).toBeInTheDocument();
expect(screen.queryByText("Second Agent")).not.toBeInTheDocument();
await user.clear(search);
await user.paste("ffffffff-0000-4000-8000-000000000000");
expect(screen.queryByText("Billing Router")).not.toBeInTheDocument();
expect(screen.queryByText("Second Agent")).not.toBeInTheDocument();
expect(screen.getByText("No matching agents")).toBeInTheDocument();
});
it("shows the no-match empty state when the search matches nothing", async () => {
const user = userEvent.setup();
render(<AgentsTable agents={[makeAgent()]} {...baseProps} />);
await user.type(screen.getByPlaceholderText("Search agent names or descriptions..."), "zzzz");
await user.type(screen.getByPlaceholderText("Search agents by name, ID, or description..."), "zzzz");
expect(screen.queryByText("Test Agent")).not.toBeInTheDocument();
expect(screen.getByText("No matching agents")).toBeInTheDocument();
});

View file

@ -55,7 +55,12 @@ const AgentsTable: React.FC<AgentsTableProps> = ({
const [sorting, setSorting] = useState<SortingState>(DEFAULT_SORTING);
const [searchTerm, setSearchTerm] = useState("");
const filteredAgents = useMemo(
() => filterBySearchTerm(agents, searchTerm, (agent) => [agent.agent_name, agent.agent_card_params?.description]),
() =>
filterBySearchTerm(agents, searchTerm, (agent) => [
agent.agent_name,
agent.agent_id,
agent.agent_card_params?.description,
]),
[agents, searchTerm],
);
@ -83,7 +88,7 @@ const AgentsTable: React.FC<AgentsTableProps> = ({
<SearchIcon className="size-4 text-muted-foreground" />
</InputGroupAddon>
<InputGroupInput
placeholder="Search agent names or descriptions..."
placeholder="Search agents by name, ID, or description..."
value={searchTerm}
onChange={(e) => setSearchTerm(e.target.value)}
/>

View file

@ -518,6 +518,24 @@ describe("useKeys", () => {
const callUrl = mockFetch.mock.calls[0][0];
expect(callUrl).not.toContain("agent_id");
});
it("sends the combined alias-or-ID search as the search param, separate from key_alias and key_hash", async () => {
mockFetch.mockResolvedValueOnce({
ok: true,
json: async () => mockKeysResponse,
});
const { result } = renderHook(() => useKeys(1, 10, { search: "pasted-key-id" }), { wrapper });
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
});
const callUrl = new URL(mockFetch.mock.calls[0][0], "http://localhost");
expect(callUrl.searchParams.get("search")).toBe("pasted-key-id");
expect(callUrl.searchParams.has("key_alias")).toBe(false);
expect(callUrl.searchParams.has("key_hash")).toBe(false);
});
});
describe("useDeletedKeys", () => {

View file

@ -40,6 +40,7 @@ export interface KeyListCallOptions {
selectedKeyAlias?: string | null;
userID?: string | null;
keyHash?: string | null;
search?: string | null;
sortBy?: string | null;
sortOrder?: string | null;
expand?: string | null;
@ -61,6 +62,7 @@ const keyListCall = async (accessToken: string, page: number, pageSize: number,
organization_id: options.organizationID,
key_alias: options.selectedKeyAlias,
key_hash: options.keyHash,
search: options.search,
user_id: options.userID,
page,
size: pageSize,

View file

@ -95,6 +95,7 @@ describe("MemoryTable", () => {
it("shows the filtered-empty copy when a search is active", () => {
render(<MemoryTable {...baseProps} data={[]} rowCount={0} hasActiveSearch={true} />);
expect(screen.getByText("No matching memories")).toBeInTheDocument();
expect(screen.getByText("No memories match your search.")).toBeInTheDocument();
expect(screen.queryByText("No memories stored yet")).not.toBeInTheDocument();
});
@ -128,6 +129,7 @@ describe("MemoryTable", () => {
const onRefresh = vi.fn();
render(<MemoryTable {...baseProps} onSearchChange={onSearchChange} onRefresh={onRefresh} />);
expect(screen.getByPlaceholderText("Search by key prefix or memory ID…")).toBeInTheDocument();
fireEvent.change(screen.getByTestId("datatable-search"), { target: { value: "u" } });
expect(onSearchChange).toHaveBeenCalledWith("u");

View file

@ -36,7 +36,7 @@ function MemoryEmptyState({ hasActiveSearch }: { hasActiveSearch: boolean }) {
</div>
<div className="text-sm text-muted-foreground">
{hasActiveSearch
? "No memories have keys starting with your search."
? "No memories match your search."
: "Memories your agents store under /v1/memory will appear here."}
</div>
</div>
@ -81,7 +81,7 @@ export function MemoryTable({
table={table}
searchValue={searchValue}
onSearchChange={onSearchChange}
searchPlaceholder='Filter by key prefix, e.g. "user:"'
searchPlaceholder="Search by key prefix or memory ID…"
onRefresh={onRefresh}
isRefreshing={isRefreshing}
showViewOptions={false}

View file

@ -1,8 +1,9 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { act, render, screen } from "@testing-library/react";
import type { PaginationState } from "@tanstack/react-table";
import { act, render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import React from "react";
import { describe, expect, it, vi } from "vitest";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { MemoryRow } from "@/components/networking";
@ -13,10 +14,13 @@ interface CapturedTableProps {
rowCount: number;
data: MemoryRow[];
hasActiveSearch: boolean;
onSearchChange: (value: string) => void;
onPaginationChange: (state: PaginationState) => void;
onViewClick: (row: MemoryRow) => void;
}
const captured = vi.hoisted(() => ({ current: null as CapturedTableProps | null }));
const fetchMemoryListMock = vi.hoisted(() => vi.fn());
vi.mock("./MemoryTable", () => ({
MemoryTable: function MemoryTableMock(props: CapturedTableProps) {
@ -25,6 +29,15 @@ vi.mock("./MemoryTable", () => ({
},
}));
vi.mock("@/components/networking", async (importOriginal) => ({
...(await importOriginal<typeof import("@/components/networking")>()),
fetchMemoryList: fetchMemoryListMock,
}));
vi.mock("@tanstack/react-pacer/debouncer", () => ({
useDebouncedValue: (value: unknown) => [value, { cancel: vi.fn(), flush: vi.fn() }],
}));
const renderView = (accessToken: string | null) => {
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
return render(
@ -35,6 +48,28 @@ const renderView = (accessToken: string | null) => {
};
describe("MemoryView", () => {
beforeEach(() => {
fetchMemoryListMock.mockReset();
fetchMemoryListMock.mockResolvedValue({ memories: [], total: 0 });
});
it("queries the server with the search box value as `search` and resets to page 1", async () => {
renderView("token");
await waitFor(() => expect(fetchMemoryListMock).toHaveBeenCalled());
act(() => captured.current?.onPaginationChange({ pageIndex: 2, pageSize: 50 }));
await waitFor(() =>
expect(fetchMemoryListMock).toHaveBeenLastCalledWith("token", expect.objectContaining({ page: 3 })),
);
act(() => captured.current?.onSearchChange("mem-abc123"));
await waitFor(() =>
expect(fetchMemoryListMock).toHaveBeenLastCalledWith("token", { search: "mem-abc123", page: 1, pageSize: 50 }),
);
expect(captured.current?.hasActiveSearch).toBe(true);
});
it("keeps the table out of the skeleton state when the token is null (disabled query)", () => {
renderView(null);

View file

@ -43,10 +43,8 @@ export const MemoryView: React.FC<MemoryViewProps> = ({ accessToken }) => {
queryKey: [MEMORY_LIST_KEY, debouncedSearch, pagination.pageIndex, pagination.pageSize],
queryFn: () => {
if (!accessToken) throw new Error("Access token required");
// Prefix search matches the Redis-style mental model (namespace scan):
// typing "user:" finds "user:profile", "user:prefs", etc.
return fetchMemoryList(accessToken, {
keyPrefix: debouncedSearch || undefined,
search: debouncedSearch || undefined,
page: pagination.pageIndex + 1,
pageSize: pagination.pageSize,
});

View file

@ -542,6 +542,19 @@ describe("server-side filtering – the LIT-4080 regression guard", () => {
expect((lastCall[2] ?? {}).userID).toBeUndefined();
});
});
it("sends the search box as the combined alias-or-ID search rather than the key-alias filter", async () => {
renderWithProviders(<VirtualKeysTable />);
fireEvent.change(screen.getByPlaceholderText(/Search by key alias or ID/), { target: { value: mockKey.token } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ search: mockKey.token }));
});
const lastOptions = mockUseKeys.mock.calls.at(-1)?.[2];
expect(lastOptions?.selectedKeyAlias).toBeUndefined();
expect(lastOptions?.keyHash).toBeUndefined();
});
});
describe("pagination display – total count comes from useKeys", () => {
@ -663,7 +676,7 @@ describe("table state lives in the URL so it survives leaving and returning to t
expect(mockUseKeys).toHaveBeenLastCalledWith(
3,
25,
expect.objectContaining({ selectedKeyAlias: "prod", sortBy: "spend", sortOrder: "asc" }),
expect.objectContaining({ search: "prod", sortBy: "spend", sortOrder: "asc" }),
);
});
expect(screen.getByPlaceholderText(/Search by key alias/)).toHaveValue("prod");
@ -736,7 +749,7 @@ describe("table state lives in the URL so it survives leaving and returning to t
fireEvent.change(screen.getByPlaceholderText(/Search by key alias/), { target: { value: "prod" } });
await waitFor(() => {
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ selectedKeyAlias: "prod" }));
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ search: "prod" }));
});
await waitFor(() => {
expect(lastSearchParam(onUrlUpdate, "page")).toBeNull();

View file

@ -118,7 +118,7 @@ export function VirtualKeysTable({ headerActions }: VirtualKeysTableProps) {
const keyListOptions = {
teamID: appliedFilters.team_id || undefined,
organizationID: appliedFilters.org_id || undefined,
selectedKeyAlias: searchQuery.trim() || undefined,
search: searchQuery.trim() || undefined,
userID: appliedFilters.user_id || undefined,
keyHash: appliedFilters.key_hash || undefined,
sortBy,
@ -291,7 +291,7 @@ export function VirtualKeysTable({ headerActions }: VirtualKeysTableProps) {
table={table}
searchValue={searchInput}
onSearchChange={handleSearchChange}
searchPlaceholder="Search by key alias…"
searchPlaceholder="Search by key alias or ID…"
onRefresh={() => refetch?.()}
isRefreshing={isFetching}
onOpenFilters={() => setFiltersOpen(true)}

View file

@ -854,3 +854,48 @@ describe("userListCall search serialization", () => {
expect(lastParams(mockFetch).get("user_email")).toBe("ada@example.com");
});
});
describe("fetchMemoryList search serialization", () => {
const originalFetch = global.fetch;
afterEach(() => {
global.fetch = originalFetch;
});
const mockOkFetch = () => {
const emptyPage = { memories: [], total: 0 };
const mockFetch = vi.fn().mockResolvedValue({ ok: true, json: vi.fn().mockResolvedValue(emptyPage) } as any);
global.fetch = mockFetch as any;
return mockFetch;
};
const lastParams = (mockFetch: ReturnType<typeof vi.fn>) => {
const [url] = mockFetch.mock.calls.at(-1) ?? [];
return new URL(url as string, "http://example.com").searchParams;
};
it("sends the search box value as search and omits key_prefix and key", async () => {
const mockFetch = mockOkFetch();
await Networking.fetchMemoryList("token", { search: "mem-abc123", page: 1, pageSize: 50 });
const params = lastParams(mockFetch);
expect(params.get("search")).toBe("mem-abc123");
expect(params.has("key_prefix")).toBe(false);
expect(params.has("key")).toBe(false);
expect(params.get("page")).toBe("1");
expect(params.get("page_size")).toBe("50");
});
it("keeps key_prefix and key working when no search is given", async () => {
const mockFetch = mockOkFetch();
await Networking.fetchMemoryList("token", { keyPrefix: "user:" });
expect(lastParams(mockFetch).get("key_prefix")).toBe("user:");
expect(lastParams(mockFetch).has("search")).toBe(false);
await Networking.fetchMemoryList("token", { key: "user:profile" });
expect(lastParams(mockFetch).get("key")).toBe("user:profile");
expect(lastParams(mockFetch).has("search")).toBe(false);
});
});

View file

@ -2056,6 +2056,7 @@ interface UiSpendLogsParams {
exclude_internal_health_checks?: boolean;
group_by_session?: boolean;
session_cursor?: string;
search?: string;
}
interface UiSpendLogsCallOptions {
@ -6563,6 +6564,7 @@ interface UiAuditLogsParams {
changed_by_api_key?: string;
object_team_id?: string;
object_key_hash?: string;
search?: string | null;
sort_by?: string;
sort_order?: "asc" | "desc";
}
@ -8061,15 +8063,18 @@ export const fetchMemoryList = async (
options: {
key?: string;
keyPrefix?: string;
search?: string;
page?: number;
pageSize?: number;
} = {},
): Promise<MemoryListResponse> => {
const base = proxyBaseUrl ? `${proxyBaseUrl}/v1/memory` : `/v1/memory`;
const params = new URLSearchParams();
// keyPrefix takes precedence — backend also does, but we omit `key`
// Backend precedence is search > key_prefix > key; only the winner is sent
// to keep the URL clean and intent obvious.
if (options.keyPrefix) {
if (options.search) {
params.append("search", options.search);
} else if (options.keyPrefix) {
params.append("key_prefix", options.keyPrefix);
} else if (options.key) {
params.append("key", options.key);

View file

@ -30,6 +30,8 @@ vi.mock("@tanstack/react-pacer/debouncer", () => ({
const mockUseKeys = useKeys as MockedFunction<typeof useKeys>;
const KEY_HASH = "88a145505dd6e87e2ea166fcef1e4b53948dbdb32af6431dfd05ec06b571ee52";
const createMockKey = (overrides: Partial<KeyResponse> = {}): KeyResponse =>
({
token: "sk-test123",
@ -277,7 +279,7 @@ describe("TeamVirtualKeysTable", () => {
);
});
it("maps the search box to a server-side key-alias query", async () => {
it("maps the Key ID drawer filter to a server-side useKeys query and clears it", async () => {
const user = userEvent.setup();
mockUseKeys.mockReturnValue({
data: { keys: [createMockKey()], total_count: 1, current_page: 1, total_pages: 1 },
@ -288,11 +290,42 @@ describe("TeamVirtualKeysTable", () => {
renderWithProviders(<TeamVirtualKeysTable {...defaultProps} />);
fireEvent.change(await screen.findByTestId("datatable-search"), { target: { value: "check-002" } });
await user.click(await screen.findByTestId("datatable-filters-trigger"));
const drawerBody = await screen.findByTestId("filter-drawer-body");
fireEvent.change(within(drawerBody).getByPlaceholderText("Enter Key ID…"), { target: { value: KEY_HASH } });
await user.click(screen.getByTestId("filter-drawer-apply"));
await waitFor(() =>
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ selectedKeyAlias: "check-002" })),
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ keyHash: KEY_HASH })),
);
expect(screen.getByTestId("filter-chip-key_hash")).toHaveTextContent("Key ID");
await user.click(screen.getByTestId("datatable-clear-filters"));
await waitFor(() =>
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ keyHash: undefined })),
);
});
it("maps the search box to the combined alias-or-ID search rather than the key-alias filter", async () => {
mockUseKeys.mockReturnValue({
data: { keys: [createMockKey()], total_count: 1, current_page: 1, total_pages: 1 },
isPending: false,
isFetching: false,
refetch: vi.fn(),
} as unknown as ReturnType<typeof useKeys>);
renderWithProviders(<TeamVirtualKeysTable {...defaultProps} />);
const searchBox = await screen.findByTestId("datatable-search");
expect(searchBox).toHaveAttribute("placeholder", "Search by key alias or ID…");
fireEvent.change(searchBox, { target: { value: KEY_HASH } });
await waitFor(() =>
expect(mockUseKeys).toHaveBeenLastCalledWith(1, 50, expect.objectContaining({ search: KEY_HASH })),
);
const lastOptions = mockUseKeys.mock.calls.at(-1)?.[2];
expect(lastOptions?.selectedKeyAlias).toBeUndefined();
expect(lastOptions?.keyHash).toBeUndefined();
});
it("should show Loading keys when isPending", async () => {

View file

@ -68,19 +68,17 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi
const pageIndex = tablePagination.pageIndex;
const pageSize = tablePagination.pageSize;
const {
data: keys,
isPending: isLoading,
isFetching,
refetch,
} = useKeys(pageIndex + 1, pageSize, {
const keyListOptions = {
teamID: teamId,
selectedKeyAlias: searchQuery.trim() || undefined,
search: searchQuery.trim() || undefined,
userID: getFilterValue("user_id"),
keyHash: getFilterValue("key_hash"),
sortBy: sortBy || undefined,
sortOrder: sortOrder || undefined,
expand: "user",
});
};
const { data: keys, isPending: isLoading, isFetching, refetch } = useKeys(pageIndex + 1, pageSize, keyListOptions);
const displayKeys = useMemo(() => {
const kList = keys?.keys || [];
@ -481,11 +479,11 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi
table={table}
searchValue={searchInput}
onSearchChange={handleSearchChange}
searchPlaceholder="Search by key alias…"
searchPlaceholder="Search by key alias or ID…"
onRefresh={() => refetch?.()}
isRefreshing={isFetching}
onOpenFilters={() => setFiltersOpen(true)}
filterLabels={{ user_id: "User ID" }}
filterLabels={{ user_id: "User ID", key_hash: "Key ID" }}
/>
<DataTableFilterDrawer
table={table}
@ -495,13 +493,22 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi
description={`Narrow down keys for ${teamAlias ?? "this team"}`}
>
{({ get, set }) => (
<DataTableFilterField label="User ID">
<Input
value={(get("user_id") as string) ?? ""}
onChange={(event) => set("user_id", event.target.value)}
placeholder="Filter by user ID…"
/>
</DataTableFilterField>
<>
<DataTableFilterField label="User ID">
<Input
value={(get("user_id") as string) ?? ""}
onChange={(event) => set("user_id", event.target.value)}
placeholder="Filter by user ID…"
/>
</DataTableFilterField>
<DataTableFilterField label="Key ID">
<Input
value={(get("key_hash") as string) ?? ""}
onChange={(event) => set("key_hash", event.target.value)}
placeholder="Enter Key ID…"
/>
</DataTableFilterField>
</>
)}
</DataTableFilterDrawer>
</>

View file

@ -0,0 +1,145 @@
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { chooseSelectOption } from "../../../tests/test-utils";
import AuditLogsPanel from "./AuditLogsPanel";
vi.mock("../networking", async (importOriginal) => {
const actual = await importOriginal<typeof import("../networking")>();
return { ...actual, uiAuditLogsCall: vi.fn() };
});
// Resolve the debounced search synchronously so typed input reaches the query within the test tick.
vi.mock("@tanstack/react-pacer/debouncer", () => ({
useDebouncedValue: (value: unknown) => [value, { cancel: vi.fn(), flush: vi.fn() }],
}));
import { uiAuditLogsCall } from "../networking";
type AuditLogsParams = NonNullable<Parameters<typeof uiAuditLogsCall>[0]["params"]>;
const PAGE_SIZE = 50;
const ID_PARAM_KEYS = [
"search",
"object_id",
"changed_by",
"object_team_id",
"object_key_hash",
"action",
"table_name",
] as const satisfies readonly (keyof AuditLogsParams)[];
const respondWith = (total: number) => {
const response = { audit_logs: [], total, page: 1, page_size: PAGE_SIZE, total_pages: Math.ceil(total / PAGE_SIZE) };
return vi.mocked(uiAuditLogsCall).mockResolvedValue(response);
};
const lastCall = () => vi.mocked(uiAuditLogsCall).mock.calls.at(-1)?.[0];
const sentIdParams = () => ID_PARAM_KEYS.filter((key) => lastCall()?.params?.[key] !== undefined);
const defaultProps = {
accessToken: "sk-test",
token: "jwt-test",
userRole: "Admin",
userID: "user-1",
isActive: true,
premiumUser: true,
};
const renderPanel = () => {
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
return render(
<QueryClientProvider client={queryClient}>
<AuditLogsPanel {...defaultProps} />
</QueryClientProvider>,
);
};
const TEXT_FILTERS: { filterId: string; placeholder: string; paramKey: keyof AuditLogsParams }[] = [
{ filterId: "object_id", placeholder: "Enter object ID…", paramKey: "object_id" },
{ filterId: "changed_by", placeholder: "Enter user ID…", paramKey: "changed_by" },
{ filterId: "team_id", placeholder: "Enter team ID…", paramKey: "object_team_id" },
{ filterId: "key_hash", placeholder: "Enter key hash…", paramKey: "object_key_hash" },
];
const SELECT_FILTERS: {
label: string;
comboboxIndex: number;
option: string;
paramKey: keyof AuditLogsParams;
value: string;
}[] = [
{ label: "Action", comboboxIndex: 0, option: "Created", paramKey: "action", value: "created" },
{ label: "Table", comboboxIndex: 1, option: "Teams", paramKey: "table_name", value: "LiteLLM_TeamTable" },
];
describe("AuditLogsPanel", () => {
beforeEach(() => {
vi.clearAllMocks();
respondWith(0);
});
it("sends the typed search as params.search and returns to the first page", async () => {
const user = userEvent.setup();
respondWith(120);
renderPanel();
await waitFor(() => expect(uiAuditLogsCall).toHaveBeenCalled());
expect(lastCall()?.params?.search).toBeUndefined();
await user.click(screen.getByTestId("pagination-next"));
await waitFor(() => expect(lastCall()?.page).toBe(2));
await user.type(screen.getByTestId("datatable-search"), "team-abc");
await waitFor(() => expect(lastCall()?.params?.search).toBe("team-abc"));
expect(lastCall()?.page).toBe(1);
expect(sentIdParams()).toEqual(["search"]);
});
it("trims the search and drops params.search once the box is cleared", async () => {
const user = userEvent.setup();
renderPanel();
const input = await screen.findByTestId("datatable-search");
await user.type(input, " abc");
await waitFor(() => expect(lastCall()?.params?.search).toBe("abc"));
await user.clear(input);
await waitFor(() => expect(lastCall()?.params?.search).toBeUndefined());
expect(sentIdParams()).toEqual([]);
});
it.each(TEXT_FILTERS)("maps the $filterId drawer filter to params.$paramKey", async ({ placeholder, paramKey }) => {
const user = userEvent.setup();
renderPanel();
await waitFor(() => expect(uiAuditLogsCall).toHaveBeenCalled());
await user.click(screen.getByTestId("datatable-filters-trigger"));
fireEvent.change(await screen.findByPlaceholderText(placeholder), { target: { value: "val-1" } });
await user.click(screen.getByTestId("filter-drawer-apply"));
await waitFor(() => expect(lastCall()?.params?.[paramKey]).toBe("val-1"));
expect(sentIdParams()).toEqual([paramKey]);
});
it.each(SELECT_FILTERS)(
"maps the $label drawer select to params.$paramKey",
async ({ comboboxIndex, option, paramKey, value }) => {
const user = userEvent.setup();
renderPanel();
await waitFor(() => expect(uiAuditLogsCall).toHaveBeenCalled());
await user.click(screen.getByTestId("datatable-filters-trigger"));
const triggers = await screen.findAllByRole("combobox");
await chooseSelectOption(user, triggers[comboboxIndex], option);
await user.click(screen.getByTestId("filter-drawer-apply"));
await waitFor(() => expect(lastCall()?.params?.[paramKey]).toBe(value));
expect(sentIdParams()).toEqual([paramKey]);
},
);
});

View file

@ -1,7 +1,9 @@
import { useCallback, useState } from "react";
import { useDebouncedValue } from "@tanstack/react-pacer/debouncer";
import { useQuery, keepPreviousData } from "@tanstack/react-query";
import { ColumnFiltersState, OnChangeFn, PaginationState } from "@tanstack/react-table";
import { resolveLogoSrc } from "@/lib/assetPaths";
import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants";
import { uiAuditLogsCall } from "../networking";
import { AuditLogEntry } from "./AuditLogsTableColumns";
import { AuditLogsTable } from "./AuditLogsTable";
@ -39,9 +41,13 @@ export default function AuditLogsPanel({
}: AuditLogsProps) {
const [pagination, setPagination] = useState<PaginationState>({ pageIndex: 0, pageSize: PAGE_SIZE });
const [columnFilters, setColumnFilters] = useState<ColumnFiltersState>([]);
const [searchInput, setSearchInput] = useState("");
const [debouncedSearch] = useDebouncedValue(searchInput, { wait: DEBOUNCE_WAIT_MS });
const [selectedLog, setSelectedLog] = useState<AuditLogEntry | null>(null);
const [drawerOpen, setDrawerOpen] = useState(false);
const searchTerm = debouncedSearch.trim();
const getFilterValue = (columnId: string): string | undefined => {
const entry = columnFilters.find((filter) => filter.id === columnId);
return typeof entry?.value === "string" && entry.value.trim() ? entry.value.trim() : undefined;
@ -50,7 +56,7 @@ export default function AuditLogsPanel({
const canQueryAuditLogs = !!accessToken && !!token && !!userRole && !!userID && isActive && premiumUser;
const query = useQuery<AuditLogsResponse>({
queryKey: ["audit_logs", pagination.pageIndex, pagination.pageSize, columnFilters],
queryKey: ["audit_logs", pagination.pageIndex, pagination.pageSize, columnFilters, searchTerm],
queryFn: async () => {
if (!accessToken) {
return { audit_logs: [], total: 0, page: 1, page_size: pagination.pageSize, total_pages: 0 };
@ -60,6 +66,7 @@ export default function AuditLogsPanel({
page: pagination.pageIndex + 1,
page_size: pagination.pageSize,
params: {
search: searchTerm || undefined,
object_id: getFilterValue("object_id"),
changed_by: getFilterValue("changed_by"),
object_key_hash: getFilterValue("key_hash"),
@ -80,6 +87,11 @@ export default function AuditLogsPanel({
setPagination((prev) => ({ ...prev, pageIndex: 0 }));
}, []);
const handleSearchChange = useCallback((value: string) => {
setSearchInput(value);
setPagination((prev) => ({ ...prev, pageIndex: 0 }));
}, []);
const handleViewLog = useCallback((log: AuditLogEntry) => {
setSelectedLog(log);
setDrawerOpen(true);
@ -128,6 +140,8 @@ export default function AuditLogsPanel({
onPaginationChange={setPagination}
columnFilters={columnFilters}
onColumnFiltersChange={handleColumnFiltersChange}
searchValue={searchInput}
onSearchChange={handleSearchChange}
onRefresh={() => query.refetch()}
onViewLog={handleViewLog}
/>

View file

@ -120,6 +120,24 @@ describe("AuditLogsTable", () => {
expect(screen.getByText("No matching audit logs")).toBeInTheDocument();
});
it("renders the toolbar search box from the search props and forwards typed input", () => {
const onSearchChange = vi.fn();
renderTable({ searchValue: "team-", onSearchChange });
const input = screen.getByPlaceholderText("Search audit logs by ID…");
expect(input).toHaveValue("team-");
fireEvent.change(input, { target: { value: "team-7" } });
expect(onSearchChange).toHaveBeenCalledWith("team-7");
});
it("treats an active search as a filter for the empty state", () => {
const emptySearchResult = { data: [], rowCount: 0, searchValue: "zzz", onSearchChange: vi.fn() };
renderTable(emptySearchResult);
expect(screen.getByText("No matching audit logs")).toBeInTheDocument();
});
it("renders active filter chips with human-readable labels", () => {
const filters: ColumnFiltersState = [{ id: "action", value: "created" }];
renderTable({ columnFilters: filters });

View file

@ -24,6 +24,8 @@ interface AuditLogsTableProps {
onPaginationChange: OnChangeFn<PaginationState>;
columnFilters: ColumnFiltersState;
onColumnFiltersChange: OnChangeFn<ColumnFiltersState>;
searchValue?: string;
onSearchChange?: (value: string) => void;
onRefresh: () => void;
onViewLog: (log: AuditLogEntry) => void;
}
@ -102,11 +104,14 @@ export function AuditLogsTable({
onPaginationChange,
columnFilters,
onColumnFiltersChange,
searchValue,
onSearchChange,
onRefresh,
onViewLog,
}: AuditLogsTableProps) {
const [filtersOpen, setFiltersOpen] = useState(false);
const columns = useMemo(() => getAuditLogsTableColumns({ onViewLog }), [onViewLog]);
const hasActiveSearch = Boolean(searchValue?.trim());
return (
<DataTable
@ -122,12 +127,15 @@ export function AuditLogsTable({
onColumnFiltersChange={onColumnFiltersChange}
isLoading={isLoading}
loadingMessage="Loading audit logs…"
noDataMessage={<AuditLogsEmptyState filtered={columnFilters.length > 0} />}
noDataMessage={<AuditLogsEmptyState filtered={columnFilters.length > 0 || hasActiveSearch} />}
size="compact"
toolbar={(table) => (
<>
<DataTableToolbar
table={table}
searchValue={searchValue}
onSearchChange={onSearchChange}
searchPlaceholder="Search audit logs by ID…"
onRefresh={onRefresh}
isRefreshing={isRefreshing}
onOpenFilters={() => setFiltersOpen(true)}

View file

@ -54,6 +54,14 @@ vi.mock("./LogDetailsDrawer", () => ({
},
}));
const debounce = vi.hoisted(() => ({ settled: null as string | null }));
vi.mock("@tanstack/react-pacer/debouncer", () => ({
useDebouncedValue: vi.fn((value: unknown) => [debounce.settled ?? value, { cancel: vi.fn(), flush: vi.fn() }]),
}));
import { useDebouncedValue } from "@tanstack/react-pacer/debouncer";
import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants";
import { uiSpendLogsCall } from "../networking";
const logEntry = (overrides: Partial<LogEntry>): LogEntry => ({
@ -136,6 +144,7 @@ describe("RequestLogsPanel", () => {
sessionStorage.clear();
testQueryClient.clear();
respondWith([]);
debounce.settled = null;
});
describe("server-grouped session pagination (#38060)", () => {
@ -322,9 +331,8 @@ describe("RequestLogsPanel", () => {
});
});
describe("search by request id (LIT-3981)", () => {
it("sends the typed request id to the server on the first page instead of filtering the loaded rows", async () => {
const user = userEvent.setup();
describe("search by any id (LIT-3981, LIT-4741)", () => {
it("sends the typed id to the server as search on the first page instead of filtering the loaded rows", async () => {
renderPanel();
await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled());
@ -334,9 +342,54 @@ describe("RequestLogsPanel", () => {
await waitFor(() => {
const call = lastCall();
if (!call) throw new Error("uiSpendLogsCall was not called");
expect(call.params?.request_id).toBe("req-on-another-page");
expect(call.params?.search).toBe("req-on-another-page");
expect(call.page).toBe(1);
});
expect(lastCall()?.params?.request_id).toBeUndefined();
expect(lastCall()?.params?.session_cursor).toBeUndefined();
});
it("sends the debounced value to the server while the box shows what is being typed", async () => {
debounce.settled = "settled-id";
renderPanel();
await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled());
fireEvent.change(screen.getByTestId("datatable-search"), { target: { value: "still-typing" } });
expect(screen.getByTestId("datatable-search")).toHaveValue("still-typing");
await waitFor(() =>
expect(useDebouncedValue).toHaveBeenLastCalledWith("still-typing", { wait: DEBOUNCE_WAIT_MS }),
);
await waitFor(() => expect(lastCall()?.params?.search).toBe("settled-id"));
const sentLiveValue = vi
.mocked(uiSpendLogsCall)
.mock.calls.some(([options]) => options.params?.search === "still-typing");
expect(sentLiveValue).toBe(false);
});
it("shows a Search chip whose remove button clears the box and restores the unsearched listing", async () => {
const user = userEvent.setup();
vi.mocked(uiSpendLogsCall).mockImplementation(async ({ params }) => {
const data =
params?.search === "sess-42"
? [logEntry({ request_id: "req-sess", session_id: "sess-42" })]
: [logEntry({ request_id: "req-initial" })];
return { data, total: data.length, page: 1, page_size: 50, total_pages: 1 };
});
renderPanel();
await waitFor(() => expect(row("req-initial")).not.toBeNull());
fireEvent.change(screen.getByTestId("datatable-search"), { target: { value: "sess-42" } });
await waitFor(() => expect(row("req-sess")).not.toBeNull());
expect(row("req-initial")).toBeNull();
expect(screen.getByTestId("filter-chip-search")).toHaveTextContent("Search:sess-42");
await user.click(screen.getByRole("button", { name: "Remove Search filter" }));
expect(screen.getByTestId("datatable-search")).toHaveValue("");
await waitFor(() => expect(row("req-initial")).not.toBeNull());
expect(row("req-sess")).toBeNull();
});
});

View file

@ -1,11 +1,13 @@
"use client";
import { useDebouncedValue } from "@tanstack/react-pacer/debouncer";
import { useQuery, type UseQueryOptions } from "@tanstack/react-query";
import type { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table";
import moment from "moment";
import { useCallback, useEffect, useMemo, useState } from "react";
import { AutoRouterModelGroupsProvider } from "@/components/shared/table_cells";
import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants";
import type { KeyResponse } from "../key_team_helpers/key_list";
import { keyInfoV1Call, uiSpendLogsCall } from "../networking";
import KeyInfoView from "../templates/key_info_view";
@ -75,12 +77,22 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID,
sessionStorage.setItem("excludeInternalHealthChecks", JSON.stringify(excludeInternalHealthChecks));
}, [excludeInternalHealthChecks]);
const searchTerm = useMemo(() => {
const entry = columnFilters.find((filter) => filter.id === LOG_FILTER_IDS.SEARCH);
return typeof entry?.value === "string" ? entry.value : "";
}, [columnFilters]);
const [debouncedSearch] = useDebouncedValue(searchTerm, { wait: DEBOUNCE_WAIT_MS });
const queryColumnFilters = useMemo<ColumnFiltersState>(() => {
const others = columnFilters.filter((filter) => filter.id !== LOG_FILTER_IDS.SEARCH);
return debouncedSearch === "" ? others : [...others, { id: LOG_FILTER_IDS.SEARCH, value: debouncedSearch }];
}, [columnFilters, debouncedSearch]);
const { logsQuery, filteredLogs, allTeams, usesSessionCursor } = useLogFilterLogic({
accessToken,
token,
userRole,
userID,
columnFilters,
columnFilters: queryColumnFilters,
activeTab: isActive ? "request logs" : "inactive",
isLiveTail,
excludeInternalHealthChecks,
@ -155,15 +167,10 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID,
const rows: LogEntry[] = filteredLogs.data;
const searchTerm = useMemo(() => {
const entry = columnFilters.find((filter) => filter.id === LOG_FILTER_IDS.REQUEST_ID);
return typeof entry?.value === "string" ? entry.value : "";
}, [columnFilters]);
const handleSearchChange = useCallback((value: string) => {
setColumnFilters((previous) => {
const others = previous.filter((filter) => filter.id !== LOG_FILTER_IDS.REQUEST_ID);
return value === "" ? others : [...others, { id: LOG_FILTER_IDS.REQUEST_ID, value }];
const others = previous.filter((filter) => filter.id !== LOG_FILTER_IDS.SEARCH);
return value === "" ? others : [...others, { id: LOG_FILTER_IDS.SEARCH, value }];
});
setSessionCursors({});
setPagination((previous) => ({ ...previous, pageIndex: 0 }));

View file

@ -108,7 +108,7 @@ export function RequestLogsTable({
table={table}
searchValue={searchValue}
onSearchChange={onSearchChange}
searchPlaceholder="Search by Request ID"
searchPlaceholder="Search logs by ID…"
onRefresh={onRefresh}
isRefreshing={isRefreshing}
onOpenFilters={() => setFiltersOpen(true)}

View file

@ -91,6 +91,7 @@ describe("useLogFilterLogic", () => {
{ id: LOG_FILTER_IDS.ERROR_CODE, value: "429", param: "error_code" },
{ id: LOG_FILTER_IDS.ERROR_MESSAGE, value: "rate limited", param: "error_message" },
{ id: LOG_FILTER_IDS.USER_ID, value: "user-9", param: "user_id" },
{ id: LOG_FILTER_IDS.SEARCH, value: "any-id", param: "search" },
];
it.each(cases)("sends $id as $param", async ({ id, value, param }) => {

View file

@ -33,6 +33,7 @@ export const LOG_FILTER_IDS = {
PUBLIC_MODEL_OR_SEARCH_TOOL: "model",
REQUEST_ID: "request_id",
USER_ID: "user_id",
SEARCH: "search",
} as const;
export const LOG_FILTER_LABELS: Record<string, string> = {
@ -48,6 +49,7 @@ export const LOG_FILTER_LABELS: Record<string, string> = {
[LOG_FILTER_IDS.SESSION_ID]: "Session ID",
[LOG_FILTER_IDS.MODEL_ID]: "Model",
[LOG_FILTER_IDS.PUBLIC_MODEL_OR_SEARCH_TOOL]: "Public model / search tool",
[LOG_FILTER_IDS.SEARCH]: "Search",
};
export interface LogsWindow {
@ -175,6 +177,7 @@ export function useLogFilterLogic({
api_key: getFilterValue(columnFilters, LOG_FILTER_IDS.KEY_HASH),
team_id: getFilterValue(columnFilters, LOG_FILTER_IDS.TEAM_ID),
request_id: getFilterValue(columnFilters, LOG_FILTER_IDS.REQUEST_ID),
search: getFilterValue(columnFilters, LOG_FILTER_IDS.SEARCH),
session_id: getFilterValue(columnFilters, LOG_FILTER_IDS.SESSION_ID),
user_id: userIdFilter,
end_user: getFilterValue(columnFilters, LOG_FILTER_IDS.END_USER),

View file

@ -40904,6 +40904,8 @@ export interface operations {
object_team_id?: string | null;
/** @description Filter by token (key hash) present in before_value or updated_values JSON (PostgreSQL only) */
object_key_hash?: string | null;
/** @description Match a row whose id, object_id, changed_by, or changed_by_api_key equals this value */
search?: string | null;
/** @description Column to sort by (e.g. 'updated_at', 'action', 'table_name') */
sort_by?: string | null;
/** @description Sort order ('asc' or 'desc') */
@ -49622,6 +49624,8 @@ export interface operations {
key_hash?: string | null;
/** @description Filter keys by key alias. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching. */
key_alias?: string | null;
/** @description Combined search: matches keys whose token (key hash) equals the value OR whose key_alias contains it (case-insensitive). */
search?: string | null;
/** @description Return full key object */
return_full_object?: boolean;
/** @description Include all keys for teams that user is an admin of. */
@ -56880,6 +56884,8 @@ export interface operations {
group_by_session?: boolean;
/** @description Keyset cursor '<last_activity>|<api_key>|<session_key>' from a previous group_by_session page. UI route only, honored when sorting by startTime */
session_cursor?: string | null;
/** @description Match a log whose request_id, api_key (hash), team_id, user, end_user, session_id, or model_id equals this value. request_id matches across all time; the other columns match inside start_date/end_date, which stay required */
search?: string | null;
};
header?: never;
path?: never;
@ -56996,6 +57002,8 @@ export interface operations {
group_by_session?: boolean;
/** @description Keyset cursor '<last_activity>|<api_key>|<session_key>' from a previous group_by_session page. UI route only, honored when sorting by startTime */
session_cursor?: string | null;
/** @description Match a log whose request_id, api_key (hash), team_id, user, end_user, session_id, or model_id equals this value. request_id matches across all time; the other columns match inside start_date/end_date, which stay required */
search?: string | null;
};
header?: never;
path?: never;
@ -63308,6 +63316,8 @@ export interface operations {
key?: string | null;
/** @description Filter by key prefix (Redis-style namespace scan). Mutually exclusive with `key`; if both are provided, `key_prefix` wins. */
key_prefix?: string | null;
/** @description Match entries whose key starts with this value or whose memory_id equals it. Takes precedence over `key_prefix` and `key` when provided. */
search?: string | null;
page?: number;
page_size?: number;
};