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

# Conflicts:
#	litellm/llms/openai/workload_identity.py
#	tests/test_litellm/llms/openai/test_openai_workload_identity.py
This commit is contained in:
mateo-berri 2026-09-03 16:39:40 -07:00
commit 37f2de1b0e
78 changed files with 3802 additions and 388 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

@ -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

@ -9,7 +9,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
@ -17,8 +17,6 @@ if TYPE_CHECKING:
from openai.auth import SubjectTokenProvider, WorkloadIdentity, WorkloadIdentityAuth
OPENAI_WIF_CLIENT_ID: Final = "litellm"
_OPENAI_API_HOST: Final = "api.openai.com"
_OPENAI_REGIONAL_HOST_SUFFIX: Final = f".{_OPENAI_API_HOST}"
_SDK_UPGRADE_MESSAGE: Final = (
"OpenAI workload identity federation requires openai>=2.32.0. "
"Upgrade the installed openai package to use OPENAI_IDENTITY_PROVIDER_ID / "
@ -87,9 +85,7 @@ def _targets_openai_api(api_base: str | None) -> bool:
if api_base is None:
return True
parsed: Final = urlparse(api_base)
if parsed.scheme != "https" or parsed.hostname is None:
return False
return parsed.hostname == _OPENAI_API_HOST or parsed.hostname.endswith(_OPENAI_REGIONAL_HOST_SUFFIX)
return parsed.scheme == "https" and is_openai_backed_api_base(api_base)
@lru_cache(maxsize=16)

View file

@ -448,6 +448,17 @@ def _append_query_params(url: str, params: dict[str, str]) -> str:
return urlunparse(parsed._replace(query=urlencode(query_params)))
def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCPServer | None:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
by_name: Final = global_mcp_server_manager.get_mcp_server_by_name(lookup, client_ip=client_ip)
if by_name is not None:
return by_name
return global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
def _resolve_oauth2_server_for_root_endpoints(
client_ip: str | None = None,
) -> MCPServer | None:
@ -1766,10 +1777,6 @@ async def authorize(
resource: str | None = None,
):
# Redirect to real OAuth provider with PKCE support
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id):
if is_proxy_api_resource(request, resource):
return await native_client_authorize(
@ -1797,9 +1804,7 @@ async def authorize(
lookup_name: Final[str | None] = mcp_server_name or client_id
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = (
global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip) if lookup_name else None
)
mcp_server = _resolve_mcp_server_by_name_or_id(lookup_name, client_ip) if lookup_name else None
if mcp_server is None and mcp_server_name is None:
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if mcp_server is None:
@ -1855,10 +1860,6 @@ async def token_endpoint(
3. Return the token
4. Return a virtual key in this response
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
if mcp_server_name is None and is_gateway_dcr_client_id(client_id):
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load
master_key,
@ -1882,7 +1883,7 @@ async def token_endpoint(
lookup_name: Final = mcp_server_name or client_id
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip)
mcp_server = _resolve_mcp_server_by_name_or_id(lookup_name, client_ip)
if mcp_server is None and mcp_server_name is None:
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
if mcp_server is None:
@ -2288,10 +2289,6 @@ async def _build_oauth_protected_resource_response(
Returns:
OAuth protected resource metadata dict
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
request_base_url: Final = get_request_base_url(request)
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
explicitly_named: Final = mcp_server_name is not None
@ -2304,7 +2301,7 @@ async def _build_oauth_protected_resource_response(
mcp_server: MCPServer | None = None
if mcp_server_name:
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
mcp_server = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
# Build resource URL based on the pattern
if mcp_server_name:
@ -2562,10 +2559,6 @@ def _build_oauth_authorization_server_response(
registry lookups; unlike :func:`_build_oauth_protected_resource_response`
it does not need to await any upstream IO.
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
request_base_url: Final = get_request_base_url(request)
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
explicitly_named: Final = mcp_server_name is not None
@ -2583,7 +2576,7 @@ def _build_oauth_authorization_server_response(
mcp_server: MCPServer | None = None
if mcp_server_name:
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
mcp_server = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth authorization server")
@ -2709,10 +2702,6 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
@router.post("/{mcp_server_name}/register")
@router.post("/register")
async def register_client(request: Request, mcp_server_name: str | None = None):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
# Get the correct base URL considering X-Forwarded-* headers
request_base_url: Final = get_request_base_url(request)
@ -2748,7 +2737,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
)
return dummy_return
mcp_server: Final = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
mcp_server: Final = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
if mcp_server is None:
return dummy_return
return await register_client_with_server(

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
@ -6147,13 +6145,13 @@ class MCPServerManager:
internal_networks = IPAddressUtils.parse_internal_networks(general_settings.get("mcp_internal_ip_ranges"))
return IPAddressUtils.is_internal_ip(client_ip, internal_networks)
def get_mcp_server_by_id(self, server_id: str) -> MCPServer | None:
"""
Get the MCP Server from the server id
"""
def get_mcp_server_by_id(self, server_id: str, client_ip: str | None = None) -> MCPServer | None:
"""Get the MCP Server from the server id."""
registry: Final = self.get_registry()
for server in registry.values():
if server.server_id == server_id:
if not self._is_server_accessible_from_ip(server, client_ip):
return None
return server
return None

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

@ -5,10 +5,13 @@ Handles agent permission checking for keys and teams using object_permission_id.
Follows the same pattern as MCP permission handling.
"""
import asyncio
from collections.abc import Awaitable, Callable, Sequence
from dataclasses import dataclass
from typing import Final, TypeAlias
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
from litellm.proxy._types import (
UI_TEAM_ID,
LiteLLM_ObjectPermissionTable,
@ -443,15 +446,47 @@ class AgentRequestHandler:
return []
async def accessible_agents(user_api_key_auth: UserAPIKeyAuth) -> tuple[AgentResponse, ...]:
"""Every registry agent for proxy admins, else the agents the key's and team's grants reach."""
def _granted_ids(access: AgentAccess) -> frozenset[str]:
match access:
case UnrestrictedAgentAccess():
return frozenset()
case RestrictedAgentAccess(agent_ids):
return agent_ids
ResolveAgentAccess: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[AgentAccess]]
EffectiveAuthContexts: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[Sequence[UserAPIKeyAuth]]]
async def _granted_agent_ids(
user_api_key_auth: UserAPIKeyAuth,
resolve_access: ResolveAgentAccess,
effective_contexts: EffectiveAuthContexts,
) -> frozenset[str]:
"""Union of the explicit grants reachable from the key, its team, or (for a dashboard session)
the user's real teams and user row. No grant anywhere yields the empty set, unlike the
open-by-default ``resolve_agent_access`` that guards direct access."""
accesses: Final = await asyncio.gather(
*(resolve_access(auth_context) for auth_context in await effective_contexts(user_api_key_auth))
)
return frozenset().union(*(_granted_ids(access) for access in accesses))
async def accessible_agents(
user_api_key_auth: UserAPIKeyAuth,
all_agents: tuple[AgentResponse, ...] | None = None,
resolve_access: ResolveAgentAccess | None = None,
effective_contexts: EffectiveAuthContexts = build_effective_auth_contexts,
) -> tuple[AgentResponse, ...]:
"""Every registry agent for proxy admins, else only the agents the caller was granted."""
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
all_agents: Final = global_agent_registry.get_agent_list()
agents: Final = global_agent_registry.get_agent_list() if all_agents is None else all_agents
if user_api_key_auth.user_role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN.value):
return all_agents
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth=user_api_key_auth):
case UnrestrictedAgentAccess():
return all_agents
case RestrictedAgentAccess(allowed_agent_ids):
return tuple(agent for agent in all_agents if agent.agent_id in allowed_agent_ids)
return agents
allowed_agent_ids: Final = await _granted_agent_ids(
user_api_key_auth,
AgentRequestHandler.resolve_agent_access if resolve_access is None else resolve_access,
effective_contexts,
)
return tuple(agent for agent in agents if agent.agent_id in allowed_agent_ids)

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

@ -5,10 +5,11 @@ This hook uses the DBSpendUpdateWriter to batch-write response IDs to the databa
instead of writing immediately on each request.
"""
from collections.abc import AsyncGenerator
from collections.abc import AsyncGenerator, Callable, Mapping
from typing import TYPE_CHECKING, Any, Final, cast
from fastapi import HTTPException
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
@ -32,6 +33,44 @@ _RESPONSES_API_PROVIDER_PREFIX: Final = "/openai"
_RESPONSES_API_CREATE_ROUTES: Final = frozenset({"/v1/responses", "/responses"})
_RESPONSE_PAYLOAD_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def _response_payload(response_obj: object) -> Mapping[str, object] | None:
try:
return _RESPONSE_PAYLOAD_ADAPTER.validate_python(response_obj)
except ValidationError:
return None
def _rewrite_advertised_id(
event: BaseLiteLLMOpenAIResponseObject,
rewrite: Callable[[str], str],
) -> BaseLiteLLMOpenAIResponseObject:
event_id: Final = getattr(event, "id", None)
if isinstance(event_id, str) and event_id.startswith("resp_"):
setattr(event, "id", rewrite(event_id))
return event
nested: Final = getattr(event, "response", None)
if isinstance(nested, ResponsesAPIResponse):
setattr(nested, "id", rewrite(nested.id))
setattr(event, "response", nested)
return event
payload: Final = _response_payload(nested)
if payload is None:
return event
payload_id: Final = payload.get("id")
if not isinstance(payload_id, str):
return event
rewritten: Final = {**payload, "id": rewrite(payload_id)} # mutable-ok: pydantic cannot serialize a frozen map
setattr(event, "response", rewritten)
return event
def _is_responses_api_create_route(request_route: str | None) -> bool:
if request_route is None:
return False
@ -196,10 +235,6 @@ class ResponsesIDSecurity(CustomLogger):
user_api_key_dict: "UserAPIKeyAuth",
request_cache: dict[str, str] | None = None,
) -> BaseLiteLLMOpenAIResponseObject:
# encrypt the response id using the symmetric key
# encrypt the response id, and encode the user id and response id in base64
# Check if signing key is available
signing_key: Final = self._get_signing_key()
if signing_key is None:
verbose_proxy_logger.debug(
@ -210,43 +245,22 @@ class ResponsesIDSecurity(CustomLogger):
)
return response
response_id: Final = getattr(response, "id", None)
response_obj: Final = getattr(response, "response", None)
def encrypt(original_id: str) -> str:
cached: Final = request_cache.get(original_id) if request_cache is not None else None
if cached is not None:
return cached
if response_id and isinstance(response_id, str) and response_id.startswith("resp_"):
# Check request-scoped cache first (for streaming consistency)
if request_cache is not None and response_id in request_cache:
setattr(response, "id", request_cache[response_id])
else:
encrypted_response_id = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
response_id,
user_api_key_dict.user_id or "",
user_api_key_dict.team_id or "",
)
managed_id: Final = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
original_id,
user_api_key_dict.user_id or "",
user_api_key_dict.team_id or "",
)
encrypted_id: Final = f"resp_{encrypt_value_helper(value=managed_id)}"
if request_cache is not None:
request_cache[original_id] = encrypted_id
return encrypted_id
encoded_user_id_and_response_id = encrypt_value_helper(value=encrypted_response_id)
encrypted_id = f"resp_{encoded_user_id_and_response_id}"
if request_cache is not None:
request_cache[response_id] = encrypted_id
setattr(response, "id", encrypted_id)
elif response_obj and isinstance(response_obj, ResponsesAPIResponse):
# Check request-scoped cache first (for streaming consistency)
if request_cache is not None and response_obj.id in request_cache:
setattr(response_obj, "id", request_cache[response_obj.id])
else:
encrypted_response_id = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
response_obj.id,
user_api_key_dict.user_id or "",
user_api_key_dict.team_id or "",
)
encoded_user_id_and_response_id = encrypt_value_helper(value=encrypted_response_id)
encrypted_id = f"resp_{encoded_user_id_and_response_id}"
if request_cache is not None:
request_cache[response_obj.id] = encrypted_id
setattr(response_obj, "id", encrypted_id)
setattr(response, "response", response_obj)
return response
return _rewrite_advertised_id(response, encrypt)
async def async_post_call_success_hook(
self,

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

@ -10,6 +10,7 @@ All /team management endpoints
"""
import asyncio
import copy
import json
import math
import traceback
@ -2189,6 +2190,26 @@ async def update_team(
if field in updated_kv
}
_writes_metadata_backed_field: Final = any(
field in updated_kv
for field in (
*LiteLLM_ManagementEndpoint_MetadataFields,
*LiteLLM_ManagementEndpoint_MetadataFields_Premium,
)
)
if isinstance(existing_team_row.metadata, dict):
if "metadata" not in updated_kv and (_team_member_fields_in_request or _writes_metadata_backed_field):
updated_kv["metadata"] = copy.deepcopy(existing_team_row.metadata)
elif isinstance(updated_kv.get("metadata"), dict):
updated_kv["metadata"] = {
**updated_kv["metadata"],
**{
key: existing_team_row.metadata[key]
for key in TeamMemberBudgetHandler.SYSTEM_MANAGED_METADATA_KEYS
if key in existing_team_row.metadata
},
}
if _team_member_fields_in_request and TeamMemberBudgetHandler.should_create_budget(
team_member_budget=data.team_member_budget,
team_member_rpm_limit=data.team_member_rpm_limit,

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

@ -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

@ -42,7 +42,7 @@
"limit": 52
},
"B010": {
"limit": 190
"limit": 187
},
"B018": {
"limit": 2

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

@ -9,14 +9,16 @@ be forced by a sequential script: it needs one request to be genuinely mid-fligh
other commits. A mocked prisma cannot arbitrate that either, since the property under test
is whether Postgres's own advisory lock actually serializes the two requests.
These tests pin the interleaving the same way test_access_group_team_sync.py does: a second
real connection holds the team's advisory lock in its own transaction, so the function under
test is provably blocked on it rather than hoping a sleep lands in the right gap.
These tests pin the interleaving without a timing assumption: a second real connection holds
the team's advisory lock in its own transaction, and the test then waits for Postgres itself
to report the endpoint queued behind that exact lock. A sleep can only guess whether the
endpoint has reached the lock yet; pg_locks answers it.
"""
import asyncio
import json
import os
import time
import uuid
from contextlib import asynccontextmanager
from datetime import timedelta
@ -39,6 +41,47 @@ _DELETE_SEEDED = 'DELETE FROM "LiteLLM_TeamMembership" WHERE team_id = $1'
_DELETE_USER = 'DELETE FROM "LiteLLM_UserTable" WHERE user_id = $1'
_DELETE_TEAM = 'DELETE FROM "LiteLLM_TeamTable" WHERE team_id = $1'
_LOCK_SQL = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked"
_HELD_LOCK_KEY_SQL = (
"SELECT classid::bigint AS classid, objid::bigint AS objid FROM pg_locks "
"WHERE locktype = 'advisory' AND granted AND pid = pg_backend_pid()"
)
_LOCK_WAITER_SQL = (
"SELECT count(*)::int AS waiters FROM pg_locks "
"WHERE locktype = 'advisory' AND NOT granted "
"AND classid::bigint = $1 AND objid::bigint = $2"
)
_LOCK_WAIT_TIMEOUT_SECONDS = 20.0
_LOCK_POLL_SECONDS = 0.01
async def _hold_team_lock(held, team_id: str) -> tuple[int, int]:
"""Take the team's advisory lock and return its pg_locks key.
Reading the key back off our own backend avoids re-deriving hashtext()'s signed
32-bit split here, and pins the watcher to this lock rather than to any advisory
lock another xdist worker happens to hold on the same database."""
await held.query_raw(_LOCK_SQL, team_id)
rows = await held.query_raw(_HELD_LOCK_KEY_SQL)
assert len(rows) == 1, f"expected exactly one advisory lock on the blocking connection, got {rows}"
return rows[0]["classid"], rows[0]["objid"]
async def _await_lock_contention(watcher, lock_key: tuple[int, int], task, what: str) -> None:
"""Block until Postgres reports `task` queued behind the held lock.
This is the assertion that the endpoint serializes on the team's advisory lock, and it
is what a fixed sleep was standing in for: the endpoint is only provably waiting once a
non-granted advisory lock on the same key exists. `watcher` must be a connection that is
not itself blocked, so it can observe the queue."""
classid, objid = lock_key
deadline = time.monotonic() + _LOCK_WAIT_TIMEOUT_SECONDS
while time.monotonic() < deadline:
if task.done():
raise AssertionError(f"{what} returned without waiting on the team's advisory lock") from task.exception()
if (await watcher.query_raw(_LOCK_WAITER_SQL, classid, objid))[0]["waiters"]:
return
await asyncio.sleep(_LOCK_POLL_SECONDS)
raise AssertionError(f"{what} never queued on the team's advisory lock within {_LOCK_WAIT_TIMEOUT_SECONDS}s")
def _race_ids() -> tuple[str, str]:
@ -110,10 +153,8 @@ async def test_member_add_blocked_by_delete_writes_no_dangling_reference():
blocker = Prisma()
await blocker.connect()
lock_acquired = asyncio.Event()
async def add_member():
lock_acquired.set()
await _add_team_members_to_team(
data=TeamMemberAddRequest(
team_id=team_id,
@ -128,11 +169,9 @@ async def test_member_add_blocked_by_delete_writes_no_dangling_reference():
try:
async with blocker.tx(timeout=timedelta(seconds=30)) as held:
await held.query_raw(_LOCK_SQL, team_id)
lock_key = await _hold_team_lock(held, team_id)
task = asyncio.create_task(add_member())
await lock_acquired.wait()
await asyncio.sleep(0.2)
assert not task.done(), "member_add did not wait on the team's advisory lock"
await _await_lock_contention(db, lock_key, task, "member_add")
# the delete wins the race: strip the team row while the lock is held
await held.execute_raw(_DELETE_TEAM, team_id)
@ -185,10 +224,8 @@ async def test_member_delete_blocked_by_member_add_removes_from_the_fresh_roster
blocker = Prisma()
await blocker.connect()
lock_acquired = asyncio.Event()
async def run_delete():
lock_acquired.set()
return await team_member_delete(
data=TeamMemberDeleteRequest(team_id=team_id, user_id=user_id),
user_api_key_dict=_admin_auth(),
@ -196,11 +233,9 @@ async def test_member_delete_blocked_by_member_add_removes_from_the_fresh_roster
try:
async with blocker.tx(timeout=timedelta(seconds=30)) as held:
await held.query_raw(_LOCK_SQL, team_id)
lock_key = await _hold_team_lock(held, team_id)
task = asyncio.create_task(run_delete())
await lock_acquired.wait()
await asyncio.sleep(0.2)
assert not task.done(), "member_delete did not wait on the team's advisory lock"
await _await_lock_contention(db, lock_key, task, "member_delete")
# member_add wins the race: it adds `other_user` while holding the lock
await held.litellm_teamtable.update(
@ -265,10 +300,8 @@ async def test_delete_blocked_by_member_add_sweeps_the_fresh_reference():
blocker = Prisma()
await blocker.connect()
lock_acquired = asyncio.Event()
async def run_delete():
lock_acquired.set()
return await delete_team(
data=DeleteTeamRequest(team_ids=[team_id]),
http_request=MagicMock(),
@ -278,11 +311,9 @@ async def test_delete_blocked_by_member_add_sweeps_the_fresh_reference():
try:
async with blocker.tx(timeout=timedelta(seconds=30)) as held:
await held.query_raw(_LOCK_SQL, team_id)
lock_key = await _hold_team_lock(held, team_id)
task = asyncio.create_task(run_delete())
await lock_acquired.wait()
await asyncio.sleep(0.3)
assert not task.done(), "delete_team did not wait on the team's advisory lock"
await _await_lock_contention(db, lock_key, task, "delete_team")
# member_add wins the race: write the reference while holding the lock
await held.litellm_usertable.upsert(

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

@ -97,22 +97,30 @@ 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("regional_host", ("eu.api.openai.com", "us.api.openai.com"))
def test_regional_openai_api_base_allows(
self, wif_env: OpenAIWorkloadIdentityConfig, regional_host: str
) -> None:
assert (
resolve_openai_workload_identity_config(api_key=None, api_base=f"https://{regional_host}/v1") == wif_env
)
@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(
"lookalike_base",
("https://api.openai.com.evil.example/v1", "https://openai.com/v1", "https://euapi.openai.com/v1"),
"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_openai_lookalike_api_base_disables(
self, wif_env: OpenAIWorkloadIdentityConfig, lookalike_base: str
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=lookalike_base) is 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
@ -187,6 +195,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)
@ -260,6 +276,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

@ -3508,6 +3508,169 @@ def _create_oauth2_server(
)
def _create_id_lookup_oauth2_server():
return _create_oauth2_server(
server_id="oauth-server-id",
name="oauth-server-name",
server_name="oauth-server-name",
alias="oauth-server-alias",
)
@pytest.mark.asyncio
async def test_authorize_resolves_server_by_id_when_name_lookup_fails():
from fastapi import Request
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
server = _create_id_lookup_oauth2_server()
request = MagicMock(spec=Request)
request.base_url = "https://llm.example.com/"
request.headers = {}
with (
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
patch.object(discoverable_endpoints, "encrypt_value_helper", return_value="encrypted-state"), # test-quality-ok: flow seam
):
response = await discoverable_endpoints.authorize(
request=request,
client_id=server.client_id,
mcp_server_name=server.server_id,
redirect_uri="http://localhost:62646/callback",
state="test_state",
)
assert response.status_code == 307
assert "https://provider.com/oauth/authorize" in response.headers["location"]
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
@pytest.mark.asyncio
async def test_token_endpoint_resolves_server_by_id_when_name_lookup_fails():
from fastapi import Request
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
server = _create_id_lookup_oauth2_server()
request = MagicMock(spec=Request)
request.base_url = "https://llm.example.com/"
request.headers = {}
response = MagicMock()
response.json.return_value = {"access_token": "token", "token_type": "Bearer"}
response.raise_for_status = MagicMock()
client = MagicMock()
client.post = AsyncMock(return_value=response)
with (
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
patch.object(discoverable_endpoints, "get_async_httpx_client", return_value=client), # test-quality-ok: HTTP seam
):
result = await discoverable_endpoints.token_endpoint(
request=request,
grant_type="authorization_code",
code="test_code",
redirect_uri="http://localhost:62646/callback",
client_id=server.client_id,
mcp_server_name=server.server_id,
client_secret=server.client_secret,
)
assert json.loads(result.body)["access_token"] == "token"
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
@pytest.mark.asyncio
async def test_register_client_resolves_server_by_id_when_name_lookup_fails():
from fastapi import Request
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
server = _create_id_lookup_oauth2_server().model_copy(
update={"client_id": None, "client_secret": None, "registration_url": "https://provider.com/oauth/register"}
)
request = MagicMock(spec=Request)
request.base_url = "https://llm.example.com/"
request.headers = {}
response = MagicMock()
response.json.return_value = {"client_id": "registered-client", "client_secret": "registered-secret"}
response.raise_for_status = MagicMock()
client = MagicMock()
client.post = AsyncMock(return_value=response)
with (
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
patch.object(discoverable_endpoints, "_read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: request seam
patch.object(discoverable_endpoints, "get_async_httpx_client", return_value=client), # test-quality-ok: HTTP seam
):
result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id)
assert json.loads(result.body)["client_id"] == "registered-client"
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
@pytest.mark.asyncio
async def test_protected_resource_metadata_resolves_server_by_id_when_name_lookup_fails():
from fastapi import Request
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
server = _create_id_lookup_oauth2_server()
request = MagicMock(spec=Request)
request.base_url = "https://llm.example.com/"
request.headers = {}
with (
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
):
result = await discoverable_endpoints._build_oauth_protected_resource_response(
request=request,
mcp_server_name=server.server_id,
use_standard_pattern=True,
)
assert result["authorization_servers"] == ["https://llm.example.com/mcp"]
assert result["resource"] == f"https://llm.example.com/mcp/{server.server_id}"
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fails():
from fastapi import Request
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
server = _create_id_lookup_oauth2_server()
request = MagicMock(spec=Request)
request.base_url = "https://llm.example.com/"
request.headers = {}
with (
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
):
result = discoverable_endpoints._build_oauth_authorization_server_response(
request=request,
mcp_server_name=server.server_id,
)
assert result["scopes_supported"] == server.scopes
assert result["issuer"] == f"https://llm.example.com/{server.server_id}"
by_name.assert_called_once_with(server.server_id, client_ip=None)
by_id.assert_called_once_with(server.server_id, client_ip=None)
@pytest.mark.asyncio
async def test_authorize_root_resolves_single_oauth2_server():
"""When /authorize is hit without server name and exactly 1 OAuth2 server exists, resolve it."""

View file

@ -129,6 +129,32 @@ class TestMCPServerManager:
assert added_server.args == ["-m", "server"]
assert added_server.env == {"DEBUG": "1", "TEST": "1"}
def test_get_mcp_server_by_id_allows_internal_or_unspecified_client_ip(self):
manager = MCPServerManager()
server = MCPServer(
server_id="private-server",
name="private-server",
transport=MCPTransport.http,
available_on_public_internet=False,
)
manager.registry[server.server_id] = server
assert manager.get_mcp_server_by_id(server.server_id) is server
assert manager.get_mcp_server_by_id(server.server_id, client_ip="10.0.0.1") is server
def test_get_mcp_server_by_id_rejects_private_server_for_public_ip(self):
manager = MCPServerManager()
server = MCPServer(
server_id="private-server",
name="private-server",
transport=MCPTransport.http,
available_on_public_internet=False,
)
manager.registry[server.server_id] = server
with patch.object(manager, "_get_general_settings", return_value={}):
assert manager.get_mcp_server_by_id(server.server_id, client_ip="8.8.8.8") is None
async def test_create_mcp_client_stdio(self):
"""Test creating MCP client for stdio transport"""
manager = MCPServerManager()
@ -2590,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."""
@ -4545,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

@ -10,15 +10,42 @@ from unittest.mock import AsyncMock, patch
import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
AgentAccess,
AgentRequestHandler,
RestrictedAgentAccess,
UnrestrictedAgentAccess,
accessible_agents,
)
def _registry_with(*agent_names: str) -> AgentRegistry:
registry: Final = AgentRegistry()
registry.load_agents_from_config(
[
{
"agent_name": name,
"agent_card_params": {"name": name, "url": "http://localhost", "version": "1.0.0"},
}
for name in agent_names
]
)
return registry
def _agent_id(registry: AgentRegistry, agent_name: str) -> str:
agent: Final = registry.get_agent_by_name(agent_name)
assert agent is not None
return agent.agent_id
async def _single_context(user_api_key_auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
return [user_api_key_auth]
@pytest.mark.asyncio
class TestAgentRequestHandler:
"""
@ -265,6 +292,78 @@ class TestAgentRequestHandler:
)
assert result == UnrestrictedAgentAccess()
async def test_accessible_agents_hides_ungranted_agents_from_non_admins(self):
"""LIT-6862: a key with no agent grant on itself or its team must list nothing,
while a proxy admin with the same lack of grants still lists every agent."""
registry: Final = _registry_with("alpha", "beta")
internal_user: Final = UserAPIKeyAuth(
api_key="test-key", user_id="alice", team_id="team-no-perms", user_role=LitellmUserRoles.INTERNAL_USER
)
proxy_admin: Final = UserAPIKeyAuth(
api_key="admin-key", user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
)
async def no_grant_anywhere(user_api_key_auth: UserAPIKeyAuth) -> AgentAccess:
return UnrestrictedAgentAccess()
assert (
await accessible_agents(internal_user, registry.get_agent_list(), no_grant_anywhere, _single_context) == ()
)
assert {
agent.agent_name
for agent in await accessible_agents(
proxy_admin, registry.get_agent_list(), no_grant_anywhere, _single_context
)
} == {"alpha", "beta"}
async def test_accessible_agents_lists_only_granted_agents(self):
"""A grant for one agent lists that agent and hides the ungranted one."""
registry: Final = _registry_with("alpha", "beta")
granted_user: Final = UserAPIKeyAuth(
api_key="test-key", user_id="bob", team_id="team-granted", user_role=LitellmUserRoles.INTERNAL_USER
)
async def alpha_only(user_api_key_auth: UserAPIKeyAuth) -> AgentAccess:
return RestrictedAgentAccess(frozenset({_agent_id(registry, "alpha")}))
listed: Final = await accessible_agents(granted_user, registry.get_agent_list(), alpha_only, _single_context)
assert [agent.agent_name for agent in listed] == ["alpha"]
async def test_accessible_agents_resolves_dashboard_session_through_real_teams_and_user(self):
"""LIT-6862: a dashboard session carries the shared litellm-dashboard team id, which holds no
grants. Listing must union the grants of the user's real teams and of the user row instead
of treating the session as ungranted or as unrestricted."""
registry: Final = _registry_with("alpha", "beta", "gamma")
session: Final = UserAPIKeyAuth(
api_key="session-key",
user_id="alice",
team_id=UI_SESSION_TOKEN_TEAM_ID,
user_role=LitellmUserRoles.INTERNAL_USER,
)
admitted_user: Final = UserAPIKeyAuth(user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER)
grants: Final = {
"team-granted": RestrictedAgentAccess(frozenset({_agent_id(registry, "alpha")})),
"team-no-perms": UnrestrictedAgentAccess(),
UI_SESSION_TOKEN_TEAM_ID: UnrestrictedAgentAccess(),
}
async def effective_contexts(user_api_key_auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
assert user_api_key_auth is session
return [
session.model_copy(update={"team_id": "team-granted"}),
session.model_copy(update={"team_id": "team-no-perms"}),
admitted_user,
]
async def resolve_access(user_api_key_auth: UserAPIKeyAuth) -> AgentAccess:
if user_api_key_auth is admitted_user:
return RestrictedAgentAccess(frozenset({_agent_id(registry, "beta")}))
assert user_api_key_auth.team_id is not None
return grants[user_api_key_auth.team_id]
listed: Final = await accessible_agents(session, registry.get_agent_list(), resolve_access, effective_contexts)
assert {agent.agent_name for agent in listed} == {"alpha", "beta"}
async def test_get_allowed_agents_for_key_via_access_group_ids(self):
"""
Test that _get_allowed_agents_for_key includes agents from key's access_group_ids

View file

@ -11,7 +11,6 @@ from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.agent_endpoints import endpoints as agent_endpoints
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
RestrictedAgentAccess,
UnrestrictedAgentAccess,
)
from litellm.proxy.agent_endpoints.endpoints import (
_attach_keys_to_agents,
@ -550,9 +549,9 @@ class TestAgentRBACProxyAdminViewOnly:
self.allowed_agents_spy.assert_awaited_once()
def test_should_still_redact_secrets_for_view_only_admin(self):
"""An unrestricted viewer sees the same agents as an admin but with keys
"""A viewer granted every agent sees the same agents as an admin but with keys
stripped; litellm_params secrets never appear in either response."""
self.allowed_agents_spy.return_value = UnrestrictedAgentAccess()
self.allowed_agents_spy.return_value = RestrictedAgentAccess(frozenset({"agent-1", "agent-2"}))
viewer_resp = self._list_agents(self.viewer_client)
admin_resp = self._list_agents(self.admin_client)

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

@ -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

@ -2910,6 +2910,7 @@ async def test_update_team_with_team_member_budget_duration(
"metadata": {"team_member_budget_id": "budget_123"},
}
mock_existing_team.metadata = {"team_member_budget_id": "budget_123"}
mock_existing_team.members_with_roles = []
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(
return_value=mock_existing_team
)
@ -11290,6 +11291,78 @@ async def test_patch_preserves_required_metadata_key_that_post_would_wipe():
assert patch_meta == {"cost_center": "FINOPS-1", "team_notes": "edited"} # preserved by PATCH
_STORED_METADATA_WITH_BUDGET: Final = {
"team_member_budget_id": "budget-existing-123",
"team_member_key_duration": "30d",
"logging": [{"callback_name": "langfuse", "callback_type": "success"}],
"cost_center": "cc-1234",
}
async def _written_metadata_with_budget(kind, body):
"""Like ``_written_metadata`` but the team already owns a member budget row."""
from litellm.proxy._types import LiteLLM_BudgetTable
with (
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy_server module global is the endpoint's only injection point
patch( # test-quality-ok: update_team imports update_budget at call time; the module attribute is its only seam
"litellm.proxy.management_endpoints.budget_management_endpoints.update_budget",
AsyncMock(return_value=LiteLLM_BudgetTable(budget_id="budget-existing-123")),
),
):
return await _written_metadata(kind, dict(_STORED_METADATA_WITH_BUDGET), body)
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["post", "patch"])
@pytest.mark.parametrize(
"body",
[
{"team_member_budget": 50.0},
{"team_member_budget_duration": "1d"},
{"team_member_tpm_limit": 500},
{"team_member_rpm_limit": 5},
],
ids=lambda body: next(iter(body)),
)
async def test_team_member_budget_only_update_preserves_stored_metadata(kind, body):
"""LIT-5150: a budget-only update that omits ``metadata`` must not replace the
stored metadata JSON with just ``{"team_member_budget_id": ...}``."""
assert await _written_metadata_with_budget(kind, body) == _STORED_METADATA_WITH_BUDGET
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["post", "patch"])
async def test_team_member_key_duration_only_update_preserves_stored_metadata(kind):
"""LIT-5150: a metadata-backed field sent alone is merged into the stored
metadata instead of becoming the whole metadata JSON."""
written = await _written_metadata_with_budget(kind, {"team_member_key_duration": "7d"})
assert written == {**_STORED_METADATA_WITH_BUDGET, "team_member_key_duration": "7d"}
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["post", "patch"])
async def test_explicit_null_metadata_with_budget_field_still_clears_metadata(kind):
"""``metadata: null`` is an explicit clear, so only the server-owned budget link survives."""
written = await _written_metadata_with_budget(kind, {"metadata": None, "team_member_budget": 7.0})
assert written == {"team_member_budget_id": "budget-existing-123"}
@pytest.mark.asyncio
async def test_metadata_only_update_keeps_team_member_budget_link():
"""LIT-5150: rewriting metadata without any team member field must not drop the
server-owned ``team_member_budget_id``, or the member budget silently resets."""
body = {"metadata": {"cost_center": "cc-9999"}}
post_meta = await _written_metadata_with_budget("post", body)
patch_meta = await _written_metadata_with_budget("patch", body)
assert post_meta == {"cost_center": "cc-9999", "team_member_budget_id": "budget-existing-123"}
assert patch_meta == {**_STORED_METADATA_WITH_BUDGET, "cost_center": "cc-9999"}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"body, field, expected",
@ -11318,17 +11391,15 @@ async def test_top_level_fields_identical_post_and_patch(body, field, expected):
@pytest.mark.asyncio
async def test_patch_strips_system_managed_metadata_key_like_post():
"""A caller cannot inject/overwrite server-owned keys via PATCH any more than
via POST: team_member_budget_id is stripped from the write in both."""
via POST: the stored team_member_budget_id wins over the caller's value in both."""
existing = {"team_member_budget_id": "budget-123", "cost_center": "1234"}
body = {"metadata": {"team_member_budget_id": "HACKED", "cost_center": "9999"}}
post_meta = await _written_metadata("post", existing, body)
patch_meta = await _written_metadata("patch", existing, body)
assert "team_member_budget_id" not in post_meta
assert "team_member_budget_id" not in patch_meta
assert post_meta == {"cost_center": "9999"}
assert patch_meta == {"cost_center": "9999"}
assert post_meta == {"cost_center": "9999", "team_member_budget_id": "budget-123"}
assert patch_meta == {"cost_center": "9999", "team_member_budget_id": "budget-123"}
@pytest.mark.parametrize(

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

@ -14,7 +14,9 @@ from litellm.proxy.hooks.responses_id_security import (
_is_responses_api_create_route,
)
from litellm.types.llms.openai import (
GenericEvent,
ResponseCompletedEvent,
ResponseCreatedEvent,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
@ -691,6 +693,126 @@ class TestAsyncPostCallStreamingIteratorHook:
assert not responses_id_security._is_encrypted_response_id(streamed_id)
class TestStreamedGenericEventIdEncryption:
"""A background stream carries event types with no typed model, which arrive as
GenericEvent holding a plain dict. Those used to skip encryption while their typed
siblings were encrypted, so one stream advertised two ids and the unencrypted one
skipped the ownership check. Asserts the property rather than one event type: every
id a client can see is the same encrypted id, and the raw one appears in no frame."""
RAW_ID = "resp_rawprovider123"
@staticmethod
async def _agen(chunks):
for chunk in chunks:
yield chunk
@classmethod
def _typed_event(cls, event_type):
return {
ResponsesAPIStreamEvents.RESPONSE_CREATED: ResponseCreatedEvent,
ResponsesAPIStreamEvents.RESPONSE_COMPLETED: ResponseCompletedEvent,
}[event_type](
type=event_type,
response=ResponsesAPIResponse(
id=cls.RAW_ID,
created_at=0,
model="gpt-5.1",
object="response",
output=[],
parallel_tool_calls=False,
tool_choice="auto",
tools=[],
),
)
@classmethod
def _background_stream(cls):
return [
cls._typed_event(ResponsesAPIStreamEvents.RESPONSE_CREATED),
GenericEvent(
type="response.queued",
response={"id": cls.RAW_ID, "status": "queued"},
),
GenericEvent(type="keepalive"),
GenericEvent(
type="response.some_event_openai_adds_later",
response={"id": cls.RAW_ID, "status": "in_progress"},
),
cls._typed_event(ResponsesAPIStreamEvents.RESPONSE_COMPLETED),
]
@staticmethod
def _advertised_ids(events):
nested = (getattr(event, "response", None) for event in events)
return [
payload["id"] if isinstance(payload, dict) else payload.id
for payload in nested
if payload is not None
] + [
event.id for event in events if isinstance(getattr(event, "id", None), str)
]
async def _drain(self, responses_id_security, monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-abcdefghij")
mock_auth = MagicMock()
mock_auth.user_id = "user-a"
mock_auth.team_id = "team-a"
mock_auth.request_route = "/v1/responses"
return [
out
async for out in responses_id_security.async_post_call_streaming_iterator_hook(
user_api_key_dict=mock_auth,
response=self._agen(self._background_stream()),
request_data={},
)
]
@pytest.mark.asyncio
async def test_every_event_advertises_the_same_encrypted_id(
self, responses_id_security, monkeypatch
):
events = await self._drain(responses_id_security, monkeypatch)
advertised = self._advertised_ids(events)
assert len(advertised) == 4
assert len(set(advertised)) == 1
streamed_id = advertised[0]
assert streamed_id != self.RAW_ID
assert responses_id_security._is_encrypted_response_id(streamed_id)
assert responses_id_security._decrypt_response_id(streamed_id) == (
self.RAW_ID,
"user-a",
"team-a",
)
@pytest.mark.asyncio
async def test_raw_provider_id_never_reaches_the_client(
self, responses_id_security, monkeypatch
):
events = await self._drain(responses_id_security, monkeypatch)
assert [self.RAW_ID in event.model_dump_json() for event in events] == [
False
] * len(events)
@pytest.mark.asyncio
async def test_sibling_fields_survive_the_rewrite(
self, responses_id_security, monkeypatch
):
_, queued, keepalive, later, _ = await self._drain(
responses_id_security, monkeypatch
)
assert queued.response["status"] == "queued"
assert later.response["status"] == "in_progress"
assert keepalive.type == "keepalive"
assert getattr(keepalive, "response", None) is None
class TestAsyncPostCallSuccessHook:
"""Test async_post_call_success_hook function"""

View file

@ -7392,6 +7392,71 @@ def test_get_configured_token_limits_coerces_numeric_strings():
assert router.get_configured_token_limits("quoted-limits-model") == (32000, 8000)
def test_get_configured_mode_reads_deployment_model_info():
router = litellm.Router(
model_list=[
{
"model_name": "tts-model",
"litellm_params": {"model": "openai/some-unmapped-tts-model"},
"model_info": {"mode": "audio_speech"},
}
]
)
assert router.get_configured_mode("tts-model") == "audio_speech"
def test_get_configured_mode_returns_none_for_unset_or_unknown():
router = litellm.Router(
model_list=[
{
"model_name": "no-mode-model",
"litellm_params": {"model": "openai/some-unmapped-model"},
}
]
)
assert router.get_configured_mode("no-mode-model") is None
assert router.get_configured_mode("not-a-real-model") is None
def test_get_configured_mode_skips_wildcard_pattern_matching():
router = litellm.Router(
model_list=[
{
"model_name": "bedrock/*",
"litellm_params": {"model": "bedrock/*"},
"model_info": {"mode": "chat"},
}
]
)
with patch.object(
router.pattern_router, "route", side_effect=AssertionError("pattern route called")
):
assert (
router.get_configured_mode("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0")
is None
)
def test_get_configured_mode_treats_malformed_values_as_absent():
malformed = ["", " ", 12345, ["chat"], {"mode": "chat"}, True]
router = litellm.Router(
model_list=[
{
"model_name": f"bad-mode-{i}",
"litellm_params": {"model": "openai/some-unmapped-model"},
"model_info": {"mode": bad},
}
for i, bad in enumerate(malformed)
]
)
for i in range(len(malformed)):
assert router.get_configured_mode(f"bad-mode-{i}") is None
def test_get_configured_display_name_reads_deployment_model_info():
router = litellm.Router(
model_list=[
@ -12695,33 +12760,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

@ -27,10 +27,10 @@
"limit": 0
},
"LIT010": {
"limit": 16476
"limit": 16470
},
"LIT011": {
"limit": 5518
"limit": 5516
},
"LIT012": {
"limit": 4489

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

@ -2134,6 +2134,7 @@ interface UiSpendLogsParams {
exclude_internal_health_checks?: boolean;
group_by_session?: boolean;
session_cursor?: string;
search?: string;
}
interface UiSpendLogsCallOptions {
@ -6641,6 +6642,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";
}
@ -8139,15 +8141,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

@ -1538,6 +1538,35 @@ describe("TeamInfoView", () => {
});
});
describe("team member settings", () => {
it("should populate Default Key Duration from the team's stored metadata", async () => {
const user = userEvent.setup({ delay: null });
vi.mocked(networking.teamInfoCall).mockResolvedValue(
createMockTeamData({ metadata: { team_member_key_duration: "30d" } }),
);
renderWithProviders(<TeamInfoView {...defaultProps} />);
await waitFor(() => {
expect(screen.queryAllByText("Test Team").length).toBeGreaterThan(0);
});
await user.click(screen.getByRole("tab", { name: "Settings" }));
await waitFor(() => {
expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument();
});
await user.click(screen.getByRole("button", { name: /edit settings/i }));
await user.click(await screen.findByRole("button", { name: /team member settings/i }));
await waitFor(() => {
expect(screen.getByLabelText(/^Default Key Duration/)).toHaveValue("30d");
});
});
});
describe("guardrails dropdown grouping", () => {
const guardrail = (name: string, defaultOn: boolean) => ({
guardrail_name: name,

View file

@ -335,7 +335,7 @@ const toTeamFormValues = (info: TeamInfoRecord, effectiveGuardrails: string[]):
default_team_member_models: info.default_team_member_models || [],
team_member_budget: info.team_member_budget_table?.max_budget,
team_member_budget_duration: info.team_member_budget_table?.budget_duration,
team_member_key_duration: info.team_member_key_duration,
team_member_key_duration: info.metadata?.team_member_key_duration,
team_member_tpm_limit: info.team_member_budget_table?.tpm_limit,
team_member_rpm_limit: info.team_member_budget_table?.rpm_limit,
budget_duration: info.budget_duration,

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

@ -41118,6 +41118,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') */
@ -49868,6 +49870,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. */
@ -57159,6 +57163,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;
@ -57275,6 +57281,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;
@ -63587,6 +63595,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;
};