Merge remote-tracking branch 'berri/litellm_internal_staging' into litellm_lit6899_vertex_batch_tuned_endpoints

This commit is contained in:
mubashir1osmani 2026-09-03 19:29:20 -04:00
commit 27689c5919
48 changed files with 2180 additions and 222 deletions

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

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

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

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

View file

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

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

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

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

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -7362,14 +7362,14 @@ def test_get_configured_mode_reads_deployment_model_info():
router = litellm.Router(
model_list=[
{
"model_name": "chat-model",
"litellm_params": {"model": "openai/some-unmapped-model"},
"model_info": {"mode": "chat"},
"model_name": "tts-model",
"litellm_params": {"model": "openai/some-unmapped-tts-model"},
"model_info": {"mode": "audio_speech"},
}
]
)
assert router.get_configured_mode("chat-model") == "chat"
assert router.get_configured_mode("tts-model") == "audio_speech"
def test_get_configured_mode_returns_none_for_unset_or_unknown():
@ -12701,33 +12701,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

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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