mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_internal_copy_38013
# Conflicts: # litellm/llms/openai/workload_identity.py # tests/test_litellm/llms/openai/test_openai_workload_identity.py
This commit is contained in:
commit
37f2de1b0e
78 changed files with 3802 additions and 388 deletions
|
|
@ -55,7 +55,10 @@ RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/
|
|||
|
||||
ENV PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries
|
||||
|
||||
COPY docker/prod_entrypoint.sh /app/docker/prod_entrypoint.sh
|
||||
RUN sed -i 's/\r$//' /app/docker/prod_entrypoint.sh && chmod +x /app/docker/prod_entrypoint.sh
|
||||
|
||||
EXPOSE 4000/tcp
|
||||
|
||||
ENTRYPOINT ["litellm"]
|
||||
ENTRYPOINT ["/app/docker/prod_entrypoint.sh"]
|
||||
CMD ["--port", "4000"]
|
||||
|
|
|
|||
|
|
@ -1,8 +1,10 @@
|
|||
#!/bin/sh
|
||||
|
||||
if [ "$USE_DDTRACE" = "true" ]; then
|
||||
export DD_TRACE_OPENAI_ENABLED="False"
|
||||
exec ddtrace-run "$@"
|
||||
fi
|
||||
case "$USE_DDTRACE" in
|
||||
[Tt][Rr][Uu][Ee])
|
||||
export DD_TRACE_OPENAI_ENABLED="False"
|
||||
exec ddtrace-run "$@"
|
||||
;;
|
||||
esac
|
||||
|
||||
exec "$@"
|
||||
|
|
|
|||
|
|
@ -1,8 +1,10 @@
|
|||
#!/bin/sh
|
||||
|
||||
if [ "$USE_DDTRACE" = "true" ]; then
|
||||
export DD_TRACE_OPENAI_ENABLED="False"
|
||||
exec ddtrace-run litellm "$@"
|
||||
else
|
||||
exec litellm "$@"
|
||||
fi
|
||||
case "$USE_DDTRACE" in
|
||||
[Tt][Rr][Uu][Ee])
|
||||
export DD_TRACE_OPENAI_ENABLED="False"
|
||||
exec ddtrace-run litellm "$@"
|
||||
;;
|
||||
esac
|
||||
|
||||
exec litellm "$@"
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -50,10 +50,12 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
MAX_FILE_LIST_LIMIT,
|
||||
_is_base64_encoded_unified_file_id,
|
||||
apply_unified_file_ids,
|
||||
decode_model_from_file_id,
|
||||
ensure_batch_response_managed_file_ids,
|
||||
get_batch_id_from_unified_batch_id,
|
||||
get_content_type_from_file_object,
|
||||
get_model_id_from_unified_batch_id,
|
||||
get_original_file_id,
|
||||
map_raw_file_ids_to_unified,
|
||||
normalize_mime_type_for_provider,
|
||||
resolve_managed_output_file_model_name,
|
||||
|
|
@ -427,6 +429,103 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
detail=f"Object not found: {unified_object_id}",
|
||||
)
|
||||
|
||||
async def enforce_batch_object_access(
|
||||
self, object_id: str, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> None:
|
||||
"""Deny access to a provider-format batch id owned by another caller.
|
||||
|
||||
Ids with no ownership row (batches created before ownership tracking,
|
||||
or directly on the provider account) stay accessible so pass-through
|
||||
reads keep working.
|
||||
"""
|
||||
if self.prisma_client is None:
|
||||
return
|
||||
managed_object = (
|
||||
await self.prisma_client.db.litellm_managedobjecttable.find_first(
|
||||
where={"OR": [{"unified_object_id": object_id}, {"model_object_id": object_id}]}
|
||||
)
|
||||
)
|
||||
if managed_object is None:
|
||||
return
|
||||
if not can_access_resource(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
created_by=managed_object.created_by,
|
||||
resource_team_id=managed_object.team_id,
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"User {user_api_key_dict.user_id} does not have access to the object {object_id}",
|
||||
)
|
||||
|
||||
async def enforce_provider_file_access(
|
||||
self, file_id: str, user_api_key_dict: UserAPIKeyAuth
|
||||
) -> None:
|
||||
"""Deny access to a provider-format file id owned by another caller.
|
||||
|
||||
Ownership rows for provider-format ids are written when a managed
|
||||
batch's output/error files are first synced; ids with no row stay
|
||||
accessible so pass-through reads keep working.
|
||||
"""
|
||||
if self.prisma_client is None:
|
||||
return
|
||||
managed_file = (
|
||||
await self.prisma_client.db.litellm_managedfiletable.find_first(
|
||||
where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]}
|
||||
)
|
||||
)
|
||||
if managed_file is None:
|
||||
return
|
||||
if not can_access_resource(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
created_by=managed_file.created_by,
|
||||
resource_team_id=managed_file.team_id,
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"User {user_api_key_dict.user_id} does not have access to the file {file_id}",
|
||||
)
|
||||
|
||||
async def store_batch_output_file_ownership(
|
||||
self, response: LiteLLMBatch, litellm_parent_otel_span: Optional[Span]
|
||||
) -> None:
|
||||
"""Record ownership rows for a batch's provider-format output/error
|
||||
file ids, inherited from the owning batch row (never the caller), so
|
||||
file reads can be isolation-checked."""
|
||||
provider_file_ids = tuple(
|
||||
file_id
|
||||
for file_id in (
|
||||
getattr(response, "output_file_id", None),
|
||||
getattr(response, "error_file_id", None),
|
||||
)
|
||||
if file_id and not _is_base64_encoded_unified_file_id(file_id)
|
||||
)
|
||||
if not provider_file_ids:
|
||||
return
|
||||
if self.prisma_client is None:
|
||||
return
|
||||
batch_row = (
|
||||
await self.prisma_client.db.litellm_managedobjecttable.find_first(
|
||||
where={"unified_object_id": response.id}
|
||||
)
|
||||
)
|
||||
if batch_row is None or (
|
||||
batch_row.created_by is None and batch_row.team_id is None
|
||||
):
|
||||
return
|
||||
owner_identity = UserAPIKeyAuth(
|
||||
user_id=batch_row.created_by, team_id=batch_row.team_id
|
||||
)
|
||||
for file_id in provider_file_ids:
|
||||
model_name = decode_model_from_file_id(file_id)
|
||||
raw_file_id = get_original_file_id(file_id)
|
||||
await self.store_unified_file_id(
|
||||
file_id=file_id,
|
||||
file_object=None,
|
||||
litellm_parent_otel_span=litellm_parent_otel_span,
|
||||
model_mappings={model_name: raw_file_id} if model_name else {},
|
||||
user_api_key_dict=owner_identity,
|
||||
)
|
||||
|
||||
async def list_user_batches(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -613,6 +712,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
status_code=403,
|
||||
detail=f"User {user_api_key_dict.user_id} does not have access to the file {retrieve_file_id}",
|
||||
)
|
||||
if retrieve_file_id:
|
||||
await self.enforce_provider_file_access(
|
||||
retrieve_file_id, user_api_key_dict
|
||||
)
|
||||
return False
|
||||
|
||||
async def check_file_ids_access(self, file_ids: List[str], user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
|
|
@ -765,6 +868,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
)
|
||||
data["model"] = potential_model_id
|
||||
data[accessor_key] = get_batch_id_from_unified_batch_id(potential_llm_object_id)
|
||||
elif retrieve_object_id and accessor_key == "batch_id":
|
||||
await self.enforce_batch_object_access(retrieve_object_id, user_api_key_dict)
|
||||
elif call_type == CallTypes.acreate_fine_tuning_job.value:
|
||||
input_file_id = cast(Optional[str], data.get("training_file"))
|
||||
if input_file_id:
|
||||
|
|
@ -1297,7 +1402,13 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
request_tags=request_tags_from_metadata(request_metadata if isinstance(request_metadata, dict) else {}),
|
||||
persist_attribution=is_batch_create,
|
||||
create_if_missing=is_batch_create,
|
||||
)
|
||||
if not is_batch_create:
|
||||
await self.store_batch_output_file_ownership(
|
||||
response=response,
|
||||
litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
)
|
||||
|
||||
# Only record batch creation metric on actual create (not retrieve/cancel).
|
||||
# unified_file_id in _hidden_params is only set by the create_batch endpoint.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -70,7 +70,7 @@ def to_basic_auth(auth_value: str) -> str:
|
|||
|
||||
|
||||
def strip_auth_scheme(auth_value: str, scheme: str) -> str:
|
||||
"""Return ``auth_value`` with a leading ``<scheme> `` removed, or unchanged when absent.
|
||||
"""Return ``auth_value`` with a leading ``<scheme>`` and separator removed, or unchanged when absent.
|
||||
|
||||
Callers supply both a bare credential and a complete header value, so prefixing
|
||||
unconditionally yields ``Bearer Bearer <jwt>``. Scheme names are case-insensitive per
|
||||
|
|
@ -78,10 +78,9 @@ def strip_auth_scheme(auth_value: str, scheme: str) -> str:
|
|||
with the scheme text and a scheme with nothing behind it are returned untouched.
|
||||
Surrounding whitespace is left to ``_strip_header_whitespace`` at header-build time.
|
||||
"""
|
||||
scheme_name, _, remainder = auth_value.lstrip().partition(" ")
|
||||
credential: Final = remainder.lstrip()
|
||||
if credential and scheme_name.lower() == scheme.lower():
|
||||
return credential
|
||||
parts: Final = auth_value.split(None, 1)
|
||||
if len(parts) == 2 and parts[0].lower() == scheme.lower():
|
||||
return parts[1]
|
||||
return auth_value
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ from urllib.parse import urlparse
|
|||
import litellm
|
||||
from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str
|
||||
|
||||
from .common_utils import OpenAIError
|
||||
from .common_utils import OpenAIError, is_openai_backed_api_base
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
|
@ -17,8 +17,6 @@ if TYPE_CHECKING:
|
|||
from openai.auth import SubjectTokenProvider, WorkloadIdentity, WorkloadIdentityAuth
|
||||
|
||||
OPENAI_WIF_CLIENT_ID: Final = "litellm"
|
||||
_OPENAI_API_HOST: Final = "api.openai.com"
|
||||
_OPENAI_REGIONAL_HOST_SUFFIX: Final = f".{_OPENAI_API_HOST}"
|
||||
_SDK_UPGRADE_MESSAGE: Final = (
|
||||
"OpenAI workload identity federation requires openai>=2.32.0. "
|
||||
"Upgrade the installed openai package to use OPENAI_IDENTITY_PROVIDER_ID / "
|
||||
|
|
@ -87,9 +85,7 @@ def _targets_openai_api(api_base: str | None) -> bool:
|
|||
if api_base is None:
|
||||
return True
|
||||
parsed: Final = urlparse(api_base)
|
||||
if parsed.scheme != "https" or parsed.hostname is None:
|
||||
return False
|
||||
return parsed.hostname == _OPENAI_API_HOST or parsed.hostname.endswith(_OPENAI_REGIONAL_HOST_SUFFIX)
|
||||
return parsed.scheme == "https" and is_openai_backed_api_base(api_base)
|
||||
|
||||
|
||||
@lru_cache(maxsize=16)
|
||||
|
|
|
|||
|
|
@ -448,6 +448,17 @@ def _append_query_params(url: str, params: dict[str, str]) -> str:
|
|||
return urlunparse(parsed._replace(query=urlencode(query_params)))
|
||||
|
||||
|
||||
def _resolve_mcp_server_by_name_or_id(lookup: str, client_ip: str | None) -> MCPServer | None:
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
by_name: Final = global_mcp_server_manager.get_mcp_server_by_name(lookup, client_ip=client_ip)
|
||||
if by_name is not None:
|
||||
return by_name
|
||||
return global_mcp_server_manager.get_mcp_server_by_id(lookup, client_ip=client_ip)
|
||||
|
||||
|
||||
def _resolve_oauth2_server_for_root_endpoints(
|
||||
client_ip: str | None = None,
|
||||
) -> MCPServer | None:
|
||||
|
|
@ -1766,10 +1777,6 @@ async def authorize(
|
|||
resource: str | None = None,
|
||||
):
|
||||
# Redirect to real OAuth provider with PKCE support
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
if mcp_server_name is None and client_id and is_gateway_dcr_client_id(client_id):
|
||||
if is_proxy_api_resource(request, resource):
|
||||
return await native_client_authorize(
|
||||
|
|
@ -1797,9 +1804,7 @@ async def authorize(
|
|||
|
||||
lookup_name: Final[str | None] = mcp_server_name or client_id
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
mcp_server = (
|
||||
global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip) if lookup_name else None
|
||||
)
|
||||
mcp_server = _resolve_mcp_server_by_name_or_id(lookup_name, client_ip) if lookup_name else None
|
||||
if mcp_server is None and mcp_server_name is None:
|
||||
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
|
||||
if mcp_server is None:
|
||||
|
|
@ -1855,10 +1860,6 @@ async def token_endpoint(
|
|||
3. Return the token
|
||||
4. Return a virtual key in this response
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
if mcp_server_name is None and is_gateway_dcr_client_id(client_id):
|
||||
from litellm.proxy.proxy_server import ( # noqa: PLC0415 # circular import at module load
|
||||
master_key,
|
||||
|
|
@ -1882,7 +1883,7 @@ async def token_endpoint(
|
|||
|
||||
lookup_name: Final = mcp_server_name or client_id
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(lookup_name, client_ip=client_ip)
|
||||
mcp_server = _resolve_mcp_server_by_name_or_id(lookup_name, client_ip)
|
||||
if mcp_server is None and mcp_server_name is None:
|
||||
mcp_server = _resolve_oauth2_server_for_root_endpoints(client_ip=client_ip)
|
||||
if mcp_server is None:
|
||||
|
|
@ -2288,10 +2289,6 @@ async def _build_oauth_protected_resource_response(
|
|||
Returns:
|
||||
OAuth protected resource metadata dict
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
explicitly_named: Final = mcp_server_name is not None
|
||||
|
|
@ -2304,7 +2301,7 @@ async def _build_oauth_protected_resource_response(
|
|||
|
||||
mcp_server: MCPServer | None = None
|
||||
if mcp_server_name:
|
||||
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
|
||||
mcp_server = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
|
||||
|
||||
# Build resource URL based on the pattern
|
||||
if mcp_server_name:
|
||||
|
|
@ -2562,10 +2559,6 @@ def _build_oauth_authorization_server_response(
|
|||
registry lookups; unlike :func:`_build_oauth_protected_resource_response`
|
||||
it does not need to await any upstream IO.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(request)
|
||||
explicitly_named: Final = mcp_server_name is not None
|
||||
|
|
@ -2583,7 +2576,7 @@ def _build_oauth_authorization_server_response(
|
|||
|
||||
mcp_server: MCPServer | None = None
|
||||
if mcp_server_name:
|
||||
mcp_server = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
|
||||
mcp_server = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
|
||||
|
||||
_raise_unless_oauth2_discovery_server(mcp_server, mcp_server_name, "not an OAuth authorization server")
|
||||
|
||||
|
|
@ -2709,10 +2702,6 @@ async def oauth_authorization_server_legacy(request: Request, mcp_server_name: s
|
|||
@router.post("/{mcp_server_name}/register")
|
||||
@router.post("/register")
|
||||
async def register_client(request: Request, mcp_server_name: str | None = None):
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
# Get the correct base URL considering X-Forwarded-* headers
|
||||
request_base_url: Final = get_request_base_url(request)
|
||||
|
||||
|
|
@ -2748,7 +2737,7 @@ async def register_client(request: Request, mcp_server_name: str | None = None):
|
|||
)
|
||||
return dummy_return
|
||||
|
||||
mcp_server: Final = global_mcp_server_manager.get_mcp_server_by_name(mcp_server_name, client_ip=client_ip)
|
||||
mcp_server: Final = _resolve_mcp_server_by_name_or_id(mcp_server_name, client_ip)
|
||||
if mcp_server is None:
|
||||
return dummy_return
|
||||
return await register_client_with_server(
|
||||
|
|
|
|||
|
|
@ -47,7 +47,7 @@ from litellm.constants import (
|
|||
MCP_TOOL_LISTING_TIMEOUT,
|
||||
)
|
||||
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
|
||||
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme
|
||||
from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme, to_basic_credentials
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
_sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic
|
||||
)
|
||||
|
|
@ -2299,13 +2299,13 @@ class MCPServerManager:
|
|||
from litellm.types.mcp import MCPAuth
|
||||
|
||||
if server.auth_type == MCPAuth.bearer_token:
|
||||
headers["Authorization"] = f"Bearer {server.authentication_token}"
|
||||
headers["Authorization"] = f"Bearer {strip_auth_scheme(server.authentication_token, 'Bearer')}"
|
||||
elif server.auth_type == MCPAuth.api_key:
|
||||
headers["Authorization"] = f"ApiKey {server.authentication_token}"
|
||||
headers["Authorization"] = f"ApiKey {strip_auth_scheme(server.authentication_token, 'ApiKey')}"
|
||||
elif server.auth_type == MCPAuth.basic:
|
||||
headers["Authorization"] = f"Basic {server.authentication_token}"
|
||||
headers["Authorization"] = f"Basic {to_basic_credentials(server.authentication_token)}"
|
||||
elif server.auth_type == MCPAuth.token:
|
||||
headers["Authorization"] = f"token {server.authentication_token}"
|
||||
headers["Authorization"] = f"token {strip_auth_scheme(server.authentication_token, 'token')}"
|
||||
|
||||
# Add any static headers from server config.
|
||||
#
|
||||
|
|
@ -3346,9 +3346,7 @@ class MCPServerManager:
|
|||
normalized: Final = {k.lower(): v for k, v in raw_headers.items()}
|
||||
auth_value = normalized.get("authorization")
|
||||
if auth_value:
|
||||
if auth_value.startswith("Bearer "):
|
||||
return auth_value[len("Bearer ") :]
|
||||
return auth_value
|
||||
return strip_auth_scheme(auth_value, "Bearer")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -6147,13 +6145,13 @@ class MCPServerManager:
|
|||
internal_networks = IPAddressUtils.parse_internal_networks(general_settings.get("mcp_internal_ip_ranges"))
|
||||
return IPAddressUtils.is_internal_ip(client_ip, internal_networks)
|
||||
|
||||
def get_mcp_server_by_id(self, server_id: str) -> MCPServer | None:
|
||||
"""
|
||||
Get the MCP Server from the server id
|
||||
"""
|
||||
def get_mcp_server_by_id(self, server_id: str, client_ip: str | None = None) -> MCPServer | None:
|
||||
"""Get the MCP Server from the server id."""
|
||||
registry: Final = self.get_registry()
|
||||
for server in registry.values():
|
||||
if server.server_id == server_id:
|
||||
if not self._is_server_accessible_from_ip(server, client_ip):
|
||||
return None
|
||||
return server
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from fastapi import HTTPException
|
|||
from pydantic import SecretStr
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm.experimental_mcp_client.client import strip_auth_scheme, to_basic_credentials
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import resolve_upstream_resource
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
|
|
@ -213,7 +214,9 @@ def _shared_key_spec(
|
|||
token: Final = server.authentication_token
|
||||
if not token:
|
||||
return None # no key configured -> defer to v1 (parity-safe)
|
||||
value: Final = base64.b64encode(token.encode("utf-8")).decode() if encode else token
|
||||
value: Final = (
|
||||
to_basic_credentials(token) if encode else strip_auth_scheme(token, value_prefix) if value_prefix else token
|
||||
)
|
||||
return ServerSpec(
|
||||
server_id=server.server_id,
|
||||
resource=resource,
|
||||
|
|
|
|||
|
|
@ -5,10 +5,13 @@ Handles agent permission checking for keys and teams using object_permission_id.
|
|||
Follows the same pattern as MCP permission handling.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, TypeAlias
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
|
||||
from litellm.proxy._types import (
|
||||
UI_TEAM_ID,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
|
|
@ -443,15 +446,47 @@ class AgentRequestHandler:
|
|||
return []
|
||||
|
||||
|
||||
async def accessible_agents(user_api_key_auth: UserAPIKeyAuth) -> tuple[AgentResponse, ...]:
|
||||
"""Every registry agent for proxy admins, else the agents the key's and team's grants reach."""
|
||||
def _granted_ids(access: AgentAccess) -> frozenset[str]:
|
||||
match access:
|
||||
case UnrestrictedAgentAccess():
|
||||
return frozenset()
|
||||
case RestrictedAgentAccess(agent_ids):
|
||||
return agent_ids
|
||||
|
||||
|
||||
ResolveAgentAccess: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[AgentAccess]]
|
||||
EffectiveAuthContexts: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[Sequence[UserAPIKeyAuth]]]
|
||||
|
||||
|
||||
async def _granted_agent_ids(
|
||||
user_api_key_auth: UserAPIKeyAuth,
|
||||
resolve_access: ResolveAgentAccess,
|
||||
effective_contexts: EffectiveAuthContexts,
|
||||
) -> frozenset[str]:
|
||||
"""Union of the explicit grants reachable from the key, its team, or (for a dashboard session)
|
||||
the user's real teams and user row. No grant anywhere yields the empty set, unlike the
|
||||
open-by-default ``resolve_agent_access`` that guards direct access."""
|
||||
accesses: Final = await asyncio.gather(
|
||||
*(resolve_access(auth_context) for auth_context in await effective_contexts(user_api_key_auth))
|
||||
)
|
||||
return frozenset().union(*(_granted_ids(access) for access in accesses))
|
||||
|
||||
|
||||
async def accessible_agents(
|
||||
user_api_key_auth: UserAPIKeyAuth,
|
||||
all_agents: tuple[AgentResponse, ...] | None = None,
|
||||
resolve_access: ResolveAgentAccess | None = None,
|
||||
effective_contexts: EffectiveAuthContexts = build_effective_auth_contexts,
|
||||
) -> tuple[AgentResponse, ...]:
|
||||
"""Every registry agent for proxy admins, else only the agents the caller was granted."""
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
|
||||
all_agents: Final = global_agent_registry.get_agent_list()
|
||||
agents: Final = global_agent_registry.get_agent_list() if all_agents is None else all_agents
|
||||
if user_api_key_auth.user_role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN.value):
|
||||
return all_agents
|
||||
match await AgentRequestHandler.resolve_agent_access(user_api_key_auth=user_api_key_auth):
|
||||
case UnrestrictedAgentAccess():
|
||||
return all_agents
|
||||
case RestrictedAgentAccess(allowed_agent_ids):
|
||||
return tuple(agent for agent in all_agents if agent.agent_id in allowed_agent_ids)
|
||||
return agents
|
||||
allowed_agent_ids: Final = await _granted_agent_ids(
|
||||
user_api_key_auth,
|
||||
AgentRequestHandler.resolve_agent_access if resolve_access is None else resolve_access,
|
||||
effective_contexts,
|
||||
)
|
||||
return tuple(agent for agent in agents if agent.agent_id in allowed_agent_ids)
|
||||
|
|
|
|||
|
|
@ -5788,7 +5788,7 @@ async def vector_store_access_check(
|
|||
vector_store_ids_to_run: Final = litellm.vector_store_registry.get_vector_store_ids_to_run(
|
||||
non_default_params=request_body, tools=request_body.get("tools", None)
|
||||
)
|
||||
if vector_store_ids_to_run is None:
|
||||
if not vector_store_ids_to_run:
|
||||
verbose_proxy_logger.debug("Vector store to run not found, skipping vector store access check")
|
||||
return True
|
||||
|
||||
|
|
|
|||
|
|
@ -5,10 +5,11 @@ This hook uses the DBSpendUpdateWriter to batch-write response IDs to the databa
|
|||
instead of writing immediately on each request.
|
||||
"""
|
||||
|
||||
from collections.abc import AsyncGenerator
|
||||
from collections.abc import AsyncGenerator, Callable, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -32,6 +33,44 @@ _RESPONSES_API_PROVIDER_PREFIX: Final = "/openai"
|
|||
_RESPONSES_API_CREATE_ROUTES: Final = frozenset({"/v1/responses", "/responses"})
|
||||
|
||||
|
||||
_RESPONSE_PAYLOAD_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def _response_payload(response_obj: object) -> Mapping[str, object] | None:
|
||||
try:
|
||||
return _RESPONSE_PAYLOAD_ADAPTER.validate_python(response_obj)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _rewrite_advertised_id(
|
||||
event: BaseLiteLLMOpenAIResponseObject,
|
||||
rewrite: Callable[[str], str],
|
||||
) -> BaseLiteLLMOpenAIResponseObject:
|
||||
event_id: Final = getattr(event, "id", None)
|
||||
if isinstance(event_id, str) and event_id.startswith("resp_"):
|
||||
setattr(event, "id", rewrite(event_id))
|
||||
return event
|
||||
|
||||
nested: Final = getattr(event, "response", None)
|
||||
if isinstance(nested, ResponsesAPIResponse):
|
||||
setattr(nested, "id", rewrite(nested.id))
|
||||
setattr(event, "response", nested)
|
||||
return event
|
||||
|
||||
payload: Final = _response_payload(nested)
|
||||
if payload is None:
|
||||
return event
|
||||
|
||||
payload_id: Final = payload.get("id")
|
||||
if not isinstance(payload_id, str):
|
||||
return event
|
||||
|
||||
rewritten: Final = {**payload, "id": rewrite(payload_id)} # mutable-ok: pydantic cannot serialize a frozen map
|
||||
setattr(event, "response", rewritten)
|
||||
return event
|
||||
|
||||
|
||||
def _is_responses_api_create_route(request_route: str | None) -> bool:
|
||||
if request_route is None:
|
||||
return False
|
||||
|
|
@ -196,10 +235,6 @@ class ResponsesIDSecurity(CustomLogger):
|
|||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
request_cache: dict[str, str] | None = None,
|
||||
) -> BaseLiteLLMOpenAIResponseObject:
|
||||
# encrypt the response id using the symmetric key
|
||||
# encrypt the response id, and encode the user id and response id in base64
|
||||
|
||||
# Check if signing key is available
|
||||
signing_key: Final = self._get_signing_key()
|
||||
if signing_key is None:
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -210,43 +245,22 @@ class ResponsesIDSecurity(CustomLogger):
|
|||
)
|
||||
return response
|
||||
|
||||
response_id: Final = getattr(response, "id", None)
|
||||
response_obj: Final = getattr(response, "response", None)
|
||||
def encrypt(original_id: str) -> str:
|
||||
cached: Final = request_cache.get(original_id) if request_cache is not None else None
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
if response_id and isinstance(response_id, str) and response_id.startswith("resp_"):
|
||||
# Check request-scoped cache first (for streaming consistency)
|
||||
if request_cache is not None and response_id in request_cache:
|
||||
setattr(response, "id", request_cache[response_id])
|
||||
else:
|
||||
encrypted_response_id = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
|
||||
response_id,
|
||||
user_api_key_dict.user_id or "",
|
||||
user_api_key_dict.team_id or "",
|
||||
)
|
||||
managed_id: Final = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
|
||||
original_id,
|
||||
user_api_key_dict.user_id or "",
|
||||
user_api_key_dict.team_id or "",
|
||||
)
|
||||
encrypted_id: Final = f"resp_{encrypt_value_helper(value=managed_id)}"
|
||||
if request_cache is not None:
|
||||
request_cache[original_id] = encrypted_id
|
||||
return encrypted_id
|
||||
|
||||
encoded_user_id_and_response_id = encrypt_value_helper(value=encrypted_response_id)
|
||||
encrypted_id = f"resp_{encoded_user_id_and_response_id}"
|
||||
if request_cache is not None:
|
||||
request_cache[response_id] = encrypted_id
|
||||
setattr(response, "id", encrypted_id)
|
||||
|
||||
elif response_obj and isinstance(response_obj, ResponsesAPIResponse):
|
||||
# Check request-scoped cache first (for streaming consistency)
|
||||
if request_cache is not None and response_obj.id in request_cache:
|
||||
setattr(response_obj, "id", request_cache[response_obj.id])
|
||||
else:
|
||||
encrypted_response_id = SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
|
||||
response_obj.id,
|
||||
user_api_key_dict.user_id or "",
|
||||
user_api_key_dict.team_id or "",
|
||||
)
|
||||
encoded_user_id_and_response_id = encrypt_value_helper(value=encrypted_response_id)
|
||||
encrypted_id = f"resp_{encoded_user_id_and_response_id}"
|
||||
if request_cache is not None:
|
||||
request_cache[response_obj.id] = encrypted_id
|
||||
setattr(response_obj, "id", encrypted_id)
|
||||
setattr(response, "response", response_obj)
|
||||
return response
|
||||
return _rewrite_advertised_id(response, encrypt)
|
||||
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ All /team management endpoints
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
import math
|
||||
import traceback
|
||||
|
|
@ -2189,6 +2190,26 @@ async def update_team(
|
|||
if field in updated_kv
|
||||
}
|
||||
|
||||
_writes_metadata_backed_field: Final = any(
|
||||
field in updated_kv
|
||||
for field in (
|
||||
*LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
*LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
)
|
||||
)
|
||||
if isinstance(existing_team_row.metadata, dict):
|
||||
if "metadata" not in updated_kv and (_team_member_fields_in_request or _writes_metadata_backed_field):
|
||||
updated_kv["metadata"] = copy.deepcopy(existing_team_row.metadata)
|
||||
elif isinstance(updated_kv.get("metadata"), dict):
|
||||
updated_kv["metadata"] = {
|
||||
**updated_kv["metadata"],
|
||||
**{
|
||||
key: existing_team_row.metadata[key]
|
||||
for key in TeamMemberBudgetHandler.SYSTEM_MANAGED_METADATA_KEYS
|
||||
if key in existing_team_row.metadata
|
||||
},
|
||||
}
|
||||
|
||||
if _team_member_fields_in_request and TeamMemberBudgetHandler.should_create_budget(
|
||||
team_member_budget=data.team_member_budget,
|
||||
team_member_rpm_limit=data.team_member_rpm_limit,
|
||||
|
|
|
|||
|
|
@ -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]}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@
|
|||
"limit": 52
|
||||
},
|
||||
"B010": {
|
||||
"limit": 190
|
||||
"limit": 187
|
||||
},
|
||||
"B018": {
|
||||
"limit": 2
|
||||
|
|
|
|||
|
|
@ -214,8 +214,8 @@ locals {
|
|||
gateway_uvicorn_args = "--host 0.0.0.0 --port 4000 --workers ${var.gateway_num_workers}"
|
||||
backend_uvicorn_args = "--host 0.0.0.0 --port 4001"
|
||||
|
||||
gateway_launch_cmd = "if [ \"$USE_DDTRACE\" = \"true\" ]; then export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args}; else exec uvicorn gateway.main:app ${local.gateway_uvicorn_args}; fi"
|
||||
backend_launch_cmd = "if [ \"$USE_DDTRACE\" = \"true\" ]; then export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args}; else exec uvicorn backend.main:app ${local.backend_uvicorn_args}; fi"
|
||||
gateway_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args};; *) exec uvicorn gateway.main:app ${local.gateway_uvicorn_args};; esac"
|
||||
backend_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args};; *) exec uvicorn backend.main:app ${local.backend_uvicorn_args};; esac"
|
||||
|
||||
gateway_proxy_overrides = local.proxy_config_enabled ? {
|
||||
entryPoint = ["sh", "-c"]
|
||||
|
|
|
|||
|
|
@ -138,8 +138,8 @@ locals {
|
|||
gateway_uvicorn_args = "--host 0.0.0.0 --port 4000 --workers ${var.gateway_num_workers}"
|
||||
backend_uvicorn_args = "--host 0.0.0.0 --port 4001"
|
||||
|
||||
gateway_launch_cmd = "if [ \"$USE_DDTRACE\" = \"true\" ]; then export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args}; else exec uvicorn gateway.main:app ${local.gateway_uvicorn_args}; fi"
|
||||
backend_launch_cmd = "if [ \"$USE_DDTRACE\" = \"true\" ]; then export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args}; else exec uvicorn backend.main:app ${local.backend_uvicorn_args}; fi"
|
||||
gateway_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn gateway.main:app ${local.gateway_uvicorn_args};; *) exec uvicorn gateway.main:app ${local.gateway_uvicorn_args};; esac"
|
||||
backend_launch_cmd = "case \"$USE_DDTRACE\" in [Tt][Rr][Uu][Ee]) export DD_TRACE_OPENAI_ENABLED=\"False\"; exec ddtrace-run uvicorn backend.main:app ${local.backend_uvicorn_args};; *) exec uvicorn backend.main:app ${local.backend_uvicorn_args};; esac"
|
||||
|
||||
gateway_args = join(" && ", concat(
|
||||
local.redis_ca_fragment,
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from litellm.caching import DualCache
|
|||
from litellm.proxy._types import CallTypes
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
_is_base64_encoded_unified_file_id,
|
||||
encode_file_id_with_model,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -3153,6 +3154,421 @@ async def test_same_user_different_keys_can_access_batch():
|
|||
assert result1["batch_id"] == result2["batch_id"]
|
||||
|
||||
|
||||
MODEL_ENCODED_BATCH_ID = encode_file_id_with_model(
|
||||
"batch_provider123", "gpt-4o-team-alias", id_type="batch"
|
||||
)
|
||||
MODEL_ENCODED_OUTPUT_FILE_ID = encode_file_id_with_model(
|
||||
"file-output456", "gpt-4o-team-alias", id_type="file"
|
||||
)
|
||||
RAW_PROVIDER_BATCH_ID = "batch_provider123"
|
||||
RAW_PROVIDER_FILE_ID = "file-output456"
|
||||
|
||||
|
||||
def _owned_record(created_by, team_id):
|
||||
record = MagicMock()
|
||||
record.created_by = created_by
|
||||
record.team_id = team_id
|
||||
return record
|
||||
|
||||
|
||||
def _batch_response(batch_id, output_file_id=None, is_create=False):
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
batch = LiteLLMBatch(
|
||||
id=batch_id,
|
||||
completion_window="24h",
|
||||
created_at=1700000000,
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="file-input789",
|
||||
object="batch",
|
||||
status="completed",
|
||||
output_file_id=output_file_id,
|
||||
)
|
||||
if is_create:
|
||||
batch._hidden_params["unified_file_id"] = "unified-input-file-id"
|
||||
return batch
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("call_type", ["aretrieve_batch", "acancel_batch"])
|
||||
@pytest.mark.parametrize(
|
||||
"batch_id", [MODEL_ENCODED_BATCH_ID, RAW_PROVIDER_BATCH_ID]
|
||||
)
|
||||
async def test_team_b_cannot_access_team_a_provider_format_batch(
|
||||
call_type, batch_id
|
||||
):
|
||||
"""
|
||||
Cross-team retrieve/cancel of a model-encoded or raw provider batch id
|
||||
must 403 when an ownership row exists for another team.
|
||||
|
||||
Regression test: before this check only unified (litellm_proxy-prefixed)
|
||||
batch ids were enforced, so any key could read any model-encoded or raw
|
||||
provider batch.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
||||
_owned_record(created_by="user_a", team_id="team_a")
|
||||
)
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
DualCache(), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await proxy_managed_files.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
||||
),
|
||||
cache=MagicMock(),
|
||||
data={"batch_id": batch_id},
|
||||
call_type=call_type,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
prisma_client.db.litellm_managedobjecttable.find_first.assert_awaited_once_with(
|
||||
where={"OR": [{"unified_object_id": batch_id}, {"model_object_id": batch_id}]}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"caller_kwargs",
|
||||
[
|
||||
{"user_id": "user_a", "team_id": "team_a"},
|
||||
{"user_id": "teammate_of_a", "team_id": "team_a"},
|
||||
{"user_id": "admin_user", "user_role": "proxy_admin"},
|
||||
],
|
||||
)
|
||||
async def test_authorized_callers_can_access_provider_format_batch(caller_kwargs):
|
||||
"""
|
||||
The creator, a same-team member, and a proxy admin can all retrieve a
|
||||
model-encoded batch owned by team_a. Data must pass through unmodified so
|
||||
the endpoint's own routing still applies.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
||||
_owned_record(created_by="user_a", team_id="team_a")
|
||||
)
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
DualCache(), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
result = await proxy_managed_files.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
parent_otel_span=MagicMock(), **caller_kwargs
|
||||
),
|
||||
cache=MagicMock(),
|
||||
data={"batch_id": MODEL_ENCODED_BATCH_ID},
|
||||
call_type="aretrieve_batch",
|
||||
)
|
||||
|
||||
assert result["batch_id"] == MODEL_ENCODED_BATCH_ID
|
||||
assert "model" not in result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_format_batch_without_ownership_row_stays_accessible():
|
||||
"""
|
||||
A provider-format batch id with no ownership row (created before ownership
|
||||
tracking, or directly on the provider account) must stay retrievable.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedobjecttable.find_first.return_value = None
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
DualCache(), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
result = await proxy_managed_files.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
||||
),
|
||||
cache=MagicMock(),
|
||||
data={"batch_id": RAW_PROVIDER_BATCH_ID},
|
||||
call_type="aretrieve_batch",
|
||||
)
|
||||
|
||||
assert result["batch_id"] == RAW_PROVIDER_BATCH_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fine_tuning_provider_format_id_not_enforced():
|
||||
"""
|
||||
Provider-format fine-tuning job ids are deliberately out of scope for
|
||||
ownership enforcement; only unified fine-tuning ids are checked.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
||||
_owned_record(created_by="user_a", team_id="team_a")
|
||||
)
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
DualCache(), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
result = await proxy_managed_files.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
||||
),
|
||||
cache=MagicMock(),
|
||||
data={"fine_tuning_job_id": "ftjob-abc123"},
|
||||
call_type="aretrieve_fine_tuning_job",
|
||||
)
|
||||
|
||||
assert result["fine_tuning_job_id"] == "ftjob-abc123"
|
||||
prisma_client.db.litellm_managedobjecttable.find_first.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"call_type", ["afile_content", "afile_retrieve", "afile_delete"]
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"file_id", [MODEL_ENCODED_OUTPUT_FILE_ID, RAW_PROVIDER_FILE_ID]
|
||||
)
|
||||
async def test_team_b_cannot_access_team_a_provider_format_file(
|
||||
call_type, file_id
|
||||
):
|
||||
"""
|
||||
Cross-team content/retrieve/delete of a model-encoded or raw provider
|
||||
file id must 403 when an ownership row exists for another team.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedfiletable.find_first.return_value = (
|
||||
_owned_record(created_by="user_a", team_id="team_a")
|
||||
)
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
MagicMock(), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await proxy_managed_files.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
||||
),
|
||||
cache=MagicMock(),
|
||||
data={"file_id": file_id},
|
||||
call_type=call_type,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
prisma_client.db.litellm_managedfiletable.find_first.assert_awaited_once_with(
|
||||
where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_same_team_can_access_provider_format_file():
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedfiletable.find_first.return_value = (
|
||||
_owned_record(created_by="user_a", team_id="team_a")
|
||||
)
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
MagicMock(), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
result = await proxy_managed_files.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="teammate_of_a", team_id="team_a", parent_otel_span=MagicMock()
|
||||
),
|
||||
cache=MagicMock(),
|
||||
data={"file_id": MODEL_ENCODED_OUTPUT_FILE_ID},
|
||||
call_type="afile_content",
|
||||
)
|
||||
|
||||
assert result["file_id"] == MODEL_ENCODED_OUTPUT_FILE_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_format_file_without_ownership_row_stays_accessible():
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedfiletable.find_first.return_value = None
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
MagicMock(), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
result = await proxy_managed_files.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
||||
),
|
||||
cache=MagicMock(),
|
||||
data={"file_id": RAW_PROVIDER_FILE_ID},
|
||||
call_type="afile_content",
|
||||
)
|
||||
|
||||
assert result["file_id"] == RAW_PROVIDER_FILE_ID
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_batch_create_stores_ownership_row():
|
||||
"""
|
||||
Batch creation (response hidden params carry the unified input file id)
|
||||
must write an ownership row attributed to the creating key.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
await proxy_managed_files.async_post_call_success_hook(
|
||||
data={
|
||||
"input_file_id": "file-input789",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
|
||||
),
|
||||
response=_batch_response(MODEL_ENCODED_BATCH_ID, is_create=True),
|
||||
)
|
||||
|
||||
upsert_call = prisma_client.db.litellm_managedobjecttable.upsert.await_args
|
||||
assert upsert_call.kwargs["where"] == {
|
||||
"unified_object_id": MODEL_ENCODED_BATCH_ID
|
||||
}
|
||||
create_data = upsert_call.kwargs["data"]["create"]
|
||||
assert create_data["created_by"] == "user_a"
|
||||
assert create_data["team_id"] == "team_a"
|
||||
prisma_client.db.litellm_managedobjecttable.update_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_batch_sync_does_not_claim_ownership():
|
||||
"""
|
||||
Retrieve/cancel of a batch with no ownership row must NOT create one:
|
||||
otherwise the first foreign key to touch a legacy batch would become its
|
||||
owner and lock out the real creator once enforcement is on.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 0
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
await proxy_managed_files.async_post_call_success_hook(
|
||||
data={"batch_id": MODEL_ENCODED_BATCH_ID},
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_b", team_id="team_b", parent_otel_span=MagicMock()
|
||||
),
|
||||
response=_batch_response(MODEL_ENCODED_BATCH_ID),
|
||||
)
|
||||
|
||||
prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited()
|
||||
prisma_client.db.litellm_managedobjecttable.update_many.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_batch_sync_updates_existing_row():
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1
|
||||
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
||||
_owned_record(created_by="user_a", team_id="team_a")
|
||||
)
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
await proxy_managed_files.async_post_call_success_hook(
|
||||
data={"batch_id": MODEL_ENCODED_BATCH_ID},
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
|
||||
),
|
||||
response=_batch_response(MODEL_ENCODED_BATCH_ID),
|
||||
)
|
||||
|
||||
update_call = prisma_client.db.litellm_managedobjecttable.update_many.await_args
|
||||
assert update_call.kwargs["where"] == {
|
||||
"unified_object_id": MODEL_ENCODED_BATCH_ID
|
||||
}
|
||||
assert update_call.kwargs["data"]["status"] == "completed"
|
||||
prisma_client.db.litellm_managedobjecttable.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_batch_sync_stores_output_file_ownership_from_batch_row():
|
||||
"""
|
||||
When a synced batch reports a provider-format output file id, an
|
||||
ownership row for that file must be written with the BATCH row's
|
||||
created_by/team_id, not the caller's identity.
|
||||
"""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
prisma_client.db.litellm_managedobjecttable.update_many.return_value = 1
|
||||
prisma_client.db.litellm_managedobjecttable.find_first.return_value = (
|
||||
_owned_record(created_by="user_a", team_id="team_a")
|
||||
)
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
await proxy_managed_files.async_post_call_success_hook(
|
||||
data={"batch_id": MODEL_ENCODED_BATCH_ID},
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="admin_user",
|
||||
user_role="proxy_admin",
|
||||
parent_otel_span=MagicMock(),
|
||||
),
|
||||
response=_batch_response(
|
||||
MODEL_ENCODED_BATCH_ID, output_file_id=MODEL_ENCODED_OUTPUT_FILE_ID
|
||||
),
|
||||
)
|
||||
|
||||
file_upsert = prisma_client.db.litellm_managedfiletable.upsert.await_args
|
||||
assert file_upsert.kwargs["where"] == {
|
||||
"unified_file_id": MODEL_ENCODED_OUTPUT_FILE_ID
|
||||
}
|
||||
create_data = file_upsert.kwargs["data"]["create"]
|
||||
assert create_data["created_by"] == "user_a"
|
||||
assert create_data["team_id"] == "team_a"
|
||||
assert create_data["flat_model_file_ids"] == ["file-output456"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_batch_create_does_not_store_output_file_ownership():
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
prisma_client = AsyncMock()
|
||||
proxy_managed_files = _PROXY_LiteLLMManagedFiles(
|
||||
MagicMock(async_set_cache=AsyncMock()), prisma_client=prisma_client
|
||||
)
|
||||
|
||||
await proxy_managed_files.async_post_call_success_hook(
|
||||
data={
|
||||
"input_file_id": "file-input789",
|
||||
"endpoint": "/v1/chat/completions",
|
||||
"completion_window": "24h",
|
||||
},
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
user_id="user_a", team_id="team_a", parent_otel_span=MagicMock()
|
||||
),
|
||||
response=_batch_response(
|
||||
MODEL_ENCODED_BATCH_ID,
|
||||
output_file_id=MODEL_ENCODED_OUTPUT_FILE_ID,
|
||||
is_create=True,
|
||||
),
|
||||
)
|
||||
|
||||
prisma_client.db.litellm_managedfiletable.upsert.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_list_cursors_are_scoped_to_the_caller():
|
||||
"""A non-owner must not learn other callers' file ids through the page cursors."""
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -9,14 +9,16 @@ be forced by a sequential script: it needs one request to be genuinely mid-fligh
|
|||
other commits. A mocked prisma cannot arbitrate that either, since the property under test
|
||||
is whether Postgres's own advisory lock actually serializes the two requests.
|
||||
|
||||
These tests pin the interleaving the same way test_access_group_team_sync.py does: a second
|
||||
real connection holds the team's advisory lock in its own transaction, so the function under
|
||||
test is provably blocked on it rather than hoping a sleep lands in the right gap.
|
||||
These tests pin the interleaving without a timing assumption: a second real connection holds
|
||||
the team's advisory lock in its own transaction, and the test then waits for Postgres itself
|
||||
to report the endpoint queued behind that exact lock. A sleep can only guess whether the
|
||||
endpoint has reached the lock yet; pg_locks answers it.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import timedelta
|
||||
|
|
@ -39,6 +41,47 @@ _DELETE_SEEDED = 'DELETE FROM "LiteLLM_TeamMembership" WHERE team_id = $1'
|
|||
_DELETE_USER = 'DELETE FROM "LiteLLM_UserTable" WHERE user_id = $1'
|
||||
_DELETE_TEAM = 'DELETE FROM "LiteLLM_TeamTable" WHERE team_id = $1'
|
||||
_LOCK_SQL = "SELECT pg_advisory_xact_lock(hashtext($1)) IS NULL AS locked"
|
||||
_HELD_LOCK_KEY_SQL = (
|
||||
"SELECT classid::bigint AS classid, objid::bigint AS objid FROM pg_locks "
|
||||
"WHERE locktype = 'advisory' AND granted AND pid = pg_backend_pid()"
|
||||
)
|
||||
_LOCK_WAITER_SQL = (
|
||||
"SELECT count(*)::int AS waiters FROM pg_locks "
|
||||
"WHERE locktype = 'advisory' AND NOT granted "
|
||||
"AND classid::bigint = $1 AND objid::bigint = $2"
|
||||
)
|
||||
_LOCK_WAIT_TIMEOUT_SECONDS = 20.0
|
||||
_LOCK_POLL_SECONDS = 0.01
|
||||
|
||||
|
||||
async def _hold_team_lock(held, team_id: str) -> tuple[int, int]:
|
||||
"""Take the team's advisory lock and return its pg_locks key.
|
||||
|
||||
Reading the key back off our own backend avoids re-deriving hashtext()'s signed
|
||||
32-bit split here, and pins the watcher to this lock rather than to any advisory
|
||||
lock another xdist worker happens to hold on the same database."""
|
||||
await held.query_raw(_LOCK_SQL, team_id)
|
||||
rows = await held.query_raw(_HELD_LOCK_KEY_SQL)
|
||||
assert len(rows) == 1, f"expected exactly one advisory lock on the blocking connection, got {rows}"
|
||||
return rows[0]["classid"], rows[0]["objid"]
|
||||
|
||||
|
||||
async def _await_lock_contention(watcher, lock_key: tuple[int, int], task, what: str) -> None:
|
||||
"""Block until Postgres reports `task` queued behind the held lock.
|
||||
|
||||
This is the assertion that the endpoint serializes on the team's advisory lock, and it
|
||||
is what a fixed sleep was standing in for: the endpoint is only provably waiting once a
|
||||
non-granted advisory lock on the same key exists. `watcher` must be a connection that is
|
||||
not itself blocked, so it can observe the queue."""
|
||||
classid, objid = lock_key
|
||||
deadline = time.monotonic() + _LOCK_WAIT_TIMEOUT_SECONDS
|
||||
while time.monotonic() < deadline:
|
||||
if task.done():
|
||||
raise AssertionError(f"{what} returned without waiting on the team's advisory lock") from task.exception()
|
||||
if (await watcher.query_raw(_LOCK_WAITER_SQL, classid, objid))[0]["waiters"]:
|
||||
return
|
||||
await asyncio.sleep(_LOCK_POLL_SECONDS)
|
||||
raise AssertionError(f"{what} never queued on the team's advisory lock within {_LOCK_WAIT_TIMEOUT_SECONDS}s")
|
||||
|
||||
|
||||
def _race_ids() -> tuple[str, str]:
|
||||
|
|
@ -110,10 +153,8 @@ async def test_member_add_blocked_by_delete_writes_no_dangling_reference():
|
|||
|
||||
blocker = Prisma()
|
||||
await blocker.connect()
|
||||
lock_acquired = asyncio.Event()
|
||||
|
||||
async def add_member():
|
||||
lock_acquired.set()
|
||||
await _add_team_members_to_team(
|
||||
data=TeamMemberAddRequest(
|
||||
team_id=team_id,
|
||||
|
|
@ -128,11 +169,9 @@ async def test_member_add_blocked_by_delete_writes_no_dangling_reference():
|
|||
|
||||
try:
|
||||
async with blocker.tx(timeout=timedelta(seconds=30)) as held:
|
||||
await held.query_raw(_LOCK_SQL, team_id)
|
||||
lock_key = await _hold_team_lock(held, team_id)
|
||||
task = asyncio.create_task(add_member())
|
||||
await lock_acquired.wait()
|
||||
await asyncio.sleep(0.2)
|
||||
assert not task.done(), "member_add did not wait on the team's advisory lock"
|
||||
await _await_lock_contention(db, lock_key, task, "member_add")
|
||||
|
||||
# the delete wins the race: strip the team row while the lock is held
|
||||
await held.execute_raw(_DELETE_TEAM, team_id)
|
||||
|
|
@ -185,10 +224,8 @@ async def test_member_delete_blocked_by_member_add_removes_from_the_fresh_roster
|
|||
|
||||
blocker = Prisma()
|
||||
await blocker.connect()
|
||||
lock_acquired = asyncio.Event()
|
||||
|
||||
async def run_delete():
|
||||
lock_acquired.set()
|
||||
return await team_member_delete(
|
||||
data=TeamMemberDeleteRequest(team_id=team_id, user_id=user_id),
|
||||
user_api_key_dict=_admin_auth(),
|
||||
|
|
@ -196,11 +233,9 @@ async def test_member_delete_blocked_by_member_add_removes_from_the_fresh_roster
|
|||
|
||||
try:
|
||||
async with blocker.tx(timeout=timedelta(seconds=30)) as held:
|
||||
await held.query_raw(_LOCK_SQL, team_id)
|
||||
lock_key = await _hold_team_lock(held, team_id)
|
||||
task = asyncio.create_task(run_delete())
|
||||
await lock_acquired.wait()
|
||||
await asyncio.sleep(0.2)
|
||||
assert not task.done(), "member_delete did not wait on the team's advisory lock"
|
||||
await _await_lock_contention(db, lock_key, task, "member_delete")
|
||||
|
||||
# member_add wins the race: it adds `other_user` while holding the lock
|
||||
await held.litellm_teamtable.update(
|
||||
|
|
@ -265,10 +300,8 @@ async def test_delete_blocked_by_member_add_sweeps_the_fresh_reference():
|
|||
|
||||
blocker = Prisma()
|
||||
await blocker.connect()
|
||||
lock_acquired = asyncio.Event()
|
||||
|
||||
async def run_delete():
|
||||
lock_acquired.set()
|
||||
return await delete_team(
|
||||
data=DeleteTeamRequest(team_ids=[team_id]),
|
||||
http_request=MagicMock(),
|
||||
|
|
@ -278,11 +311,9 @@ async def test_delete_blocked_by_member_add_sweeps_the_fresh_reference():
|
|||
|
||||
try:
|
||||
async with blocker.tx(timeout=timedelta(seconds=30)) as held:
|
||||
await held.query_raw(_LOCK_SQL, team_id)
|
||||
lock_key = await _hold_team_lock(held, team_id)
|
||||
task = asyncio.create_task(run_delete())
|
||||
await lock_acquired.wait()
|
||||
await asyncio.sleep(0.3)
|
||||
assert not task.done(), "delete_team did not wait on the team's advisory lock"
|
||||
await _await_lock_contention(db, lock_key, task, "delete_team")
|
||||
|
||||
# member_add wins the race: write the reference while holding the lock
|
||||
await held.litellm_usertable.upsert(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1014,6 +1014,8 @@ class TestAuthSchemeNormalization:
|
|||
[
|
||||
("Bearer abc", "Bearer", "abc"),
|
||||
("bearer abc", "Bearer", "abc"),
|
||||
("Bearer\tabc", "Bearer", "abc"),
|
||||
("bearer\t\tabc", "Bearer", "abc"),
|
||||
(" Bearer abc ", "Bearer", "abc "),
|
||||
("abc", "Bearer", "abc"),
|
||||
("Bearerabc", "Bearer", "Bearerabc"),
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -97,22 +97,30 @@ class TestResolveConfig:
|
|||
def test_plaintext_http_api_base_disables(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
|
||||
assert resolve_openai_workload_identity_config(api_key=None, api_base="http://api.openai.com/v1") is None
|
||||
|
||||
@pytest.mark.parametrize("regional_host", ("eu.api.openai.com", "us.api.openai.com"))
|
||||
def test_regional_openai_api_base_allows(
|
||||
self, wif_env: OpenAIWorkloadIdentityConfig, regional_host: str
|
||||
) -> None:
|
||||
assert (
|
||||
resolve_openai_workload_identity_config(api_key=None, api_base=f"https://{regional_host}/v1") == wif_env
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"api_base",
|
||||
(
|
||||
"https://southcentralus.privatelink.api.openai.com/v1",
|
||||
"https://eu.api.openai.com/v1",
|
||||
"https://us.api.openai.com/v1",
|
||||
),
|
||||
)
|
||||
def test_openai_backed_api_base_allows(self, wif_env: OpenAIWorkloadIdentityConfig, api_base: str) -> None:
|
||||
assert resolve_openai_workload_identity_config(api_key=None, api_base=api_base) == wif_env
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"lookalike_base",
|
||||
("https://api.openai.com.evil.example/v1", "https://openai.com/v1", "https://euapi.openai.com/v1"),
|
||||
"api_base",
|
||||
(
|
||||
"https://api.openai.com.evil.example/v1",
|
||||
"https://openai.com/v1",
|
||||
"https://euapi.openai.com/v1",
|
||||
"http://southcentralus.privatelink.api.openai.com/v1",
|
||||
),
|
||||
)
|
||||
def test_openai_lookalike_api_base_disables(
|
||||
self, wif_env: OpenAIWorkloadIdentityConfig, lookalike_base: str
|
||||
def test_lookalike_or_plaintext_api_base_disables(
|
||||
self, wif_env: OpenAIWorkloadIdentityConfig, api_base: str
|
||||
) -> None:
|
||||
assert resolve_openai_workload_identity_config(api_key=None, api_base=lookalike_base) is None
|
||||
assert resolve_openai_workload_identity_config(api_key=None, api_base=api_base) is None
|
||||
|
||||
def test_foreign_env_base_url_disables(
|
||||
self, wif_env: OpenAIWorkloadIdentityConfig, monkeypatch: pytest.MonkeyPatch
|
||||
|
|
@ -187,6 +195,14 @@ class TestClientConstruction:
|
|||
assert client.api_key == "workload-identity-auth"
|
||||
assert client._workload_identity_auth is not None
|
||||
|
||||
def test_privatelink_client_uses_workload_identity(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
|
||||
client: Final = OpenAIChatCompletion()._get_openai_client(
|
||||
is_async=False, api_key=None, api_base="https://southcentralus.privatelink.api.openai.com/v1"
|
||||
)
|
||||
assert isinstance(client, OpenAI)
|
||||
assert client.api_key == "workload-identity-auth"
|
||||
assert client._workload_identity_auth is not None
|
||||
|
||||
def test_static_key_client_unaffected(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
|
||||
client: Final = OpenAIChatCompletion()._get_openai_client(is_async=False, api_key="sk-static", api_base=None)
|
||||
assert isinstance(client, OpenAI)
|
||||
|
|
@ -260,6 +276,16 @@ class TestResponsesValidateEnvironment:
|
|||
)
|
||||
assert headers["Authorization"] == "Bearer None"
|
||||
|
||||
@respx.mock
|
||||
def test_privatelink_api_base_mints_bearer(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
|
||||
mock_token_exchange()
|
||||
headers: Final = OpenAIResponsesAPIConfig().validate_environment(
|
||||
headers={},
|
||||
model="gpt-4o-mini",
|
||||
litellm_params=GenericLiteLLMParams(api_base="https://southcentralus.privatelink.api.openai.com/v1"),
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer exchanged-bearer-token"
|
||||
|
||||
def test_litellm_proxy_subclass_never_mints_wif(self, wif_env: OpenAIWorkloadIdentityConfig) -> None:
|
||||
headers: Final = LiteLLMProxyResponsesAPIConfig().validate_environment(
|
||||
headers={}, model="gpt-4o-mini", litellm_params=GenericLiteLLMParams()
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import pytest
|
|||
from fastapi import HTTPException
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import (
|
||||
oauth_protected_resource_path,
|
||||
raise_public,
|
||||
|
|
@ -94,6 +95,49 @@ def test_basic_scheme_base64_encodes_the_token():
|
|||
assert spec.config.key_source.value.get_secret_value() == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, authentication_token, expected_value, expected_header",
|
||||
[
|
||||
(MCPAuth.bearer_token, "Bearer abc", "abc", ("Authorization", "Bearer abc")),
|
||||
(MCPAuth.token, "token abc", "abc", ("Authorization", "token abc")),
|
||||
(MCPAuth.basic, "user:pass", "dXNlcjpwYXNz", ("Authorization", "Basic dXNlcjpwYXNz")),
|
||||
(MCPAuth.basic, "Basic dXNlcjpwYXNz", "dXNlcjpwYXNz", ("Authorization", "Basic dXNlcjpwYXNz")),
|
||||
(MCPAuth.basic, "Basic user:pass", "dXNlcjpwYXNz", ("Authorization", "Basic dXNlcjpwYXNz")),
|
||||
],
|
||||
)
|
||||
def test_shared_key_normalizes_schemed_authentication_token(
|
||||
auth_type, authentication_token, expected_value, expected_header
|
||||
):
|
||||
spec = to_server_spec(_server(auth_type=auth_type, authentication_token=authentication_token))
|
||||
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
|
||||
assert spec.config.key_source.value.get_secret_value() == expected_value
|
||||
assert spec.config.header(expected_value) == expected_header
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, authentication_token",
|
||||
[
|
||||
(MCPAuth.bearer_token, "Bearer abc"),
|
||||
(MCPAuth.bearer_token, "abc"),
|
||||
(MCPAuth.token, "token abc"),
|
||||
(MCPAuth.token, "abc"),
|
||||
(MCPAuth.basic, "user:pass"),
|
||||
(MCPAuth.basic, "Basic dXNlcjpwYXNz"),
|
||||
(MCPAuth.basic, "Basic user:pass"),
|
||||
],
|
||||
)
|
||||
def test_shared_key_authorization_matches_v1(auth_type, authentication_token):
|
||||
spec = to_server_spec(_server(auth_type=auth_type, authentication_token=authentication_token))
|
||||
assert spec is not None and isinstance(spec.config, ApiKeyConfig)
|
||||
|
||||
client = MCPClient(server_url="https://x", auth_type=auth_type)
|
||||
client.update_auth_value(authentication_token)
|
||||
|
||||
assert spec.config.header(spec.config.key_source.value.get_secret_value())[1] == client._get_auth_headers()[
|
||||
"Authorization"
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"oauth2_flow",
|
||||
[None, "authorization_code"],
|
||||
|
|
|
|||
|
|
@ -3508,6 +3508,169 @@ def _create_oauth2_server(
|
|||
)
|
||||
|
||||
|
||||
def _create_id_lookup_oauth2_server():
|
||||
return _create_oauth2_server(
|
||||
server_id="oauth-server-id",
|
||||
name="oauth-server-name",
|
||||
server_name="oauth-server-name",
|
||||
alias="oauth-server-alias",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorize_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server()
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
patch.object(discoverable_endpoints, "encrypt_value_helper", return_value="encrypted-state"), # test-quality-ok: flow seam
|
||||
):
|
||||
response = await discoverable_endpoints.authorize(
|
||||
request=request,
|
||||
client_id=server.client_id,
|
||||
mcp_server_name=server.server_id,
|
||||
redirect_uri="http://localhost:62646/callback",
|
||||
state="test_state",
|
||||
)
|
||||
|
||||
assert response.status_code == 307
|
||||
assert "https://provider.com/oauth/authorize" in response.headers["location"]
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server()
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
response = MagicMock()
|
||||
response.json.return_value = {"access_token": "token", "token_type": "Bearer"}
|
||||
response.raise_for_status = MagicMock()
|
||||
client = MagicMock()
|
||||
client.post = AsyncMock(return_value=response)
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
patch.object(discoverable_endpoints, "get_async_httpx_client", return_value=client), # test-quality-ok: HTTP seam
|
||||
):
|
||||
result = await discoverable_endpoints.token_endpoint(
|
||||
request=request,
|
||||
grant_type="authorization_code",
|
||||
code="test_code",
|
||||
redirect_uri="http://localhost:62646/callback",
|
||||
client_id=server.client_id,
|
||||
mcp_server_name=server.server_id,
|
||||
client_secret=server.client_secret,
|
||||
)
|
||||
|
||||
assert json.loads(result.body)["access_token"] == "token"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_client_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server().model_copy(
|
||||
update={"client_id": None, "client_secret": None, "registration_url": "https://provider.com/oauth/register"}
|
||||
)
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
response = MagicMock()
|
||||
response.json.return_value = {"client_id": "registered-client", "client_secret": "registered-secret"}
|
||||
response.raise_for_status = MagicMock()
|
||||
client = MagicMock()
|
||||
client.post = AsyncMock(return_value=response)
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
patch.object(discoverable_endpoints, "_read_request_body", new=AsyncMock(return_value={})), # test-quality-ok: request seam
|
||||
patch.object(discoverable_endpoints, "get_async_httpx_client", return_value=client), # test-quality-ok: HTTP seam
|
||||
):
|
||||
result = await discoverable_endpoints.register_client(request=request, mcp_server_name=server.server_id)
|
||||
|
||||
assert json.loads(result.body)["client_id"] == "registered-client"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_protected_resource_metadata_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server()
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
):
|
||||
result = await discoverable_endpoints._build_oauth_protected_resource_response(
|
||||
request=request,
|
||||
mcp_server_name=server.server_id,
|
||||
use_standard_pattern=True,
|
||||
)
|
||||
|
||||
assert result["authorization_servers"] == ["https://llm.example.com/mcp"]
|
||||
assert result["resource"] == f"https://llm.example.com/mcp/{server.server_id}"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
def test_authorization_server_metadata_resolves_server_by_id_when_name_lookup_fails():
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
|
||||
|
||||
server = _create_id_lookup_oauth2_server()
|
||||
request = MagicMock(spec=Request)
|
||||
request.base_url = "https://llm.example.com/"
|
||||
request.headers = {}
|
||||
|
||||
with (
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_name", return_value=None) as by_name, # test-quality-ok: resolver seam
|
||||
patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=server) as by_id, # test-quality-ok: resolver seam
|
||||
):
|
||||
result = discoverable_endpoints._build_oauth_authorization_server_response(
|
||||
request=request,
|
||||
mcp_server_name=server.server_id,
|
||||
)
|
||||
|
||||
assert result["scopes_supported"] == server.scopes
|
||||
assert result["issuer"] == f"https://llm.example.com/{server.server_id}"
|
||||
by_name.assert_called_once_with(server.server_id, client_ip=None)
|
||||
by_id.assert_called_once_with(server.server_id, client_ip=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorize_root_resolves_single_oauth2_server():
|
||||
"""When /authorize is hit without server name and exactly 1 OAuth2 server exists, resolve it."""
|
||||
|
|
|
|||
|
|
@ -129,6 +129,32 @@ class TestMCPServerManager:
|
|||
assert added_server.args == ["-m", "server"]
|
||||
assert added_server.env == {"DEBUG": "1", "TEST": "1"}
|
||||
|
||||
def test_get_mcp_server_by_id_allows_internal_or_unspecified_client_ip(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="private-server",
|
||||
name="private-server",
|
||||
transport=MCPTransport.http,
|
||||
available_on_public_internet=False,
|
||||
)
|
||||
manager.registry[server.server_id] = server
|
||||
|
||||
assert manager.get_mcp_server_by_id(server.server_id) is server
|
||||
assert manager.get_mcp_server_by_id(server.server_id, client_ip="10.0.0.1") is server
|
||||
|
||||
def test_get_mcp_server_by_id_rejects_private_server_for_public_ip(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="private-server",
|
||||
name="private-server",
|
||||
transport=MCPTransport.http,
|
||||
available_on_public_internet=False,
|
||||
)
|
||||
manager.registry[server.server_id] = server
|
||||
|
||||
with patch.object(manager, "_get_general_settings", return_value={}):
|
||||
assert manager.get_mcp_server_by_id(server.server_id, client_ip="8.8.8.8") is None
|
||||
|
||||
async def test_create_mcp_client_stdio(self):
|
||||
"""Test creating MCP client for stdio transport"""
|
||||
manager = MCPServerManager()
|
||||
|
|
@ -2590,6 +2616,69 @@ class TestMCPServerManager:
|
|||
await manager.preflight_token_exchange(server=server, oauth2_headers=None, user_api_key_auth=None)
|
||||
assert resolved == ["good-subject"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"authorization",
|
||||
[
|
||||
"Bearer subj-jwt",
|
||||
"bearer subj-jwt",
|
||||
"BEARER subj-jwt",
|
||||
"Bearer\tsubj-jwt",
|
||||
"Bearer subj-jwt",
|
||||
],
|
||||
)
|
||||
async def test_preflight_token_exchange_strips_inbound_authorization_scheme(self, authorization):
|
||||
"""The resolver posts inbound_token verbatim as subject_token."""
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
|
||||
StaticHeaderAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError, ServerSpec, Subject
|
||||
|
||||
resolved: Final[list[str | None]] = []
|
||||
|
||||
class _FakeProvider:
|
||||
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]:
|
||||
resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None)
|
||||
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization"))
|
||||
|
||||
manager = MCPServerManager(cred_provider=_FakeProvider())
|
||||
server = self._token_exchange_server(f"te-preflight-auth-{authorization!r}")
|
||||
|
||||
await manager.preflight_token_exchange(
|
||||
server=server,
|
||||
oauth2_headers={"Authorization": authorization},
|
||||
user_api_key_auth=None,
|
||||
)
|
||||
|
||||
assert resolved == ["subj-jwt"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_token_exchange_preserves_authorization_without_separator(self):
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
|
||||
StaticHeaderAuth,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok
|
||||
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import CredError, ServerSpec, Subject
|
||||
|
||||
resolved: Final[list[str | None]] = []
|
||||
|
||||
class _FakeProvider:
|
||||
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Ok[StaticHeaderAuth, CredError]:
|
||||
resolved.append(subject.inbound_token.get_secret_value() if subject.inbound_token else None)
|
||||
return Ok(StaticHeaderAuth("Bearer MINTED", header_name="Authorization"))
|
||||
|
||||
manager = MCPServerManager(cred_provider=_FakeProvider())
|
||||
server = self._token_exchange_server("te-preflight-auth-no-separator")
|
||||
|
||||
await manager.preflight_token_exchange(
|
||||
server=server,
|
||||
oauth2_headers={"Authorization": "Bearersubj-jwt"},
|
||||
user_api_key_auth=None,
|
||||
)
|
||||
|
||||
assert resolved == ["Bearersubj-jwt"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preflight_token_exchange_skips_discovery_for_other_auth_modes(self):
|
||||
"""Preflight must not make unrelated auth modes depend on OAuth discovery."""
|
||||
|
|
@ -4545,6 +4634,81 @@ class TestMCPServerManager:
|
|||
# auth_type is none here, so a 401 from this upstream must not be dressed up as a re-auth signal
|
||||
assert captured["relays_upstream_auth"] is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"auth_type, authentication_token, expected_authorization",
|
||||
[
|
||||
(MCPAuth.bearer_token, "Bearer abc", "Bearer abc"),
|
||||
(MCPAuth.bearer_token, "abc", "Bearer abc"),
|
||||
(MCPAuth.api_key, "ApiKey abc", "ApiKey abc"),
|
||||
(MCPAuth.token, "token abc", "token abc"),
|
||||
(MCPAuth.basic, "user:pass", "Basic dXNlcjpwYXNz"),
|
||||
(MCPAuth.basic, "Basic dXNlcjpwYXNz", "Basic dXNlcjpwYXNz"),
|
||||
(MCPAuth.basic, "Basic user:pass", "Basic dXNlcjpwYXNz"),
|
||||
],
|
||||
)
|
||||
async def test_register_openapi_tools_normalizes_authentication_token(
|
||||
self, tmp_path, monkeypatch, auth_type, authentication_token, expected_authorization
|
||||
):
|
||||
manager = MCPServerManager()
|
||||
spec_path = tmp_path / "openapi.json"
|
||||
spec_path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"openapi": "3.0.0",
|
||||
"info": {"title": "Demo", "version": "1.0.0"},
|
||||
"paths": {
|
||||
"/health": {
|
||||
"get": {
|
||||
"operationId": "health_check",
|
||||
"summary": "health",
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
server = MCPServer(
|
||||
server_id="openapi-server",
|
||||
name="openapi-server",
|
||||
server_name="openapi-server",
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=auth_type,
|
||||
authentication_token=authentication_token,
|
||||
)
|
||||
captured: dict = {}
|
||||
|
||||
def fake_create_tool_function(
|
||||
path, method, operation, base_url, headers=None, server_label=None, relays_upstream_auth=False
|
||||
):
|
||||
captured["headers"] = headers
|
||||
|
||||
async def tool_func(**kwargs):
|
||||
return "ok"
|
||||
|
||||
return tool_func
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.create_tool_function",
|
||||
fake_create_tool_function,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator.build_input_schema",
|
||||
lambda *args, **kwargs: {"type": "object", "properties": {}, "required": []},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy._experimental.mcp_server.tool_registry.global_mcp_tool_registry.register_tool",
|
||||
lambda *args, **kwargs: None,
|
||||
)
|
||||
await manager._register_openapi_tools(
|
||||
spec_path=str(spec_path),
|
||||
server=server,
|
||||
base_url="https://example.com",
|
||||
)
|
||||
|
||||
assert captured["headers"]["Authorization"] == expected_authorization
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_tool_check_allowed_tools_list_allows_tool(self):
|
||||
"""Test pre_call_tool_check allows tool when it's in allowed_tools list"""
|
||||
|
|
|
|||
|
|
@ -10,15 +10,42 @@ from unittest.mock import AsyncMock, patch
|
|||
import pytest
|
||||
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
|
||||
AgentAccess,
|
||||
AgentRequestHandler,
|
||||
RestrictedAgentAccess,
|
||||
UnrestrictedAgentAccess,
|
||||
accessible_agents,
|
||||
)
|
||||
|
||||
|
||||
def _registry_with(*agent_names: str) -> AgentRegistry:
|
||||
registry: Final = AgentRegistry()
|
||||
registry.load_agents_from_config(
|
||||
[
|
||||
{
|
||||
"agent_name": name,
|
||||
"agent_card_params": {"name": name, "url": "http://localhost", "version": "1.0.0"},
|
||||
}
|
||||
for name in agent_names
|
||||
]
|
||||
)
|
||||
return registry
|
||||
|
||||
|
||||
def _agent_id(registry: AgentRegistry, agent_name: str) -> str:
|
||||
agent: Final = registry.get_agent_by_name(agent_name)
|
||||
assert agent is not None
|
||||
return agent.agent_id
|
||||
|
||||
|
||||
async def _single_context(user_api_key_auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
|
||||
return [user_api_key_auth]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestAgentRequestHandler:
|
||||
"""
|
||||
|
|
@ -265,6 +292,78 @@ class TestAgentRequestHandler:
|
|||
)
|
||||
assert result == UnrestrictedAgentAccess()
|
||||
|
||||
async def test_accessible_agents_hides_ungranted_agents_from_non_admins(self):
|
||||
"""LIT-6862: a key with no agent grant on itself or its team must list nothing,
|
||||
while a proxy admin with the same lack of grants still lists every agent."""
|
||||
registry: Final = _registry_with("alpha", "beta")
|
||||
internal_user: Final = UserAPIKeyAuth(
|
||||
api_key="test-key", user_id="alice", team_id="team-no-perms", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
proxy_admin: Final = UserAPIKeyAuth(
|
||||
api_key="admin-key", user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
async def no_grant_anywhere(user_api_key_auth: UserAPIKeyAuth) -> AgentAccess:
|
||||
return UnrestrictedAgentAccess()
|
||||
|
||||
assert (
|
||||
await accessible_agents(internal_user, registry.get_agent_list(), no_grant_anywhere, _single_context) == ()
|
||||
)
|
||||
assert {
|
||||
agent.agent_name
|
||||
for agent in await accessible_agents(
|
||||
proxy_admin, registry.get_agent_list(), no_grant_anywhere, _single_context
|
||||
)
|
||||
} == {"alpha", "beta"}
|
||||
|
||||
async def test_accessible_agents_lists_only_granted_agents(self):
|
||||
"""A grant for one agent lists that agent and hides the ungranted one."""
|
||||
registry: Final = _registry_with("alpha", "beta")
|
||||
granted_user: Final = UserAPIKeyAuth(
|
||||
api_key="test-key", user_id="bob", team_id="team-granted", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
async def alpha_only(user_api_key_auth: UserAPIKeyAuth) -> AgentAccess:
|
||||
return RestrictedAgentAccess(frozenset({_agent_id(registry, "alpha")}))
|
||||
|
||||
listed: Final = await accessible_agents(granted_user, registry.get_agent_list(), alpha_only, _single_context)
|
||||
assert [agent.agent_name for agent in listed] == ["alpha"]
|
||||
|
||||
async def test_accessible_agents_resolves_dashboard_session_through_real_teams_and_user(self):
|
||||
"""LIT-6862: a dashboard session carries the shared litellm-dashboard team id, which holds no
|
||||
grants. Listing must union the grants of the user's real teams and of the user row instead
|
||||
of treating the session as ungranted or as unrestricted."""
|
||||
registry: Final = _registry_with("alpha", "beta", "gamma")
|
||||
session: Final = UserAPIKeyAuth(
|
||||
api_key="session-key",
|
||||
user_id="alice",
|
||||
team_id=UI_SESSION_TOKEN_TEAM_ID,
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
admitted_user: Final = UserAPIKeyAuth(user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
grants: Final = {
|
||||
"team-granted": RestrictedAgentAccess(frozenset({_agent_id(registry, "alpha")})),
|
||||
"team-no-perms": UnrestrictedAgentAccess(),
|
||||
UI_SESSION_TOKEN_TEAM_ID: UnrestrictedAgentAccess(),
|
||||
}
|
||||
|
||||
async def effective_contexts(user_api_key_auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
|
||||
assert user_api_key_auth is session
|
||||
return [
|
||||
session.model_copy(update={"team_id": "team-granted"}),
|
||||
session.model_copy(update={"team_id": "team-no-perms"}),
|
||||
admitted_user,
|
||||
]
|
||||
|
||||
async def resolve_access(user_api_key_auth: UserAPIKeyAuth) -> AgentAccess:
|
||||
if user_api_key_auth is admitted_user:
|
||||
return RestrictedAgentAccess(frozenset({_agent_id(registry, "beta")}))
|
||||
assert user_api_key_auth.team_id is not None
|
||||
return grants[user_api_key_auth.team_id]
|
||||
|
||||
listed: Final = await accessible_agents(session, registry.get_agent_list(), resolve_access, effective_contexts)
|
||||
assert {agent.agent_name for agent in listed} == {"alpha", "beta"}
|
||||
|
||||
async def test_get_allowed_agents_for_key_via_access_group_ids(self):
|
||||
"""
|
||||
Test that _get_allowed_agents_for_key includes agents from key's access_group_ids
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
|||
from litellm.proxy.agent_endpoints import endpoints as agent_endpoints
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
|
||||
RestrictedAgentAccess,
|
||||
UnrestrictedAgentAccess,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.endpoints import (
|
||||
_attach_keys_to_agents,
|
||||
|
|
@ -550,9 +549,9 @@ class TestAgentRBACProxyAdminViewOnly:
|
|||
self.allowed_agents_spy.assert_awaited_once()
|
||||
|
||||
def test_should_still_redact_secrets_for_view_only_admin(self):
|
||||
"""An unrestricted viewer sees the same agents as an admin but with keys
|
||||
"""A viewer granted every agent sees the same agents as an admin but with keys
|
||||
stripped; litellm_params secrets never appear in either response."""
|
||||
self.allowed_agents_spy.return_value = UnrestrictedAgentAccess()
|
||||
self.allowed_agents_spy.return_value = RestrictedAgentAccess(frozenset({"agent-1", "agent-2"}))
|
||||
viewer_resp = self._list_agents(self.viewer_client)
|
||||
admin_resp = self._list_agents(self.admin_client)
|
||||
|
||||
|
|
|
|||
|
|
@ -1349,6 +1349,38 @@ async def test_vector_store_access_check_early_returns(
|
|||
assert result == expected_result
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_access_check_skips_db_lookup_when_no_vector_stores_requested():
|
||||
"""Registry returns [] (not None) for plain requests; no object permission DB lookup should happen."""
|
||||
valid_token = UserAPIKeyAuth(token="test-token", object_permission_id="perm-123")
|
||||
team_object = MagicMock()
|
||||
team_object.object_permission_id = "team-permission"
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
find_unique = AsyncMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = find_unique
|
||||
|
||||
mock_vector_store_registry = MagicMock()
|
||||
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = []
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: production auth reads these module globals; no dependency injection seam exists
|
||||
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
|
||||
),
|
||||
patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists
|
||||
"litellm.vector_store_registry", mock_vector_store_registry
|
||||
),
|
||||
):
|
||||
result = await vector_store_access_check(
|
||||
request_body={"messages": [{"role": "user", "content": "test"}]},
|
||||
team_object=team_object,
|
||||
valid_token=valid_token,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
find_unique.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"object_permissions,vector_store_ids,should_raise,error_type",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -2910,6 +2910,7 @@ async def test_update_team_with_team_member_budget_duration(
|
|||
"metadata": {"team_member_budget_id": "budget_123"},
|
||||
}
|
||||
mock_existing_team.metadata = {"team_member_budget_id": "budget_123"}
|
||||
mock_existing_team.members_with_roles = []
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(
|
||||
return_value=mock_existing_team
|
||||
)
|
||||
|
|
@ -11290,6 +11291,78 @@ async def test_patch_preserves_required_metadata_key_that_post_would_wipe():
|
|||
assert patch_meta == {"cost_center": "FINOPS-1", "team_notes": "edited"} # preserved by PATCH
|
||||
|
||||
|
||||
_STORED_METADATA_WITH_BUDGET: Final = {
|
||||
"team_member_budget_id": "budget-existing-123",
|
||||
"team_member_key_duration": "30d",
|
||||
"logging": [{"callback_name": "langfuse", "callback_type": "success"}],
|
||||
"cost_center": "cc-1234",
|
||||
}
|
||||
|
||||
|
||||
async def _written_metadata_with_budget(kind, body):
|
||||
"""Like ``_written_metadata`` but the team already owns a member budget row."""
|
||||
from litellm.proxy._types import LiteLLM_BudgetTable
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
patch( # test-quality-ok: update_team imports update_budget at call time; the module attribute is its only seam
|
||||
"litellm.proxy.management_endpoints.budget_management_endpoints.update_budget",
|
||||
AsyncMock(return_value=LiteLLM_BudgetTable(budget_id="budget-existing-123")),
|
||||
),
|
||||
):
|
||||
return await _written_metadata(kind, dict(_STORED_METADATA_WITH_BUDGET), body)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ["post", "patch"])
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
{"team_member_budget": 50.0},
|
||||
{"team_member_budget_duration": "1d"},
|
||||
{"team_member_tpm_limit": 500},
|
||||
{"team_member_rpm_limit": 5},
|
||||
],
|
||||
ids=lambda body: next(iter(body)),
|
||||
)
|
||||
async def test_team_member_budget_only_update_preserves_stored_metadata(kind, body):
|
||||
"""LIT-5150: a budget-only update that omits ``metadata`` must not replace the
|
||||
stored metadata JSON with just ``{"team_member_budget_id": ...}``."""
|
||||
assert await _written_metadata_with_budget(kind, body) == _STORED_METADATA_WITH_BUDGET
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ["post", "patch"])
|
||||
async def test_team_member_key_duration_only_update_preserves_stored_metadata(kind):
|
||||
"""LIT-5150: a metadata-backed field sent alone is merged into the stored
|
||||
metadata instead of becoming the whole metadata JSON."""
|
||||
written = await _written_metadata_with_budget(kind, {"team_member_key_duration": "7d"})
|
||||
|
||||
assert written == {**_STORED_METADATA_WITH_BUDGET, "team_member_key_duration": "7d"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("kind", ["post", "patch"])
|
||||
async def test_explicit_null_metadata_with_budget_field_still_clears_metadata(kind):
|
||||
"""``metadata: null`` is an explicit clear, so only the server-owned budget link survives."""
|
||||
written = await _written_metadata_with_budget(kind, {"metadata": None, "team_member_budget": 7.0})
|
||||
|
||||
assert written == {"team_member_budget_id": "budget-existing-123"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_metadata_only_update_keeps_team_member_budget_link():
|
||||
"""LIT-5150: rewriting metadata without any team member field must not drop the
|
||||
server-owned ``team_member_budget_id``, or the member budget silently resets."""
|
||||
body = {"metadata": {"cost_center": "cc-9999"}}
|
||||
|
||||
post_meta = await _written_metadata_with_budget("post", body)
|
||||
patch_meta = await _written_metadata_with_budget("patch", body)
|
||||
|
||||
assert post_meta == {"cost_center": "cc-9999", "team_member_budget_id": "budget-existing-123"}
|
||||
assert patch_meta == {**_STORED_METADATA_WITH_BUDGET, "cost_center": "cc-9999"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"body, field, expected",
|
||||
|
|
@ -11318,17 +11391,15 @@ async def test_top_level_fields_identical_post_and_patch(body, field, expected):
|
|||
@pytest.mark.asyncio
|
||||
async def test_patch_strips_system_managed_metadata_key_like_post():
|
||||
"""A caller cannot inject/overwrite server-owned keys via PATCH any more than
|
||||
via POST: team_member_budget_id is stripped from the write in both."""
|
||||
via POST: the stored team_member_budget_id wins over the caller's value in both."""
|
||||
existing = {"team_member_budget_id": "budget-123", "cost_center": "1234"}
|
||||
body = {"metadata": {"team_member_budget_id": "HACKED", "cost_center": "9999"}}
|
||||
|
||||
post_meta = await _written_metadata("post", existing, body)
|
||||
patch_meta = await _written_metadata("patch", existing, body)
|
||||
|
||||
assert "team_member_budget_id" not in post_meta
|
||||
assert "team_member_budget_id" not in patch_meta
|
||||
assert post_meta == {"cost_center": "9999"}
|
||||
assert patch_meta == {"cost_center": "9999"}
|
||||
assert post_meta == {"cost_center": "9999", "team_member_budget_id": "budget-123"}
|
||||
assert patch_meta == {"cost_center": "9999", "team_member_budget_id": "budget-123"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -7,7 +7,6 @@ import re
|
|||
import stat
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -16,11 +15,15 @@ COMPONENT_ENTRYPOINT = REPO_ROOT / "docker" / "component_entrypoint.sh"
|
|||
PROD_ENTRYPOINT = REPO_ROOT / "docker" / "prod_entrypoint.sh"
|
||||
GATEWAY_DOCKERFILE = REPO_ROOT / "gateway" / "Dockerfile"
|
||||
BACKEND_DOCKERFILE = REPO_ROOT / "backend" / "Dockerfile"
|
||||
BUILD_FROM_PIP_DOCKERFILE = REPO_ROOT / "docker" / "build_from_pip" / "Dockerfile.build_from_pip"
|
||||
TERRAFORM_ECS = REPO_ROOT / "terraform" / "litellm" / "aws" / "ecs.tf"
|
||||
TERRAFORM_CLOUDRUN = REPO_ROOT / "terraform" / "litellm" / "gcp" / "cloudrun.tf"
|
||||
|
||||
IMAGE_ENTRYPOINT_PATH = "/app/docker/component_entrypoint.sh"
|
||||
|
||||
TRUTHY_USE_DDTRACE = ("true", "True", "TRUE", "tRuE")
|
||||
FALSY_USE_DDTRACE = (None, "", "false", "False", "1", "yes", "on", "truex")
|
||||
|
||||
PYTHONPATH_SENTINEL = "/lit-entrypoint-sentinel:/app"
|
||||
|
||||
_STUB_TEMPLATE = """#!/bin/sh
|
||||
|
|
@ -34,6 +37,7 @@ _STUB_TEMPLATE = """#!/bin/sh
|
|||
|
||||
_ENTRYPOINT_RE = re.compile(r"^ENTRYPOINT\s+(\[.*\])\s*$", re.MULTILINE)
|
||||
_CMD_RE = re.compile(r"^CMD\s+(\[.*\])\s*$", re.MULTILINE)
|
||||
_COPY_RE = re.compile(r"^COPY\s+(?!--from)(\S+)\s+(\S+)\s*$", re.MULTILINE)
|
||||
_APP_TARGET_RE = re.compile(r"(?:gateway|backend)\.main:app")
|
||||
_TF_STRING_LOCAL_RE = re.compile(r'^\s*(\w+)\s*=\s*"((?:[^"\\]|\\.)*)"\s*$', re.MULTILINE)
|
||||
_TF_INTERPOLATION_RE = re.compile(r"\$\{(local|var)\.(\w+)\}")
|
||||
|
|
@ -53,7 +57,7 @@ def _write_stubs(bin_dir: Path, names: tuple[str, ...]) -> None:
|
|||
def _run_entrypoint(
|
||||
script: Path,
|
||||
argv: tuple[str, ...],
|
||||
use_ddtrace: Optional[str],
|
||||
use_ddtrace: str | None,
|
||||
tmp_path: Path,
|
||||
) -> tuple[str, ...]:
|
||||
"""Run `script` with stubbed executables on PATH and return the recorded lines."""
|
||||
|
|
@ -84,7 +88,7 @@ def _run_entrypoint(
|
|||
return tuple(record.read_text().splitlines()) if record.exists() else ()
|
||||
|
||||
|
||||
def _run_shell_command(command: str, bin_dir: Path, record: Path, use_ddtrace: Optional[str]) -> tuple[str, ...]:
|
||||
def _run_shell_command(command: str, bin_dir: Path, record: Path, use_ddtrace: str | None) -> tuple[str, ...]:
|
||||
"""Run a resolved Terraform launch command through `sh -c` and return the recorded lines."""
|
||||
env = {
|
||||
**os.environ,
|
||||
|
|
@ -184,9 +188,19 @@ def test_ddtrace_disabled_execs_the_command_directly(tmp_path: Path) -> None:
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("use_ddtrace", [None, "", "false", "True", "TRUE", "1", "yes"])
|
||||
def test_gating_matches_the_monolithic_entrypoint(use_ddtrace: Optional[str], tmp_path: Path) -> None:
|
||||
"""The componentized images must honor `USE_DDTRACE` exactly as the monolith does."""
|
||||
@pytest.mark.parametrize(
|
||||
"use_ddtrace, traced",
|
||||
[*((v, True) for v in TRUTHY_USE_DDTRACE), *((v, False) for v in FALSY_USE_DDTRACE)],
|
||||
)
|
||||
def test_gating_matches_the_monolithic_entrypoint_and_get_secret_bool(
|
||||
use_ddtrace: str | None, traced: bool, tmp_path: Path
|
||||
) -> None:
|
||||
"""Both entrypoints must accept exactly the spellings `get_secret_bool` accepts.
|
||||
|
||||
`ProxyStartupEvent._init_dd_tracer` reads `USE_DDTRACE` through `get_secret_bool`, which
|
||||
matches `true` case-insensitively. If the shell gate were stricter, `USE_DDTRACE=True` would
|
||||
give in-process LLM spans without `ddtrace-run` HTTP spans, a half-enabled state.
|
||||
"""
|
||||
component = _run_entrypoint(
|
||||
COMPONENT_ENTRYPOINT,
|
||||
("uvicorn", "gateway.main:app"),
|
||||
|
|
@ -200,8 +214,56 @@ def test_gating_matches_the_monolithic_entrypoint(use_ddtrace: Optional[str], tm
|
|||
tmp_path=tmp_path / "monolith",
|
||||
)
|
||||
|
||||
assert component[0].startswith("exec=") and monolith[0].startswith("exec=")
|
||||
assert (component[0] == "exec=ddtrace-run") == (monolith[0] == "exec=ddtrace-run")
|
||||
expected_exec = "exec=ddtrace-run" if traced else "exec=uvicorn"
|
||||
expected_openai = "DD_TRACE_OPENAI_ENABLED=False" if traced else "DD_TRACE_OPENAI_ENABLED=<unset>"
|
||||
assert component[0] == expected_exec
|
||||
assert component[2] == expected_openai
|
||||
assert monolith[0] == ("exec=ddtrace-run" if traced else "exec=litellm")
|
||||
assert monolith[2] == expected_openai
|
||||
assert monolith[1] == ("args=litellm --port 4000" if traced else "args=--port 4000")
|
||||
|
||||
|
||||
def _copied_script(dockerfile: Path, image_path: str) -> Path:
|
||||
"""Resolve the repo file a Dockerfile `COPY`s to `image_path`, so tests run what the image ships."""
|
||||
matches = _COPY_RE.findall(dockerfile.read_text())
|
||||
sources = tuple(src for src, dst in matches if dst == image_path)
|
||||
assert sources, f"{dockerfile} never COPYs anything to {image_path}"
|
||||
source = REPO_ROOT / sources[-1]
|
||||
assert source.is_file(), f"{dockerfile} COPYs {sources[-1]}, which does not exist in the build context"
|
||||
return source
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"use_ddtrace, traced",
|
||||
[*((v, True) for v in TRUTHY_USE_DDTRACE), *((v, False) for v in FALSY_USE_DDTRACE)],
|
||||
)
|
||||
def test_build_from_pip_image_launches_litellm_through_the_prod_entrypoint(
|
||||
use_ddtrace: str | None, traced: bool, tmp_path: Path
|
||||
) -> None:
|
||||
"""Run the build_from_pip image's ENTRYPOINT + CMD through the script it actually COPYs.
|
||||
|
||||
The image used to `ENTRYPOINT ["litellm"]`, so `USE_DDTRACE` was inert there at any spelling.
|
||||
Resolving the ENTRYPOINT path back to its COPY source and executing it with the Dockerfile's
|
||||
CMD checks the launch the container performs, not just that the Dockerfile mentions the script.
|
||||
"""
|
||||
entrypoint = _entrypoint_argv(BUILD_FROM_PIP_DOCKERFILE)
|
||||
assert len(entrypoint) == 1, (
|
||||
f"{BUILD_FROM_PIP_DOCKERFILE} ENTRYPOINT must be the bare script so CMD reaches litellm"
|
||||
)
|
||||
script = _copied_script(BUILD_FROM_PIP_DOCKERFILE, entrypoint[0])
|
||||
assert script == PROD_ENTRYPOINT, f"{BUILD_FROM_PIP_DOCKERFILE} bypasses the ddtrace-aware entrypoint"
|
||||
assert f"chmod +x {entrypoint[0]}" in BUILD_FROM_PIP_DOCKERFILE.read_text()
|
||||
|
||||
cmd = _cmd_argv(BUILD_FROM_PIP_DOCKERFILE)
|
||||
recorded = _run_entrypoint(script, cmd, use_ddtrace=use_ddtrace, tmp_path=tmp_path)
|
||||
|
||||
cmd_str = " ".join(cmd)
|
||||
assert recorded == (
|
||||
"exec=ddtrace-run" if traced else "exec=litellm",
|
||||
f"args=litellm {cmd_str}" if traced else f"args={cmd_str}",
|
||||
"DD_TRACE_OPENAI_ENABLED=False" if traced else "DD_TRACE_OPENAI_ENABLED=<unset>",
|
||||
f"PYTHONPATH={PYTHONPATH_SENTINEL}",
|
||||
)
|
||||
|
||||
|
||||
def test_entrypoint_script_is_executable() -> None:
|
||||
|
|
@ -239,9 +301,9 @@ def test_component_images_make_the_entrypoint_executable(dockerfile: Path) -> No
|
|||
|
||||
@pytest.mark.parametrize("terraform_file", TERRAFORM_LAUNCH_SITES, ids=lambda p: p.parent.name)
|
||||
@pytest.mark.parametrize("component", ["gateway", "backend"])
|
||||
@pytest.mark.parametrize("use_ddtrace", [None, "true", "false", "True"])
|
||||
@pytest.mark.parametrize("use_ddtrace", [*TRUTHY_USE_DDTRACE, *FALSY_USE_DDTRACE])
|
||||
def test_terraform_launch_command_matches_the_script_contract(
|
||||
terraform_file: Path, component: str, use_ddtrace: Optional[str], tmp_path: Path
|
||||
terraform_file: Path, component: str, use_ddtrace: str | None, tmp_path: Path
|
||||
) -> None:
|
||||
"""The Terraform command and `docker/component_entrypoint.sh` must decide identically.
|
||||
|
||||
|
|
@ -273,7 +335,7 @@ def test_terraform_launch_command_matches_the_script_contract(
|
|||
assert from_terraform[2] == from_script[2], f"{terraform_file} disagrees with the script on the openai integration"
|
||||
assert app_target in from_terraform[1]
|
||||
|
||||
if use_ddtrace == "true":
|
||||
if use_ddtrace in TRUTHY_USE_DDTRACE:
|
||||
assert from_terraform[0] == "exec=ddtrace-run"
|
||||
assert from_terraform[2] == "DD_TRACE_OPENAI_ENABLED=False"
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -14,7 +14,9 @@ from litellm.proxy.hooks.responses_id_security import (
|
|||
_is_responses_api_create_route,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
GenericEvent,
|
||||
ResponseCompletedEvent,
|
||||
ResponseCreatedEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
|
@ -691,6 +693,126 @@ class TestAsyncPostCallStreamingIteratorHook:
|
|||
assert not responses_id_security._is_encrypted_response_id(streamed_id)
|
||||
|
||||
|
||||
class TestStreamedGenericEventIdEncryption:
|
||||
"""A background stream carries event types with no typed model, which arrive as
|
||||
GenericEvent holding a plain dict. Those used to skip encryption while their typed
|
||||
siblings were encrypted, so one stream advertised two ids and the unencrypted one
|
||||
skipped the ownership check. Asserts the property rather than one event type: every
|
||||
id a client can see is the same encrypted id, and the raw one appears in no frame."""
|
||||
|
||||
RAW_ID = "resp_rawprovider123"
|
||||
|
||||
@staticmethod
|
||||
async def _agen(chunks):
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
@classmethod
|
||||
def _typed_event(cls, event_type):
|
||||
return {
|
||||
ResponsesAPIStreamEvents.RESPONSE_CREATED: ResponseCreatedEvent,
|
||||
ResponsesAPIStreamEvents.RESPONSE_COMPLETED: ResponseCompletedEvent,
|
||||
}[event_type](
|
||||
type=event_type,
|
||||
response=ResponsesAPIResponse(
|
||||
id=cls.RAW_ID,
|
||||
created_at=0,
|
||||
model="gpt-5.1",
|
||||
object="response",
|
||||
output=[],
|
||||
parallel_tool_calls=False,
|
||||
tool_choice="auto",
|
||||
tools=[],
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _background_stream(cls):
|
||||
return [
|
||||
cls._typed_event(ResponsesAPIStreamEvents.RESPONSE_CREATED),
|
||||
GenericEvent(
|
||||
type="response.queued",
|
||||
response={"id": cls.RAW_ID, "status": "queued"},
|
||||
),
|
||||
GenericEvent(type="keepalive"),
|
||||
GenericEvent(
|
||||
type="response.some_event_openai_adds_later",
|
||||
response={"id": cls.RAW_ID, "status": "in_progress"},
|
||||
),
|
||||
cls._typed_event(ResponsesAPIStreamEvents.RESPONSE_COMPLETED),
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _advertised_ids(events):
|
||||
nested = (getattr(event, "response", None) for event in events)
|
||||
return [
|
||||
payload["id"] if isinstance(payload, dict) else payload.id
|
||||
for payload in nested
|
||||
if payload is not None
|
||||
] + [
|
||||
event.id for event in events if isinstance(getattr(event, "id", None), str)
|
||||
]
|
||||
|
||||
async def _drain(self, responses_id_security, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-abcdefghij")
|
||||
|
||||
mock_auth = MagicMock()
|
||||
mock_auth.user_id = "user-a"
|
||||
mock_auth.team_id = "team-a"
|
||||
mock_auth.request_route = "/v1/responses"
|
||||
|
||||
return [
|
||||
out
|
||||
async for out in responses_id_security.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=mock_auth,
|
||||
response=self._agen(self._background_stream()),
|
||||
request_data={},
|
||||
)
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_every_event_advertises_the_same_encrypted_id(
|
||||
self, responses_id_security, monkeypatch
|
||||
):
|
||||
events = await self._drain(responses_id_security, monkeypatch)
|
||||
advertised = self._advertised_ids(events)
|
||||
|
||||
assert len(advertised) == 4
|
||||
assert len(set(advertised)) == 1
|
||||
|
||||
streamed_id = advertised[0]
|
||||
assert streamed_id != self.RAW_ID
|
||||
assert responses_id_security._is_encrypted_response_id(streamed_id)
|
||||
assert responses_id_security._decrypt_response_id(streamed_id) == (
|
||||
self.RAW_ID,
|
||||
"user-a",
|
||||
"team-a",
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_provider_id_never_reaches_the_client(
|
||||
self, responses_id_security, monkeypatch
|
||||
):
|
||||
events = await self._drain(responses_id_security, monkeypatch)
|
||||
|
||||
assert [self.RAW_ID in event.model_dump_json() for event in events] == [
|
||||
False
|
||||
] * len(events)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sibling_fields_survive_the_rewrite(
|
||||
self, responses_id_security, monkeypatch
|
||||
):
|
||||
_, queued, keepalive, later, _ = await self._drain(
|
||||
responses_id_security, monkeypatch
|
||||
)
|
||||
|
||||
assert queued.response["status"] == "queued"
|
||||
assert later.response["status"] == "in_progress"
|
||||
assert keepalive.type == "keepalive"
|
||||
assert getattr(keepalive, "response", None) is None
|
||||
|
||||
|
||||
class TestAsyncPostCallSuccessHook:
|
||||
"""Test async_post_call_success_hook function"""
|
||||
|
||||
|
|
|
|||
|
|
@ -7392,6 +7392,71 @@ def test_get_configured_token_limits_coerces_numeric_strings():
|
|||
assert router.get_configured_token_limits("quoted-limits-model") == (32000, 8000)
|
||||
|
||||
|
||||
def test_get_configured_mode_reads_deployment_model_info():
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "tts-model",
|
||||
"litellm_params": {"model": "openai/some-unmapped-tts-model"},
|
||||
"model_info": {"mode": "audio_speech"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert router.get_configured_mode("tts-model") == "audio_speech"
|
||||
|
||||
|
||||
def test_get_configured_mode_returns_none_for_unset_or_unknown():
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "no-mode-model",
|
||||
"litellm_params": {"model": "openai/some-unmapped-model"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert router.get_configured_mode("no-mode-model") is None
|
||||
assert router.get_configured_mode("not-a-real-model") is None
|
||||
|
||||
|
||||
def test_get_configured_mode_skips_wildcard_pattern_matching():
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "bedrock/*",
|
||||
"litellm_params": {"model": "bedrock/*"},
|
||||
"model_info": {"mode": "chat"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
router.pattern_router, "route", side_effect=AssertionError("pattern route called")
|
||||
):
|
||||
assert (
|
||||
router.get_configured_mode("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0")
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_get_configured_mode_treats_malformed_values_as_absent():
|
||||
malformed = ["", " ", 12345, ["chat"], {"mode": "chat"}, True]
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": f"bad-mode-{i}",
|
||||
"litellm_params": {"model": "openai/some-unmapped-model"},
|
||||
"model_info": {"mode": bad},
|
||||
}
|
||||
for i, bad in enumerate(malformed)
|
||||
]
|
||||
)
|
||||
|
||||
for i in range(len(malformed)):
|
||||
assert router.get_configured_mode(f"bad-mode-{i}") is None
|
||||
|
||||
|
||||
def test_get_configured_display_name_reads_deployment_model_info():
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
|
|
@ -12695,33 +12760,3 @@ async def test_prompt_management_factory_marks_injection_for_every_deployment(mo
|
|||
bucket = captured.get("litellm_metadata") or captured["metadata"]
|
||||
assert captured["model_info"]["id"] == "provisional-dep"
|
||||
assert bucket["litellm_gateway_injected_cache"] == ""
|
||||
|
||||
|
||||
def test_get_configured_mode_reads_deployment_model_info():
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "my-tts",
|
||||
"litellm_params": {"model": "openai/some-unmapped-mode-model"},
|
||||
"model_info": {"mode": "audio_speech"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert router.get_configured_mode("my-tts") == "audio_speech"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_info", [{}, {"mode": ""}, {"mode": " "}, {"mode": 123}])
|
||||
def test_get_configured_mode_returns_none_for_unset_blank_or_unknown(model_info):
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "plain-model",
|
||||
"litellm_params": {"model": "openai/some-unmapped-mode-model"},
|
||||
"model_info": model_info,
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
assert router.get_configured_mode("plain-model") is None
|
||||
assert router.get_configured_mode("unknown-model") is None
|
||||
|
|
|
|||
|
|
@ -27,10 +27,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16476
|
||||
"limit": 16470
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5518
|
||||
"limit": 5516
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4489
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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)}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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)}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -2134,6 +2134,7 @@ interface UiSpendLogsParams {
|
|||
exclude_internal_health_checks?: boolean;
|
||||
group_by_session?: boolean;
|
||||
session_cursor?: string;
|
||||
search?: string;
|
||||
}
|
||||
|
||||
interface UiSpendLogsCallOptions {
|
||||
|
|
@ -6641,6 +6642,7 @@ interface UiAuditLogsParams {
|
|||
changed_by_api_key?: string;
|
||||
object_team_id?: string;
|
||||
object_key_hash?: string;
|
||||
search?: string | null;
|
||||
sort_by?: string;
|
||||
sort_order?: "asc" | "desc";
|
||||
}
|
||||
|
|
@ -8139,15 +8141,18 @@ export const fetchMemoryList = async (
|
|||
options: {
|
||||
key?: string;
|
||||
keyPrefix?: string;
|
||||
search?: string;
|
||||
page?: number;
|
||||
pageSize?: number;
|
||||
} = {},
|
||||
): Promise<MemoryListResponse> => {
|
||||
const base = proxyBaseUrl ? `${proxyBaseUrl}/v1/memory` : `/v1/memory`;
|
||||
const params = new URLSearchParams();
|
||||
// keyPrefix takes precedence — backend also does, but we omit `key`
|
||||
// Backend precedence is search > key_prefix > key; only the winner is sent
|
||||
// to keep the URL clean and intent obvious.
|
||||
if (options.keyPrefix) {
|
||||
if (options.search) {
|
||||
params.append("search", options.search);
|
||||
} else if (options.keyPrefix) {
|
||||
params.append("key_prefix", options.keyPrefix);
|
||||
} else if (options.key) {
|
||||
params.append("key", options.key);
|
||||
|
|
|
|||
|
|
@ -1538,6 +1538,35 @@ describe("TeamInfoView", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("team member settings", () => {
|
||||
it("should populate Default Key Duration from the team's stored metadata", async () => {
|
||||
const user = userEvent.setup({ delay: null });
|
||||
vi.mocked(networking.teamInfoCall).mockResolvedValue(
|
||||
createMockTeamData({ metadata: { team_member_key_duration: "30d" } }),
|
||||
);
|
||||
|
||||
renderWithProviders(<TeamInfoView {...defaultProps} />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.queryAllByText("Test Team").length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
await user.click(screen.getByRole("tab", { name: "Settings" }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /edit settings/i }));
|
||||
|
||||
await user.click(await screen.findByRole("button", { name: /team member settings/i }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByLabelText(/^Default Key Duration/)).toHaveValue("30d");
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
describe("guardrails dropdown grouping", () => {
|
||||
const guardrail = (name: string, defaultOn: boolean) => ({
|
||||
guardrail_name: name,
|
||||
|
|
|
|||
|
|
@ -335,7 +335,7 @@ const toTeamFormValues = (info: TeamInfoRecord, effectiveGuardrails: string[]):
|
|||
default_team_member_models: info.default_team_member_models || [],
|
||||
team_member_budget: info.team_member_budget_table?.max_budget,
|
||||
team_member_budget_duration: info.team_member_budget_table?.budget_duration,
|
||||
team_member_key_duration: info.team_member_key_duration,
|
||||
team_member_key_duration: info.metadata?.team_member_key_duration,
|
||||
team_member_tpm_limit: info.team_member_budget_table?.tpm_limit,
|
||||
team_member_rpm_limit: info.team_member_budget_table?.rpm_limit,
|
||||
budget_duration: info.budget_duration,
|
||||
|
|
|
|||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
</>
|
||||
|
|
|
|||
|
|
@ -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]);
|
||||
},
|
||||
);
|
||||
});
|
||||
|
|
@ -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}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -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 });
|
||||
|
|
|
|||
|
|
@ -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)}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -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 }));
|
||||
|
|
|
|||
|
|
@ -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)}
|
||||
|
|
|
|||
|
|
@ -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 }) => {
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
10
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
10
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -41118,6 +41118,8 @@ export interface operations {
|
|||
object_team_id?: string | null;
|
||||
/** @description Filter by token (key hash) present in before_value or updated_values JSON (PostgreSQL only) */
|
||||
object_key_hash?: string | null;
|
||||
/** @description Match a row whose id, object_id, changed_by, or changed_by_api_key equals this value */
|
||||
search?: string | null;
|
||||
/** @description Column to sort by (e.g. 'updated_at', 'action', 'table_name') */
|
||||
sort_by?: string | null;
|
||||
/** @description Sort order ('asc' or 'desc') */
|
||||
|
|
@ -49868,6 +49870,8 @@ export interface operations {
|
|||
key_hash?: string | null;
|
||||
/** @description Filter keys by key alias. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching. */
|
||||
key_alias?: string | null;
|
||||
/** @description Combined search: matches keys whose token (key hash) equals the value OR whose key_alias contains it (case-insensitive). */
|
||||
search?: string | null;
|
||||
/** @description Return full key object */
|
||||
return_full_object?: boolean;
|
||||
/** @description Include all keys for teams that user is an admin of. */
|
||||
|
|
@ -57159,6 +57163,8 @@ export interface operations {
|
|||
group_by_session?: boolean;
|
||||
/** @description Keyset cursor '<last_activity>|<api_key>|<session_key>' from a previous group_by_session page. UI route only, honored when sorting by startTime */
|
||||
session_cursor?: string | null;
|
||||
/** @description Match a log whose request_id, api_key (hash), team_id, user, end_user, session_id, or model_id equals this value. request_id matches across all time; the other columns match inside start_date/end_date, which stay required */
|
||||
search?: string | null;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
|
|
@ -57275,6 +57281,8 @@ export interface operations {
|
|||
group_by_session?: boolean;
|
||||
/** @description Keyset cursor '<last_activity>|<api_key>|<session_key>' from a previous group_by_session page. UI route only, honored when sorting by startTime */
|
||||
session_cursor?: string | null;
|
||||
/** @description Match a log whose request_id, api_key (hash), team_id, user, end_user, session_id, or model_id equals this value. request_id matches across all time; the other columns match inside start_date/end_date, which stay required */
|
||||
search?: string | null;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
|
|
@ -63587,6 +63595,8 @@ export interface operations {
|
|||
key?: string | null;
|
||||
/** @description Filter by key prefix (Redis-style namespace scan). Mutually exclusive with `key`; if both are provided, `key_prefix` wins. */
|
||||
key_prefix?: string | null;
|
||||
/** @description Match entries whose key starts with this value or whose memory_id equals it. Takes precedence over `key_prefix` and `key` when provided. */
|
||||
search?: string | null;
|
||||
page?: number;
|
||||
page_size?: number;
|
||||
};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue