mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/playground-model-settings-89dd7d
This commit is contained in:
commit
d79cd29e7c
43 changed files with 2421 additions and 194 deletions
|
|
@ -0,0 +1,8 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN "baseline_model" TEXT,
|
||||
ADD COLUMN "direction" TEXT NOT NULL DEFAULT 'forward';
|
||||
|
||||
DROP INDEX IF EXISTS "LiteLLM_ShadowEvalJob_one_active_per_key";
|
||||
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_ShadowEvalJob_one_active_per_key_direction"
|
||||
ON "LiteLLM_ShadowEvalJob"("api_key_id", "direction") WHERE "stopped_at" IS NULL;
|
||||
|
|
@ -1450,15 +1450,20 @@ model LiteLLM_AutoRouterSession {
|
|||
@@index([last_turn_at], map: "idx_autorouter_session_last_turn")
|
||||
}
|
||||
|
||||
// Shadow eval: pre-adoption evaluation of an auto-router against a key's live traffic.
|
||||
// A sampled slice of requests is duplicated through the router in a detached task and an
|
||||
// LLM judge compares real vs shadow responses blind. The job row is immutable config plus
|
||||
// Shadow eval: evaluation of an auto-router against a key's live traffic, in either
|
||||
// direction. forward duplicates the requests the key did not route through the router
|
||||
// through it, answering whether the key should adopt it; reverse duplicates the requests
|
||||
// the router did serve against a fixed baseline model, answering whether a key already on
|
||||
// it still benefits. Either way a sampled slice runs in a detached task and an LLM judge
|
||||
// compares real vs shadow responses blind. The job row is immutable config plus
|
||||
// stopped_at; every count, status, and spend figure is derived from the append-only
|
||||
// attempt rows, so nothing can disagree across pods or stop races.
|
||||
model LiteLLM_ShadowEvalJob {
|
||||
id String @id @default(cuid())
|
||||
api_key_id String // hashed virtual key whose traffic is shadowed
|
||||
router_name String
|
||||
router_name String // the auto-router under evaluation, in either direction
|
||||
direction String @default("forward") // forward | reverse
|
||||
baseline_model String? // reverse only: the fixed model the router is judged against
|
||||
judge_model String
|
||||
shadow_percentage Float
|
||||
max_turns Int // sample budget: judge at most this many turns
|
||||
|
|
|
|||
|
|
@ -1572,7 +1572,7 @@ class RedisCache(BaseCache):
|
|||
async def _pipeline_rpush_helper(
|
||||
self,
|
||||
pipe: pipeline,
|
||||
rpush_list: list[RedisPipelineRpushOperation],
|
||||
rpush_list: Sequence[RedisPipelineRpushOperation],
|
||||
) -> list[int]:
|
||||
"""Helper function for pipeline rpush operations"""
|
||||
for rpush_op in rpush_list:
|
||||
|
|
@ -1588,7 +1588,7 @@ class RedisCache(BaseCache):
|
|||
@_redis_circuit_breaker_guard
|
||||
async def async_rpush_pipeline(
|
||||
self,
|
||||
rpush_list: list[RedisPipelineRpushOperation],
|
||||
rpush_list: Sequence[RedisPipelineRpushOperation],
|
||||
) -> list[int]:
|
||||
"""
|
||||
Use Redis Pipelines for bulk RPUSH operations
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Shadow Eval Logger: samples a shadowed key's successful chat requests, duplicates each
|
||||
through the auto-router in a detached task, blind-judges real vs shadow, and appends one
|
||||
against the job's other arm in a detached task (the auto-router for a forward job, the
|
||||
fixed baseline model for a reverse one), blind-judges real vs shadow, and appends one
|
||||
``LiteLLM_ShadowEvalAttempt`` row (verdict or error) as the feature's only hot-path write.
|
||||
Counts, status, and spend derive from those rows at read time, so nothing can disagree
|
||||
across pods or stop races; the hook reads active jobs through a short-TTL cache."""
|
||||
|
|
@ -10,10 +11,12 @@ import random
|
|||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from itertools import groupby
|
||||
from operator import itemgetter
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError, field_validator, model_validator
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
|
|
@ -28,6 +31,7 @@ from litellm.litellm_core_utils.llm_judge import (
|
|||
parse_json_verdict,
|
||||
)
|
||||
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
|
||||
from litellm.types.management_endpoints.auto_router_endpoints import ShadowEvalDirection
|
||||
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -161,13 +165,26 @@ async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
def _routing_decision(metadata: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""The routing decision a pre-routing strategy wrote to a call's metadata, empty when
|
||||
a plain model served it. Read off the sampled request for the control arm, and off the
|
||||
shadow call's own write-back for the shadow arm."""
|
||||
decision: Final = metadata.get("routing_decision")
|
||||
return decision if isinstance(decision, Mapping) else _EMPTY_METADATA
|
||||
|
||||
|
||||
def _routed_tier(metadata: Mapping[str, object]) -> str | None:
|
||||
decision: Final = _routing_decision(metadata)
|
||||
raw: Final = decision.get("tier_label") or decision.get("tier")
|
||||
return str(raw) if raw is not None else None
|
||||
|
||||
|
||||
def _request_was_routed_by(request_metadata: Mapping[str, object], router_name: str) -> bool:
|
||||
"""Duplicating a request the shadowed router already served compares the router to
|
||||
itself: guaranteed ties, judge spend for zero information."""
|
||||
decision: Final = request_metadata.get("routing_decision")
|
||||
if not isinstance(decision, Mapping):
|
||||
return False
|
||||
return decision.get("router_model_name") == router_name
|
||||
"""Whether the router under evaluation served this request, which is what decides
|
||||
the direction it belongs to. A forward job skips its own router's traffic, since
|
||||
duplicating it would compare the router to itself: guaranteed ties, judge spend for
|
||||
zero information. A reverse job samples exactly that traffic and nothing else."""
|
||||
return _routing_decision(request_metadata).get("router_model_name") == router_name
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -197,22 +214,53 @@ class _JudgeVerdict:
|
|||
cost: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ActiveShadowEvalJob:
|
||||
"""One active job as the sampling path needs it: immutable config plus the attempt
|
||||
count as of the cache fill (the turn budget's staleness is bounded by the cache TTL)."""
|
||||
class ActiveShadowEvalJob(BaseModel):
|
||||
"""One active job as the sampling path needs it, validated straight off the untyped
|
||||
job row: immutable config plus the attempt count as of the cache fill (the turn
|
||||
budget's staleness is bounded by the cache TTL). Every way a row can be unsamplable
|
||||
is a validation error here, so a bad row is skipped rather than sampled wrongly."""
|
||||
|
||||
model_config = ConfigDict(frozen=True, from_attributes=True)
|
||||
|
||||
id: str
|
||||
router_name: str
|
||||
direction: ShadowEvalDirection = "forward"
|
||||
baseline_model: str | None = None
|
||||
shadow_percentage: float
|
||||
judge_model: str
|
||||
max_turns: int
|
||||
ends_at: datetime
|
||||
attempts: int
|
||||
attempts: int = 0
|
||||
|
||||
@field_validator("ends_at")
|
||||
@classmethod
|
||||
def _as_utc(cls, value: datetime) -> datetime:
|
||||
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _baseline_model_matches_direction(self) -> "ActiveShadowEvalJob":
|
||||
if (self.baseline_model is not None) != (self.direction == "reverse"):
|
||||
raise ValueError("baseline_model is set for exactly the reverse jobs")
|
||||
return self
|
||||
|
||||
@property
|
||||
def shadow_target(self) -> str:
|
||||
"""The model the duplicated arm calls: the router itself for a forward job, the
|
||||
fixed baseline for a reverse one. Total because the validator above pins
|
||||
baseline_model to reverse jobs and only those."""
|
||||
return self.baseline_model or self.router_name
|
||||
|
||||
|
||||
def _as_utc(value: datetime) -> datetime:
|
||||
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value
|
||||
def _as_active_job(record: object, attempts: int) -> ActiveShadowEvalJob | None:
|
||||
"""The sampling path's view of one job row, or None for a row it cannot sample: an
|
||||
unknown direction, or a reverse job with no baseline model to duplicate against.
|
||||
Failing closed here is what keeps the dispatch path total."""
|
||||
try:
|
||||
job: Final = ActiveShadowEvalJob.model_validate(record)
|
||||
except ValidationError as e:
|
||||
verbose_logger.debug("shadow_eval: skipping unsamplable job row: %s", e)
|
||||
return None
|
||||
return job.model_copy(update={"attempts": attempts})
|
||||
|
||||
|
||||
_jobs_cache: Final = InMemoryCache(max_size_in_memory=4, default_ttl=_JOBS_CACHE_TTL_SECONDS)
|
||||
|
|
@ -238,8 +286,9 @@ class ShadowEvalLogger(CustomLogger):
|
|||
# generation; the refill absorbs written rows and resets.
|
||||
self._job_starts: dict[str, int] = {} # mutable-ok: per-generation counter
|
||||
|
||||
async def _active_jobs(self) -> Mapping[str, ActiveShadowEvalJob]:
|
||||
"""Active jobs by api_key_id, cache-first. A DB fault returns empty without
|
||||
async def _active_jobs(self) -> Mapping[str, tuple[ActiveShadowEvalJob, ...]]:
|
||||
"""Active jobs by api_key_id, cache-first. A key holds at most one job per
|
||||
direction, so the value is a collection. A DB fault returns empty without
|
||||
caching, so sampling pauses for that request and the next one retries."""
|
||||
cached: Final = await self._jobs_cache.async_get_cache(_JOBS_CACHE_KEY)
|
||||
if cached is not None:
|
||||
|
|
@ -264,18 +313,19 @@ class ShadowEvalLogger(CustomLogger):
|
|||
else ()
|
||||
)
|
||||
attempt_counts: Final = {str(row["job_id"]): int(row["_count"]["_all"]) for row in grouped or []}
|
||||
jobs: Final = {
|
||||
str(record.api_key_id): ActiveShadowEvalJob(
|
||||
id=str(record.id),
|
||||
router_name=str(record.router_name),
|
||||
shadow_percentage=float(record.shadow_percentage),
|
||||
judge_model=str(record.judge_model),
|
||||
max_turns=int(record.max_turns),
|
||||
ends_at=_as_utc(record.ends_at),
|
||||
attempts=attempt_counts.get(str(record.id), 0),
|
||||
by_key: Final = tuple(
|
||||
sorted(
|
||||
(
|
||||
(str(record.api_key_id), job)
|
||||
for record in records or []
|
||||
if (job := _as_active_job(record, attempt_counts.get(str(record.id), 0))) is not None
|
||||
),
|
||||
key=itemgetter(0),
|
||||
)
|
||||
for record in records or []
|
||||
}
|
||||
)
|
||||
jobs: Final = MappingProxyType(
|
||||
{key: tuple(job for _, job in group) for key, group in groupby(by_key, key=itemgetter(0))}
|
||||
)
|
||||
await self._jobs_cache.async_set_cache(_JOBS_CACHE_KEY, jobs)
|
||||
self._job_starts = {} # rebind-ok: new generation, counts absorbed into the fill
|
||||
return jobs
|
||||
|
|
@ -308,43 +358,46 @@ class ShadowEvalLogger(CustomLogger):
|
|||
api_key_hash: Final = metadata.get("user_api_key_hash")
|
||||
if not api_key_hash:
|
||||
return
|
||||
job: Final = (await self._active_jobs()).get(str(api_key_hash))
|
||||
if job is None:
|
||||
return
|
||||
if datetime.now(timezone.utc) >= job.ends_at:
|
||||
return
|
||||
if job.attempts + self._job_starts.get(job.id, 0) >= job.max_turns:
|
||||
return
|
||||
request_id: Final = payload.get("id") or ""
|
||||
if not request_id:
|
||||
return
|
||||
if not _sample_hits(request_id, job.id, job.shadow_percentage):
|
||||
return
|
||||
if payload.get("call_type") not in _SAMPLED_CALL_TYPES:
|
||||
return # only known chat-shaped traffic is comparable; unknown or missing types fail closed
|
||||
if _request_was_routed_by(request_metadata, job.router_name):
|
||||
return
|
||||
if self._inflight_shadow_tasks >= _MAX_CONCURRENT_SHADOW_TASKS:
|
||||
return
|
||||
raw_messages: Final = kwargs.get("messages")
|
||||
self._job_starts[job.id] = self._job_starts.get(job.id, 0) + 1
|
||||
self._inflight_shadow_tasks += 1
|
||||
task: Final = asyncio.create_task(
|
||||
self._run_shadow_eval(
|
||||
job=job,
|
||||
request_id=request_id,
|
||||
messages=tuple(m for m in raw_messages if isinstance(m, Mapping))
|
||||
if isinstance(raw_messages, Sequence)
|
||||
else (),
|
||||
response_obj=response_obj,
|
||||
real_model=payload.get("model") or "",
|
||||
model_parameters=MappingProxyType(
|
||||
dict(payload.get("model_parameters") or {}) # mutable-ok: frozen snapshot
|
||||
),
|
||||
parent_metadata=MappingProxyType(dict(request_metadata)), # mutable-ok: frozen snapshot
|
||||
)
|
||||
messages: Final = (
|
||||
tuple(m for m in raw_messages if isinstance(m, Mapping)) if isinstance(raw_messages, Sequence) else ()
|
||||
)
|
||||
task.add_done_callback(self._release_shadow_slot)
|
||||
control_tier: Final = _routed_tier(request_metadata)
|
||||
# A key can hold one job per direction, and a request routed by one job's
|
||||
# router while bypassing the other's qualifies for both. Each is separately
|
||||
# budgeted, so both fire.
|
||||
for job in (await self._active_jobs()).get(str(api_key_hash), ()):
|
||||
if datetime.now(timezone.utc) >= job.ends_at:
|
||||
continue
|
||||
if job.attempts + self._job_starts.get(job.id, 0) >= job.max_turns:
|
||||
continue
|
||||
if not _sample_hits(request_id, job.id, job.shadow_percentage):
|
||||
continue
|
||||
if _request_was_routed_by(request_metadata, job.router_name) != (job.direction == "reverse"):
|
||||
continue
|
||||
if self._inflight_shadow_tasks >= _MAX_CONCURRENT_SHADOW_TASKS:
|
||||
return
|
||||
self._job_starts[job.id] = self._job_starts.get(job.id, 0) + 1
|
||||
self._inflight_shadow_tasks += 1
|
||||
asyncio.create_task(
|
||||
self._run_shadow_eval(
|
||||
job=job,
|
||||
request_id=request_id,
|
||||
messages=messages,
|
||||
response_obj=response_obj,
|
||||
real_model=payload.get("model") or "",
|
||||
control_tier=control_tier,
|
||||
model_parameters=MappingProxyType(
|
||||
dict(payload.get("model_parameters") or {}) # mutable-ok: frozen snapshot
|
||||
),
|
||||
parent_metadata=MappingProxyType(dict(request_metadata)), # mutable-ok: frozen snapshot
|
||||
)
|
||||
).add_done_callback(self._release_shadow_slot)
|
||||
except Exception as e: # noqa: BLE001 # logging hooks must never fail the request
|
||||
verbose_logger.debug("shadow_eval: failed to schedule task: %s", e)
|
||||
|
||||
|
|
@ -360,6 +413,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
messages: Sequence[Mapping[str, object]],
|
||||
response_obj: object,
|
||||
real_model: str,
|
||||
control_tier: str | None,
|
||||
model_parameters: Mapping[str, object],
|
||||
parent_metadata: Mapping[str, object],
|
||||
) -> None:
|
||||
|
|
@ -376,9 +430,11 @@ class ShadowEvalLogger(CustomLogger):
|
|||
if await _key_or_team_is_over_budget(parent_metadata):
|
||||
return
|
||||
|
||||
shadow: Final = await self._call_router_shadow(job.router_name, messages, model_parameters, parent_metadata)
|
||||
shadow: Final = await self._call_router_shadow(
|
||||
job.shadow_target, messages, model_parameters, parent_metadata
|
||||
)
|
||||
if isinstance(shadow, _CallFailure):
|
||||
await self._record_attempt(prisma, job, request_id, outcome="error", error=shadow.error)
|
||||
await self._record_attempt(prisma, job, request_id, control_tier, outcome="error", error=shadow.error)
|
||||
return
|
||||
|
||||
verdict: Final = await self._call_judge(
|
||||
|
|
@ -393,6 +449,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
prisma,
|
||||
job,
|
||||
request_id,
|
||||
control_tier,
|
||||
outcome="error",
|
||||
error=verdict.error,
|
||||
shadow=shadow,
|
||||
|
|
@ -403,6 +460,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
prisma,
|
||||
job,
|
||||
request_id,
|
||||
control_tier,
|
||||
outcome=verdict.preference,
|
||||
shadow=shadow,
|
||||
real_model=real_model,
|
||||
|
|
@ -411,13 +469,16 @@ class ShadowEvalLogger(CustomLogger):
|
|||
)
|
||||
except Exception as e: # noqa: BLE001 # detached task: record what happened, never raise
|
||||
verbose_logger.debug("shadow_eval: pipeline failed for %s: %s", request_id, e)
|
||||
await self._record_attempt(prisma, job, request_id, outcome="error", error=f"pipeline error: {e}")
|
||||
await self._record_attempt(
|
||||
prisma, job, request_id, control_tier, outcome="error", error=f"pipeline error: {e}"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _record_attempt(
|
||||
prisma: "PrismaClient | None",
|
||||
job: ActiveShadowEvalJob,
|
||||
request_id: str,
|
||||
control_tier: str | None,
|
||||
*,
|
||||
outcome: str,
|
||||
shadow: _ShadowResponse | None = None,
|
||||
|
|
@ -434,7 +495,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
"job_id": job.id,
|
||||
"request_id": request_id,
|
||||
"outcome": outcome,
|
||||
"tier": shadow.tier if shadow else None,
|
||||
"tier": control_tier if job.direction == "reverse" else (shadow.tier if shadow else None),
|
||||
"real_model": real_model or None,
|
||||
"shadow_model": shadow.model if shadow else None,
|
||||
"confidence": confidence,
|
||||
|
|
@ -447,14 +508,15 @@ class ShadowEvalLogger(CustomLogger):
|
|||
|
||||
async def _call_router_shadow(
|
||||
self,
|
||||
router_name: str,
|
||||
target_model: str,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
model_parameters: Mapping[str, object],
|
||||
parent_metadata: Mapping[str, object],
|
||||
) -> "_ShadowResponse | _CallFailure":
|
||||
"""Send the prompt through the auto-router being evaluated. The metadata carries
|
||||
the shadowed key's identity (spend attribution) and receives the router's routing
|
||||
decision write-back, read back for tier attribution."""
|
||||
"""Send the prompt through the arm nobody was served: the auto-router under
|
||||
evaluation, or a reverse job's fixed baseline model. The metadata carries the
|
||||
shadowed key's identity (spend attribution) and receives a routing decision
|
||||
write-back, which a plain baseline model simply never makes."""
|
||||
router: Final = self._router_provider()
|
||||
if router is None:
|
||||
return _CallFailure("no router configured on this pod")
|
||||
|
|
@ -466,7 +528,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
}
|
||||
try:
|
||||
response: Final = await router.acompletion(
|
||||
model=router_name,
|
||||
model=target_model,
|
||||
messages=messages, # pyright: ignore[reportArgumentType] # snapshot of the SDK's own message dicts
|
||||
metadata=shadow_metadata,
|
||||
num_retries=0,
|
||||
|
|
@ -479,13 +541,10 @@ class ShadowEvalLogger(CustomLogger):
|
|||
text: Final = self._extract_response_text(response)
|
||||
if not text:
|
||||
return _CallFailure("shadow router returned an empty response")
|
||||
raw_decision: Final = shadow_metadata.get("routing_decision")
|
||||
routing_decision: Final = raw_decision if isinstance(raw_decision, Mapping) else _EMPTY_METADATA
|
||||
raw_tier: Final = routing_decision.get("tier_label") or routing_decision.get("tier")
|
||||
return _ShadowResponse(
|
||||
text=text,
|
||||
model=str(getattr(response, "model", None) or routing_decision.get("routed_model") or ""),
|
||||
tier=str(raw_tier) if raw_tier is not None else None,
|
||||
model=str(getattr(response, "model", None) or _routing_decision(shadow_metadata).get("routed_model") or ""),
|
||||
tier=_routed_tier(shadow_metadata),
|
||||
)
|
||||
|
||||
async def _call_judge(
|
||||
|
|
@ -552,7 +611,7 @@ class ShadowEvalLogger(CustomLogger):
|
|||
return extract_text_from_content(content)
|
||||
|
||||
|
||||
_EMPTY_JOBS: Final[Mapping[str, ActiveShadowEvalJob]] = MappingProxyType({})
|
||||
_EMPTY_JOBS: Final[Mapping[str, tuple[ActiveShadowEvalJob, ...]]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _default_prisma_provider() -> "PrismaClient | None":
|
||||
|
|
|
|||
|
|
@ -37,9 +37,32 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM):
|
|||
super().__init__()
|
||||
|
||||
def get_vector_store_endpoints_by_type(self) -> VectorStoreIndexEndpoints:
|
||||
"""
|
||||
Every ``GET`` under ``/indexes/`` is a read: get details, stats, and the
|
||||
document reads (GET-form search, ``$count``, point lookup, and the
|
||||
GET forms of suggest and autocomplete).
|
||||
|
||||
``POST`` splits by endpoint. Search, suggest, autocomplete, and analyze
|
||||
are query endpoints, so they read; ``/docs/index`` is the batch endpoint
|
||||
carrying upload, merge, mergeOrUpload, and delete actions, so it writes.
|
||||
|
||||
Patterns stay literal rather than ``{placeholder}`` templates because the
|
||||
matcher falls back to the substring before a ``{``, which here is always
|
||||
``/indexes/``. The matcher is substring-based, so an index name may
|
||||
itself contain a read fragment (an index named ``analyze*`` puts
|
||||
``/analyze`` inside the batch-write path); writes are classified before
|
||||
reads, so such a path demands the write grant rather than being
|
||||
shadowed into a read.
|
||||
"""
|
||||
return {
|
||||
"read": [("GET", "/docs/search"), ("POST", "/docs/search")],
|
||||
"write": [("PUT", "/docs")],
|
||||
"read": [
|
||||
("GET", "/indexes/"),
|
||||
("POST", "/docs/search"),
|
||||
("POST", "/docs/suggest"),
|
||||
("POST", "/docs/autocomplete"),
|
||||
("POST", "/analyze"),
|
||||
],
|
||||
"write": [("POST", "/docs/index")],
|
||||
}
|
||||
|
||||
def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials:
|
||||
|
|
|
|||
|
|
@ -18,6 +18,16 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
_PERPLEXITY_UNIFIED_PARAMS: Final[frozenset[str]] = frozenset(
|
||||
(
|
||||
"max_results",
|
||||
"search_domain_filter",
|
||||
"country",
|
||||
"max_tokens_per_page",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _search_host(url: str) -> str:
|
||||
return urlsplit(url).netloc.lower()
|
||||
|
||||
|
|
@ -96,7 +106,7 @@ class BaseSearchConfig:
|
|||
return "POST"
|
||||
|
||||
@staticmethod
|
||||
def get_supported_perplexity_optional_params() -> set:
|
||||
def get_supported_perplexity_optional_params() -> frozenset[str]:
|
||||
"""
|
||||
Get the set of Perplexity unified search parameters.
|
||||
These are the standard parameters that providers should transform from.
|
||||
|
|
@ -104,12 +114,7 @@ class BaseSearchConfig:
|
|||
Returns:
|
||||
Set of parameter names that are part of the unified spec
|
||||
"""
|
||||
return {
|
||||
"max_results",
|
||||
"search_domain_filter",
|
||||
"country",
|
||||
"max_tokens_per_page",
|
||||
}
|
||||
return _PERPLEXITY_UNIFIED_PARAMS
|
||||
|
||||
def _assert_trusted_api_base_for_server_credential(
|
||||
self,
|
||||
|
|
|
|||
3
litellm/llms/nimble/__init__.py
Normal file
3
litellm/llms/nimble/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from litellm.llms.nimble.search.transformation import NimbleSearchConfig
|
||||
|
||||
__all__ = ("NimbleSearchConfig",)
|
||||
3
litellm/llms/nimble/search/__init__.py
Normal file
3
litellm/llms/nimble/search/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from litellm.llms.nimble.search.transformation import NimbleSearchConfig
|
||||
|
||||
__all__ = ("NimbleSearchConfig",)
|
||||
264
litellm/llms/nimble/search/transformation.py
Normal file
264
litellm/llms/nimble/search/transformation.py
Normal file
|
|
@ -0,0 +1,264 @@
|
|||
"""
|
||||
Calls Nimble's /v2/search endpoint to search the web.
|
||||
|
||||
Nimble API Reference: https://docs.nimbleway.com/api-reference/search/search
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.search.transformation import (
|
||||
BaseSearchConfig,
|
||||
SearchResponse,
|
||||
SearchResult,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
_NIMBLE_DOCS_URL: Final = "https://docs.nimbleway.com/api-reference/search/search"
|
||||
|
||||
|
||||
class _NimbleResult(BaseModel):
|
||||
"""One entry of Nimble's `results` array. Every field is optional so a single degraded
|
||||
result degrades to empty strings instead of failing the whole call."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
title: str | None = None
|
||||
url: str | None = None
|
||||
content: str | None = None
|
||||
description: str | None = None
|
||||
# Free-form per Nimble's schema, so an unexpected shape must not fail the search.
|
||||
additional_data: object = None
|
||||
|
||||
|
||||
class _NimbleSearchResponse(BaseModel):
|
||||
"""Nimble's /v2/search response envelope."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
# Required: a search with no hits returns `[]`, so a null or absent `results` means the
|
||||
# body is not a search response and must not be reported as a successful empty search.
|
||||
results: tuple[_NimbleResult, ...]
|
||||
|
||||
|
||||
class _AdditionalData(BaseModel):
|
||||
"""The slice of a result's free-form `additional_data` that maps onto SearchResult."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
publish_date: str | None = None
|
||||
|
||||
|
||||
class _ErrorEnvelope(BaseModel):
|
||||
"""Nimble reports errors as either `{"detail": ...}` (validation) or
|
||||
`{"success": "false", "task_id": ..., "message": ...}` (collection)."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
detail: str | None = None
|
||||
message: str | None = None
|
||||
|
||||
|
||||
_DomainListAdapter: Final = TypeAdapter(tuple[str, ...])
|
||||
|
||||
_NOTHING: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
|
||||
def _optional(key: str, value: object) -> Mapping[str, object]:
|
||||
"""A one-entry mapping to spread into a payload, or nothing when the value is absent."""
|
||||
return MappingProxyType({key: value}) if value is not None else _NOTHING
|
||||
|
||||
|
||||
class NimbleSearchConfig(BaseSearchConfig):
|
||||
NIMBLE_API_BASE = "https://sdk.nimbleway.com/v2"
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "Nimble"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict[str, str], # mutable-ok: BaseSearchConfig.validate_environment signature
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseSearchConfig.validate_environment signature
|
||||
) -> dict[str, str]: # mutable-ok: the http handler passes this straight to httpx as headers
|
||||
"""
|
||||
Validate environment and return headers.
|
||||
|
||||
Returns a new dict rather than mutating ``headers``: the http handler calls this
|
||||
a second time after ``litellm/search/main.py`` already did, so it has to be idempotent.
|
||||
"""
|
||||
resolved_api_key: Final = self.resolve_server_api_key(
|
||||
caller_api_key=api_key,
|
||||
caller_api_base=api_base,
|
||||
key_env_vars=("NIMBLE_API_KEY",),
|
||||
base_env_var="NIMBLE_API_BASE",
|
||||
default_api_base=self.NIMBLE_API_BASE,
|
||||
)
|
||||
if not resolved_api_key:
|
||||
raise ValueError("NIMBLE_API_KEY is not set. Set `NIMBLE_API_KEY` environment variable.")
|
||||
return { # mutable-ok: httpx requires a plain dict of headers
|
||||
**headers,
|
||||
"Authorization": f"Bearer {resolved_api_key}",
|
||||
"Content-Type": "application/json",
|
||||
# Nimble's client-attribution header: names the calling software, nothing else.
|
||||
"X-Client-Source": "litellm",
|
||||
}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
optional_params: dict[str, object], # mutable-ok: BaseSearchConfig.get_complete_url signature
|
||||
data: dict[str, object] | list[dict[str, object]] | None = None, # mutable-ok: base signature
|
||||
**kwargs: object, # kwargs-ok: BaseSearchConfig.get_complete_url signature
|
||||
) -> str:
|
||||
resolved_base: Final = (api_base or get_secret_str("NIMBLE_API_BASE") or self.NIMBLE_API_BASE).rstrip("/")
|
||||
if resolved_base.endswith("/search"):
|
||||
return resolved_base
|
||||
return f"{resolved_base}/search"
|
||||
|
||||
def transform_search_request(
|
||||
self,
|
||||
query: str | list[str], # mutable-ok: BaseSearchConfig.transform_search_request signature
|
||||
optional_params: dict[str, object], # mutable-ok: base signature
|
||||
**kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_request signature
|
||||
) -> dict[str, object]: # mutable-ok: the http handler passes this straight to httpx as the JSON body
|
||||
"""
|
||||
Transform Search request to Nimble API format.
|
||||
|
||||
Nimble already uses the Perplexity unified spec's names, so this is close to a pass-through:
|
||||
- query -> query (a list is joined with spaces; Nimble takes a single string)
|
||||
- max_results -> max_results (sent unclamped so Nimble's own 1-100 validation reports the error)
|
||||
- country -> country, upper-cased to the ISO form Nimble documents
|
||||
- search_domain_filter -> include_domains, with `-`-prefixed entries going to exclude_domains
|
||||
- max_tokens_per_page -> dropped (no Nimble equivalent)
|
||||
|
||||
Everything else is forwarded as-is, so the rest of Nimble's surface stays reachable
|
||||
without LiteLLM tracking it.
|
||||
"""
|
||||
unified_params: Final = self.get_supported_perplexity_optional_params()
|
||||
country: Final = optional_params.get("country")
|
||||
|
||||
# Spread after the derived domain filters so an explicitly supplied `include_domains`
|
||||
# or `exclude_domains` wins over anything read out of `search_domain_filter`.
|
||||
passthrough: Final = MappingProxyType(
|
||||
{param: value for param, value in optional_params.items() if param not in unified_params}
|
||||
)
|
||||
|
||||
return { # mutable-ok: httpx requires a plain dict for the JSON body
|
||||
**_domain_filters(optional_params.get("search_domain_filter")),
|
||||
**passthrough,
|
||||
"query": " ".join(query) if isinstance(query, list) else query,
|
||||
**_optional("max_results", optional_params.get("max_results")),
|
||||
**_optional("country", country.upper() if isinstance(country, str) else None),
|
||||
}
|
||||
|
||||
def transform_search_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
**kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_response signature
|
||||
) -> SearchResponse:
|
||||
"""
|
||||
Transform Nimble API response to LiteLLM unified SearchResponse format.
|
||||
|
||||
`date` carries only the absolute `publish_date`. News results often carry a relative
|
||||
`publish_date_raw` ("1 day ago") instead, which is not a date, so the whole
|
||||
`additional_data` object rides through as an extra on `SearchResult` and nothing is lost.
|
||||
|
||||
Nimble ranks results itself via metadata.position, so the order is preserved as received.
|
||||
A body that does not match the documented schema raises an attributed error rather than
|
||||
being reported as a successful empty search. Parsing the response bytes rather than
|
||||
`.json()` covers the non-JSON case through that same path.
|
||||
"""
|
||||
try:
|
||||
parsed: Final = _NimbleSearchResponse.model_validate_json(raw_response.content)
|
||||
except ValidationError as e:
|
||||
raise self.get_error_class(
|
||||
error_message=f"response does not match the documented /v2/search schema: {e}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=dict(raw_response.headers), # mutable-ok: BaseSearchConfig.get_error_class signature
|
||||
)
|
||||
|
||||
return SearchResponse(
|
||||
results=[ # mutable-ok: SearchResponse.results is declared list[SearchResult]
|
||||
SearchResult(
|
||||
title=result.title or "",
|
||||
url=result.url or "",
|
||||
snippet=result.content or result.description or "",
|
||||
date=_publish_date(result.additional_data),
|
||||
last_updated=None,
|
||||
**_optional("additional_data", result.additional_data),
|
||||
)
|
||||
for result in parsed.results
|
||||
],
|
||||
object="search",
|
||||
)
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: dict[str, str], # mutable-ok: BaseSearchConfig.get_error_class signature
|
||||
) -> Exception:
|
||||
detail: Final = _unwrap_error_detail(error_message).rstrip(". ")
|
||||
return BaseLLMException(
|
||||
status_code=status_code,
|
||||
message=f"Nimble Search: {detail}. See {_NIMBLE_DOCS_URL} for details.",
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
def _unwrap_error_detail(error_message: str) -> str:
|
||||
"""
|
||||
Surface the human-readable message inside Nimble's error envelopes.
|
||||
|
||||
Falls back to the raw body for anything else (CDN HTML pages, plain text, other shapes).
|
||||
"""
|
||||
try:
|
||||
body: Final = _ErrorEnvelope.model_validate_json(error_message)
|
||||
except ValidationError:
|
||||
return error_message
|
||||
return body.detail or body.message or error_message
|
||||
|
||||
|
||||
def _domain_filters(search_domain_filter: object) -> Mapping[str, object]:
|
||||
"""
|
||||
Split the unified `search_domain_filter` into Nimble's include/exclude lists.
|
||||
|
||||
Follows the Perplexity unified spec, where a `-` prefix means "exclude this domain".
|
||||
Anything that is not a list of strings is ignored rather than raising, since it only
|
||||
ever narrows a search that is otherwise valid.
|
||||
"""
|
||||
try:
|
||||
domains: Final = _DomainListAdapter.validate_python(search_domain_filter)
|
||||
except ValidationError:
|
||||
return _NOTHING
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in (
|
||||
("include_domains", tuple(d for d in domains if d and not d.startswith("-"))),
|
||||
("exclude_domains", tuple(d[1:] for d in domains if d.startswith("-") and len(d) > 1)),
|
||||
)
|
||||
if value
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _publish_date(additional_data: object) -> str | None:
|
||||
try:
|
||||
return _AdditionalData.model_validate(additional_data).publish_date
|
||||
except ValidationError:
|
||||
return None
|
||||
|
|
@ -16295,6 +16295,14 @@
|
|||
"notes": "TinyFish Search API"
|
||||
}
|
||||
},
|
||||
"nimble/search": {
|
||||
"input_cost_per_query": 0.005,
|
||||
"litellm_provider": "nimble",
|
||||
"mode": "search",
|
||||
"metadata": {
|
||||
"notes": "Nimble Search API pay-as-you-go list price: $5 per 1,000 searches, up to 100 results per search. Volume plans price differently."
|
||||
}
|
||||
},
|
||||
"elevenlabs/scribe_v1": {
|
||||
"input_cost_per_second": 6.11e-05,
|
||||
"litellm_provider": "elevenlabs",
|
||||
|
|
|
|||
|
|
@ -880,7 +880,9 @@ _HOP_BY_HOP_HEADERS: Final = frozenset(
|
|||
}
|
||||
)
|
||||
|
||||
_SYNTHETIC_REQUEST_EXCLUDED_HEADERS: Final = _HOP_BY_HOP_HEADERS | frozenset({"content-type", "x-forwarded-for"})
|
||||
_SYNTHETIC_REQUEST_EXCLUDED_HEADERS: Final = _HOP_BY_HOP_HEADERS | frozenset(
|
||||
{"content-type", "host", "x-forwarded-for"}
|
||||
)
|
||||
|
||||
_SYNTHETIC_REQUEST_SERVER: Final = ("127.0.0.1", 4000)
|
||||
|
||||
|
|
@ -908,10 +910,57 @@ def _mcp_client_side_auth_header_name() -> str:
|
|||
return MCPRequestHandler.LITELLM_MCP_AUTH_HEADER_NAME
|
||||
|
||||
|
||||
def _identity_header_names() -> frozenset[str]:
|
||||
"""Lowercased header names the deployment reads the caller's identity out of. A name here
|
||||
is a claim about who the caller is rather than a secret, and ``get_user_from_headers``
|
||||
resolves it off the request this module reconstructs, so dropping one would lose end user
|
||||
attribution on the MCP paths that leave ``end_user_id`` unset at connect time.
|
||||
|
||||
``user_header_mappings`` is accepted as a bare mapping as well as a list of them, matching
|
||||
``get_internal_user_header_from_mapping`` and ``get_customer_user_header_from_mapping``.
|
||||
Iterating the bare form without normalizing yields its keys, which would silently exempt
|
||||
nothing."""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
except ImportError:
|
||||
return frozenset()
|
||||
if not general_settings:
|
||||
return frozenset()
|
||||
user_header: Final = general_settings.get("user_header_name")
|
||||
configured: Final = general_settings.get("user_header_mappings")
|
||||
mappings: Final = configured if isinstance(configured, list) else (configured,) if configured else ()
|
||||
mapped: Final = (mapping.get("header_name") for mapping in mappings if isinstance(mapping, Mapping))
|
||||
return frozenset(name.lower() for name in (user_header, *mapped) if isinstance(name, str) and name)
|
||||
|
||||
|
||||
def _forwarded_upstream_header_names() -> frozenset[str]:
|
||||
"""Lowercased header names that a configured MCP server forwards upstream through its
|
||||
``extra_headers`` allowlist. The names are chosen by the admin, so no prefix rule can
|
||||
recognize them, and a caller supplied value under one of them is an upstream credential.
|
||||
|
||||
``authorization`` is left out because ``clean_headers`` already strips it, and claiming it
|
||||
here would change which header ``authenticated_with_header`` resolves to on the oauth
|
||||
passthrough config, which lists it in ``extra_headers`` by design. Identity headers are
|
||||
left out for the same reason: naming one in ``extra_headers`` forwards the caller's
|
||||
identity upstream, it does not turn that identity into a secret."""
|
||||
try:
|
||||
from .mcp_server_manager import global_mcp_server_manager
|
||||
except ImportError:
|
||||
return frozenset()
|
||||
exempt: Final = _identity_header_names() | frozenset({"authorization"})
|
||||
return frozenset(
|
||||
name.lower()
|
||||
for server in global_mcp_server_manager.get_registry().values()
|
||||
for name in (server.extra_headers or ())
|
||||
if name.lower() not in exempt
|
||||
)
|
||||
|
||||
|
||||
def _upstream_credential_headers(header_names: Iterable[str]) -> frozenset[str]:
|
||||
"""Lowercased names of the headers in ``header_names`` that carry an upstream MCP
|
||||
credential rather than request context: the configured client side auth header and
|
||||
the per-server ``x-mcp-{alias}-{header}`` family. ``clean_headers`` only knows the
|
||||
credential rather than request context: the configured client side auth header, any
|
||||
header name a configured server forwards upstream via ``extra_headers``, and the
|
||||
per-server ``x-mcp-{alias}-{header}`` family. ``clean_headers`` only knows the
|
||||
credential headers of the chat completions path, so these are dropped on top of it.
|
||||
"""
|
||||
from .auth.user_api_key_auth_mcp import MCPRequestHandler
|
||||
|
|
@ -923,10 +972,13 @@ def _upstream_credential_headers(header_names: Iterable[str]) -> frozenset[str]:
|
|||
}
|
||||
)
|
||||
client_side_auth: Final = _mcp_client_side_auth_header_name().lower()
|
||||
forwarded_upstream: Final = _forwarded_upstream_header_names()
|
||||
return frozenset(
|
||||
name
|
||||
for name in (raw_name.lower() for raw_name in header_names)
|
||||
if name == client_side_auth or (name.startswith(_MCP_SERVER_AUTH_HEADER_PREFIX) and name not in non_credential)
|
||||
if name == client_side_auth
|
||||
or name in forwarded_upstream
|
||||
or (name.startswith(_MCP_SERVER_AUTH_HEADER_PREFIX) and name not in non_credential)
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -944,7 +996,9 @@ def build_synthetic_mcp_request(
|
|||
``proxy_server_request``, header-based tags, guardrails and trace correlation
|
||||
exactly as on the chat completions path. Hop-by-hop headers describe the
|
||||
original HTTP framing rather than the logical request, so they are dropped, and
|
||||
``x-forwarded-for`` comes from the resolved ``client_ip`` to avoid spoofing. Upstream
|
||||
``x-forwarded-for`` comes from the resolved ``client_ip`` to avoid spoofing. ``host`` is
|
||||
dropped for the same reason: it is what ``Request.url`` is built from, so forwarding it
|
||||
would let a caller choose the URL every logging callback records. Upstream
|
||||
MCP credentials and the deployment's proxy key header, including a custom
|
||||
``litellm_key_header_name``, are dropped so they cannot reach a callback or a guardrail
|
||||
through the derived metadata even when a caller omits ``general_settings``.
|
||||
|
|
@ -991,7 +1045,8 @@ def logging_safe_mcp_headers(raw_headers: Mapping[str, str] | None) -> Mapping[s
|
|||
too: these headers are read back out of the metadata to change proxy behaviour, so
|
||||
leaving one in place would let any MCP client turn off the redaction an admin
|
||||
configured. This path carries no key or team object to authorize an opt-out with, so
|
||||
it always strips them."""
|
||||
it always strips them. ``host`` goes too, so that a caller cannot name the deployment in
|
||||
the guardrail payload and the spend row the way it could once name the request URL."""
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.proxy.litellm_pre_call_utils import (
|
||||
|
|
@ -1003,6 +1058,7 @@ def logging_safe_mcp_headers(raw_headers: Mapping[str, str] | None) -> Mapping[s
|
|||
excluded: Final = (
|
||||
_upstream_credential_headers(raw_headers.keys() if raw_headers else ())
|
||||
| UNTRUSTED_REQUEST_HEADER_CONTROL_FIELDS
|
||||
| frozenset({"host"})
|
||||
)
|
||||
cleaned: Final = clean_headers(
|
||||
Headers(raw_headers),
|
||||
|
|
|
|||
|
|
@ -861,6 +861,8 @@ class DBSpendUpdateWriter:
|
|||
):
|
||||
verbose_proxy_logger.debug("acquired lock for spend updates")
|
||||
|
||||
uncommitted: dict[str, Any] = {} # mutable-ok: tracks popped categories still needing commit
|
||||
|
||||
try:
|
||||
(
|
||||
db_spend_update_transactions,
|
||||
|
|
@ -871,6 +873,15 @@ class DBSpendUpdateWriter:
|
|||
daily_agent_spend_update_transactions,
|
||||
) = await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline()
|
||||
|
||||
uncommitted = { # mutable-ok: drives which popped categories still need re-queuing
|
||||
"db_spend_update_transactions": db_spend_update_transactions,
|
||||
"daily_spend_update_transactions": daily_spend_update_transactions,
|
||||
"daily_team_spend_update_transactions": daily_team_spend_update_transactions,
|
||||
"daily_org_spend_update_transactions": daily_org_spend_update_transactions,
|
||||
"daily_end_user_spend_update_transactions": daily_end_user_spend_update_transactions,
|
||||
"daily_agent_spend_update_transactions": daily_agent_spend_update_transactions,
|
||||
}
|
||||
|
||||
if db_spend_update_transactions is not None:
|
||||
verbose_proxy_logger.info(
|
||||
"Spend tracking - committing spend updates from Redis to DB: "
|
||||
|
|
@ -890,6 +901,7 @@ class DBSpendUpdateWriter:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
db_spend_update_transactions=db_spend_update_transactions,
|
||||
)
|
||||
uncommitted.pop("db_spend_update_transactions", None)
|
||||
|
||||
if daily_spend_update_transactions is not None:
|
||||
await DBSpendUpdateWriter.update_daily_user_spend(
|
||||
|
|
@ -898,6 +910,8 @@ class DBSpendUpdateWriter:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
daily_spend_transactions=daily_spend_update_transactions,
|
||||
)
|
||||
uncommitted.pop("daily_spend_update_transactions", None)
|
||||
|
||||
if daily_team_spend_update_transactions is not None:
|
||||
await DBSpendUpdateWriter.update_daily_team_spend(
|
||||
n_retry_times=n_retry_times,
|
||||
|
|
@ -905,6 +919,7 @@ class DBSpendUpdateWriter:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
daily_spend_transactions=daily_team_spend_update_transactions,
|
||||
)
|
||||
uncommitted.pop("daily_team_spend_update_transactions", None)
|
||||
|
||||
if daily_org_spend_update_transactions is not None:
|
||||
await DBSpendUpdateWriter.update_daily_org_spend(
|
||||
|
|
@ -913,6 +928,7 @@ class DBSpendUpdateWriter:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
daily_spend_transactions=daily_org_spend_update_transactions,
|
||||
)
|
||||
uncommitted.pop("daily_org_spend_update_transactions", None)
|
||||
|
||||
if daily_end_user_spend_update_transactions is not None:
|
||||
await DBSpendUpdateWriter.update_daily_end_user_spend(
|
||||
|
|
@ -921,6 +937,8 @@ class DBSpendUpdateWriter:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
daily_spend_transactions=daily_end_user_spend_update_transactions,
|
||||
)
|
||||
uncommitted.pop("daily_end_user_spend_update_transactions", None)
|
||||
|
||||
if daily_agent_spend_update_transactions is not None:
|
||||
await DBSpendUpdateWriter.update_daily_agent_spend(
|
||||
n_retry_times=n_retry_times,
|
||||
|
|
@ -928,14 +946,20 @@ class DBSpendUpdateWriter:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
daily_spend_transactions=daily_agent_spend_update_transactions,
|
||||
)
|
||||
uncommitted.pop("daily_agent_spend_update_transactions", None)
|
||||
except Exception as e:
|
||||
spend_log_error(
|
||||
"Spend tracking - failed to commit spend updates from Redis to DB. "
|
||||
"Data already popped from Redis may be lost. Error: %s",
|
||||
"Re-queuing uncommitted transactions to Redis for retry on next tick. Error: %s",
|
||||
str(e),
|
||||
exc=e,
|
||||
)
|
||||
finally:
|
||||
to_restore = { # mutable-ok: transient kwargs payload consumed immediately below
|
||||
name: txns for name, txns in uncommitted.items() if txns is not None
|
||||
}
|
||||
if to_restore:
|
||||
await self.redis_update_buffer.restore_transactions_to_redis(**to_restore)
|
||||
await self.pod_lock_manager.release_lock(
|
||||
cronjob_id=DB_SPEND_UPDATE_JOB_NAME,
|
||||
)
|
||||
|
|
@ -1085,21 +1109,15 @@ class DBSpendUpdateWriter:
|
|||
):
|
||||
verbose_proxy_logger.debug("acquired lock for daily tag spend updates")
|
||||
try:
|
||||
daily_tag_spend_update_transactions: Final = (
|
||||
await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer()
|
||||
await self._drain_and_commit_daily_tag_spend_from_redis(
|
||||
prisma_client=prisma_client,
|
||||
n_retry_times=n_retry_times,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
if daily_tag_spend_update_transactions:
|
||||
await DBSpendUpdateWriter.update_daily_tag_spend(
|
||||
n_retry_times=n_retry_times,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
daily_spend_transactions=daily_tag_spend_update_transactions,
|
||||
)
|
||||
except Exception as e:
|
||||
spend_log_error(
|
||||
"Spend tracking - failed to commit daily tag spend updates from Redis to DB. "
|
||||
"Data already popped from Redis may be lost. Error: %s",
|
||||
"Re-queuing to Redis for retry on next tick. Error: %s",
|
||||
str(e),
|
||||
exc=e,
|
||||
)
|
||||
|
|
@ -1108,6 +1126,37 @@ class DBSpendUpdateWriter:
|
|||
cronjob_id=DB_DAILY_TAG_SPEND_UPDATE_JOB_NAME,
|
||||
)
|
||||
|
||||
async def _drain_and_commit_daily_tag_spend_from_redis(
|
||||
self,
|
||||
prisma_client: PrismaClient,
|
||||
n_retry_times: int,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> None:
|
||||
"""
|
||||
Drain the Redis tag spend buffer and commit it, restoring the drained transactions if the commit fails.
|
||||
|
||||
The drain is destructive, so a failed commit must push the transactions back for the next tick
|
||||
or their spend is lost permanently.
|
||||
"""
|
||||
daily_tag_spend_update_transactions: Final = (
|
||||
await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer()
|
||||
)
|
||||
if not daily_tag_spend_update_transactions:
|
||||
return
|
||||
|
||||
try:
|
||||
await DBSpendUpdateWriter.update_daily_tag_spend(
|
||||
n_retry_times=n_retry_times,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
daily_spend_transactions=daily_tag_spend_update_transactions,
|
||||
)
|
||||
except Exception:
|
||||
await self.redis_update_buffer.restore_transactions_to_redis(
|
||||
daily_tag_spend_update_transactions=daily_tag_spend_update_transactions,
|
||||
)
|
||||
raise
|
||||
|
||||
async def _flush_tool_discovery_queue(
|
||||
self,
|
||||
prisma_client: PrismaClient,
|
||||
|
|
@ -1607,9 +1656,6 @@ class DBSpendUpdateWriter:
|
|||
)
|
||||
|
||||
except Exception as e:
|
||||
if "transactions_to_process" in locals():
|
||||
for key in transactions_to_process:
|
||||
daily_spend_transactions.pop(key, None)
|
||||
_raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -6,8 +6,11 @@ This is to prevent deadlocks and improve reliability
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
from redis.exceptions import RedisError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching import RedisCache
|
||||
from litellm.constants import (
|
||||
|
|
@ -372,6 +375,59 @@ class RedisUpdateBuffer:
|
|||
if daily_txns:
|
||||
await daily_queue.update_queue.put(daily_txns)
|
||||
|
||||
async def restore_transactions_to_redis(
|
||||
self,
|
||||
db_spend_update_transactions: DBSpendUpdateTransactions | None = None,
|
||||
daily_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None,
|
||||
daily_team_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None,
|
||||
daily_org_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None,
|
||||
daily_end_user_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None,
|
||||
daily_agent_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None,
|
||||
daily_tag_spend_update_transactions: Mapping[str, BaseDailySpendTransaction] | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Re-push transactions that were popped from Redis but not committed to the DB.
|
||||
|
||||
The leader drains the buffers with a destructive ``lpop`` before committing to
|
||||
the database. When a commit fails after its retries are exhausted, the popped
|
||||
transactions must be pushed back so a later scheduler tick can retry them;
|
||||
otherwise the aggregated spend is lost permanently. The re-pushed payloads use
|
||||
the same JSON encoding as the store path, so the next drain parses them normally.
|
||||
"""
|
||||
if self.redis_cache is None:
|
||||
return
|
||||
|
||||
restore_configs: Final = (
|
||||
(db_spend_update_transactions, REDIS_UPDATE_BUFFER_KEY),
|
||||
(daily_spend_update_transactions, REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY),
|
||||
(daily_team_spend_update_transactions, REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY),
|
||||
(daily_org_spend_update_transactions, REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY),
|
||||
(daily_end_user_spend_update_transactions, REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY),
|
||||
(daily_agent_spend_update_transactions, REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY),
|
||||
(daily_tag_spend_update_transactions, REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY),
|
||||
)
|
||||
|
||||
rpush_list: Final = tuple(
|
||||
RedisPipelineRpushOperation(key=redis_key, values=(safe_dumps(transactions),))
|
||||
for transactions, redis_key in restore_configs
|
||||
if transactions
|
||||
)
|
||||
if len(rpush_list) == 0:
|
||||
return
|
||||
|
||||
try:
|
||||
await self.redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
|
||||
verbose_proxy_logger.info(
|
||||
"Spend tracking - restored %d uncommitted transaction set(s) to Redis for retry on next tick.",
|
||||
len(rpush_list),
|
||||
)
|
||||
except RedisError as e:
|
||||
verbose_proxy_logger.error(
|
||||
"Spend tracking - failed to restore uncommitted transactions to Redis. "
|
||||
"These spend updates are lost. Error: %s",
|
||||
str(e),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _number_of_transactions_to_store_in_redis(
|
||||
db_spend_update_transactions: DBSpendUpdateTransactions,
|
||||
|
|
|
|||
|
|
@ -464,35 +464,38 @@ def _is_configured_pre_routing_strategy(llm_router: "Router", router_name: str)
|
|||
)
|
||||
|
||||
|
||||
def _validate_judge_model(llm_router: "Router | None", judge_model: str) -> None:
|
||||
"""Reject a judge model the dispatch path cannot resolve, at start rather than as a
|
||||
silently growing error count once the job is already sampling and billing."""
|
||||
if llm_router is not None and _is_configured_pre_routing_strategy(llm_router, judge_model):
|
||||
def _validate_plain_model(llm_router: "Router | None", model: str, field_name: str) -> None:
|
||||
"""Reject a model the dispatch path cannot resolve, at start rather than as a silently
|
||||
growing error count once the job is already sampling and billing. Both the judge and a
|
||||
reverse job's baseline must be plain models: an auto-router in either slot would
|
||||
re-route per turn, so the comparison would have no fixed arm to attribute results to."""
|
||||
if llm_router is not None and _is_configured_pre_routing_strategy(llm_router, model):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"judge_model '{judge_model}' is an auto-router; the judge must be a plain model",
|
||||
detail=f"{field_name} '{model}' is an auto-router; it must be a plain model",
|
||||
)
|
||||
if router_resolves_model(llm_router, judge_model):
|
||||
if router_resolves_model(llm_router, model):
|
||||
return
|
||||
import litellm
|
||||
|
||||
try:
|
||||
litellm.get_llm_provider(model=judge_model)
|
||||
litellm.get_llm_provider(model=model)
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
f"judge_model '{judge_model}' is neither a model configured on this proxy nor a "
|
||||
f"{field_name} '{model}' is neither a model configured on this proxy nor a "
|
||||
"provider-qualified public model name (e.g. 'anthropic/claude-sonnet-5')"
|
||||
),
|
||||
) from e
|
||||
|
||||
|
||||
def _is_unique_violation(error: Exception) -> bool:
|
||||
"""Whether a Prisma create failed on a unique index. One active job per key lives in
|
||||
a partial unique index (raw SQL in the migration; schema.prisma cannot express partial
|
||||
indexes), so the read-then-create check above it is advisory: two concurrent starts
|
||||
pass the read, and the loser must surface as the same 409 rather than a 500."""
|
||||
"""Whether a Prisma create failed on a unique index. One active job per key and
|
||||
direction lives in a partial unique index (raw SQL in the migration; schema.prisma
|
||||
cannot express partial indexes), so the read-then-create check above it is advisory:
|
||||
two concurrent starts pass the read, and the loser must surface as the same 409
|
||||
rather than a 500."""
|
||||
try:
|
||||
from prisma.errors import UniqueViolationError
|
||||
except ImportError:
|
||||
|
|
@ -573,8 +576,10 @@ def _slices(rows: Sequence[_AttemptAggRow]) -> tuple[ShadowEvalSlice, ...]:
|
|||
|
||||
async def _shadow_eval_results(prisma_client: "PrismaClient", job_id: str) -> ShadowEvalResult | None:
|
||||
"""Both stratifications of one job's verdicts. Tier answers "where does the router do
|
||||
well"; current-model answers "which of the models this key uses today would the router
|
||||
beat". Reads are bounded by the job's own attempts (<= max_turns) via the job_id index."""
|
||||
well"; the model stratification groups by whichever model served the real arm, so it
|
||||
answers "which of the models this key uses today would the router beat" forward, and
|
||||
"for the turns the router sent to X, did X beat the baseline" in reverse. Reads are
|
||||
bounded by the job's own attempts (<= max_turns) via the job_id index."""
|
||||
by_tier: Final = _ATTEMPT_AGG_ROWS.validate_python(
|
||||
await prisma_client.db.query_raw(_ATTEMPT_AGG_BY_TIER_SQL, job_id) or ()
|
||||
)
|
||||
|
|
@ -604,9 +609,15 @@ async def start_shadow_eval(
|
|||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
) -> ShadowEvalJobResponse:
|
||||
"""
|
||||
Start a pre-adoption shadow eval: duplicate a sampled slice of a key's live traffic
|
||||
through an auto-router, judge real vs. shadow responses blind, and stratify win rates
|
||||
by the router's tier classification and by the incumbent model.
|
||||
Start a shadow eval: duplicate a sampled slice of a key's live traffic against a second
|
||||
arm, judge the two responses blind, and stratify win rates by tier and by the model that
|
||||
served the real arm.
|
||||
|
||||
A forward job answers whether the key should adopt router_name: it samples the requests
|
||||
the router did not serve and duplicates them through it. A reverse job answers whether a
|
||||
key already on the router still gains from it: it samples the requests the router did
|
||||
serve and duplicates them against baseline_model. A key can hold one active job per
|
||||
direction, so both questions can run at once.
|
||||
|
||||
Shadow responses are never served to users. The job samples until it has judged
|
||||
max_turns turns, reaches the end of its window, or is stopped; sampling changes
|
||||
|
|
@ -620,7 +631,9 @@ async def start_shadow_eval(
|
|||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
if llm_router is None or not _is_configured_pre_routing_strategy(llm_router, data.router_name):
|
||||
raise HTTPException(status_code=400, detail=f"'{data.router_name}' is not a configured auto-router")
|
||||
_validate_judge_model(llm_router, data.judge_model)
|
||||
_validate_plain_model(llm_router, data.judge_model, "judge_model")
|
||||
if data.baseline_model is not None:
|
||||
_validate_plain_model(llm_router, data.baseline_model, "baseline_model")
|
||||
key_row: Final = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": data.api_key_id} # mutable-ok: Prisma filter
|
||||
)
|
||||
|
|
@ -634,16 +647,20 @@ async def start_shadow_eval(
|
|||
)
|
||||
|
||||
# A job that expired or exhausted its turn budget stopped sampling on its own, but
|
||||
# still holds the one-active-per-key partial unique index until stamped; free it so
|
||||
# a new eval can start.
|
||||
# still holds its slot in the per-key, per-direction partial unique index until
|
||||
# stamped; free it so a new eval can start. Sweeping both directions is deliberate.
|
||||
await prisma_client.db.execute_raw(_SWEEP_FINISHED_JOBS_SQL, data.api_key_id)
|
||||
active: Final = await prisma_client.db.litellm_shadowevaljob.find_first(
|
||||
where={"api_key_id": data.api_key_id, "stopped_at": None}, # mutable-ok: Prisma filter
|
||||
where={ # mutable-ok: Prisma filter
|
||||
"api_key_id": data.api_key_id,
|
||||
"direction": data.direction,
|
||||
"stopped_at": None,
|
||||
},
|
||||
)
|
||||
if active is not None:
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=f"Key already has an active shadow eval job ({active.id}). Stop it first.",
|
||||
detail=f"Key already has an active {data.direction} shadow eval job ({active.id}). Stop it first.",
|
||||
)
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
try:
|
||||
|
|
@ -651,6 +668,8 @@ async def start_shadow_eval(
|
|||
data={ # mutable-ok: Prisma payload
|
||||
"api_key_id": data.api_key_id,
|
||||
"router_name": data.router_name,
|
||||
"direction": data.direction,
|
||||
"baseline_model": data.baseline_model,
|
||||
"judge_model": data.judge_model,
|
||||
"shadow_percentage": data.shadow_percentage,
|
||||
"max_turns": data.max_turns,
|
||||
|
|
@ -663,7 +682,9 @@ async def start_shadow_eval(
|
|||
raise
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail="Key already has an active shadow eval job (started concurrently). Stop it first.",
|
||||
detail=(
|
||||
f"Key already has an active {data.direction} shadow eval job (started concurrently). Stop it first."
|
||||
),
|
||||
) from e
|
||||
return ShadowEvalJobResponse.model_validate(job, from_attributes=True)
|
||||
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
|||
)
|
||||
from litellm.proxy.utils import is_known_model
|
||||
from litellm.proxy.vector_store_endpoints.utils import (
|
||||
assert_proxy_admin_for_vector_store_index_management,
|
||||
assert_user_can_access_vector_store,
|
||||
get_litellm_managed_vector_store,
|
||||
is_allowed_to_call_vector_store_endpoint,
|
||||
|
|
@ -1234,6 +1235,37 @@ async def assemblyai_proxy_route(
|
|||
return received_value
|
||||
|
||||
|
||||
def get_azure_ai_search_index_from_endpoint(endpoint: str) -> str | None:
|
||||
"""Return the index name in the ``/indexes/{name}`` position of an Azure AI
|
||||
Search passthrough path, or ``None`` when the path targets no index.
|
||||
|
||||
Only the segment immediately after ``indexes`` is the operable target. Any
|
||||
other segment (for example the trailing ``index`` in ``.../docs/index``) must
|
||||
never be treated as the index, otherwise a caller authorized on one index
|
||||
could have Azure apply the operation to a different index on the same service.
|
||||
"""
|
||||
segments: Final = endpoint.split("?", 1)[0].strip("/").split("/")
|
||||
for position, segment in enumerate(segments):
|
||||
if segment == "indexes" and position + 1 < len(segments):
|
||||
return segments[position + 1] or None
|
||||
return None
|
||||
|
||||
|
||||
def is_azure_ai_search_service_level_index_create(method: str, endpoint: str) -> bool:
|
||||
"""Return True for ``POST /indexes``, Azure AI Search's service-level index create.
|
||||
|
||||
No index name appears in that path, so ``get_azure_ai_search_index_from_endpoint``
|
||||
yields None and the managed-index branch can never claim the request. Without an
|
||||
explicit guard it reaches the generic Azure passthrough on the proxy's own
|
||||
credential, so a non-admin could create an index whenever ``AZURE_API_BASE``
|
||||
points at the Search service.
|
||||
"""
|
||||
if method != "POST":
|
||||
return False
|
||||
path: Final = endpoint.split("?", 1)[0].strip("/")
|
||||
return path == "indexes" or path.endswith("/indexes")
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/azure_ai/{endpoint:path}",
|
||||
methods=["GET", "POST", "PUT", "DELETE", "PATCH"],
|
||||
|
|
@ -1259,10 +1291,15 @@ async def azure_proxy_route(
|
|||
"""
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if is_azure_ai_search_service_level_index_create(method=request.method, endpoint=endpoint):
|
||||
assert_proxy_admin_for_vector_store_index_management(user_api_key_dict, operation="create")
|
||||
|
||||
parts: Final = endpoint.split(
|
||||
"/"
|
||||
) # azure model is in the url - e.g. https://{endpoint}/openai/deployments/{deployment-id}/completions?api-version=2024-10-21
|
||||
|
||||
search_index_name: Final = get_azure_ai_search_index_from_endpoint(endpoint)
|
||||
|
||||
if len(parts) > 1 and llm_router:
|
||||
for part in parts:
|
||||
# check if LLM MODEL
|
||||
|
|
@ -1271,9 +1308,9 @@ async def azure_proxy_route(
|
|||
)
|
||||
# check if vector store index
|
||||
is_vector_store_index = (
|
||||
(litellm.vector_store_index_registry.is_vector_store_index(vector_store_index_name=part))
|
||||
if litellm.vector_store_index_registry is not None
|
||||
else False
|
||||
part == search_index_name
|
||||
and litellm.vector_store_index_registry is not None
|
||||
and litellm.vector_store_index_registry.is_vector_store_index(vector_store_index_name=part)
|
||||
)
|
||||
|
||||
if is_router_model:
|
||||
|
|
|
|||
|
|
@ -1450,15 +1450,20 @@ model LiteLLM_AutoRouterSession {
|
|||
@@index([last_turn_at], map: "idx_autorouter_session_last_turn")
|
||||
}
|
||||
|
||||
// Shadow eval: pre-adoption evaluation of an auto-router against a key's live traffic.
|
||||
// A sampled slice of requests is duplicated through the router in a detached task and an
|
||||
// LLM judge compares real vs shadow responses blind. The job row is immutable config plus
|
||||
// Shadow eval: evaluation of an auto-router against a key's live traffic, in either
|
||||
// direction. forward duplicates the requests the key did not route through the router
|
||||
// through it, answering whether the key should adopt it; reverse duplicates the requests
|
||||
// the router did serve against a fixed baseline model, answering whether a key already on
|
||||
// it still benefits. Either way a sampled slice runs in a detached task and an LLM judge
|
||||
// compares real vs shadow responses blind. The job row is immutable config plus
|
||||
// stopped_at; every count, status, and spend figure is derived from the append-only
|
||||
// attempt rows, so nothing can disagree across pods or stop races.
|
||||
model LiteLLM_ShadowEvalJob {
|
||||
id String @id @default(cuid())
|
||||
api_key_id String // hashed virtual key whose traffic is shadowed
|
||||
router_name String
|
||||
router_name String // the auto-router under evaluation, in either direction
|
||||
direction String @default("forward") // forward | reverse
|
||||
baseline_model String? // reverse only: the fixed model the router is judged against
|
||||
judge_model String
|
||||
shadow_percentage Float
|
||||
max_turns Int // sample budget: judge at most this many turns
|
||||
|
|
|
|||
|
|
@ -86,7 +86,7 @@ def _is_vector_store_index_lifecycle_request(
|
|||
return True
|
||||
|
||||
# POST /indexes (create index at service level; no index name in path).
|
||||
normalized: Final = request_path.rstrip("/")
|
||||
normalized: Final = request_path.split("?", 1)[0].rstrip("/")
|
||||
if request_method == "POST" and normalized.endswith("/indexes"):
|
||||
return True
|
||||
|
||||
|
|
@ -387,17 +387,19 @@ def is_allowed_to_call_vector_store_endpoint(
|
|||
)
|
||||
return True
|
||||
|
||||
# Determine the permission type based on the request
|
||||
# Writes are classified before reads so a path matching both patterns
|
||||
# requires the stronger grant (e.g. the azure batch write on an index
|
||||
# named "analyze*" also contains the "/analyze" read fragment)
|
||||
permission_type = None
|
||||
for endpoint in provider_vector_store_endpoints["read"]:
|
||||
for endpoint in provider_vector_store_endpoints["write"]:
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "read"
|
||||
permission_type = "write"
|
||||
break
|
||||
|
||||
if permission_type is None:
|
||||
for endpoint in provider_vector_store_endpoints["write"]:
|
||||
for endpoint in provider_vector_store_endpoints["read"]:
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "write"
|
||||
permission_type = "read"
|
||||
break
|
||||
|
||||
if permission_type is None:
|
||||
|
|
@ -454,15 +456,15 @@ def is_allowed_to_call_vector_store_files_endpoint(
|
|||
request_route: Final = get_request_route(request)
|
||||
|
||||
permission_type: str | None = None
|
||||
for endpoint in provider_vector_store_endpoints.get("read", ()):
|
||||
for endpoint in provider_vector_store_endpoints.get("write", ()):
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "read"
|
||||
permission_type = "write"
|
||||
break
|
||||
|
||||
if permission_type is None:
|
||||
for endpoint in provider_vector_store_endpoints.get("write", ()):
|
||||
for endpoint in provider_vector_store_endpoints.get("read", ()):
|
||||
if request.method == endpoint[0] and _does_endpoint_match(endpoint[1], request_route):
|
||||
permission_type = "write"
|
||||
permission_type = "read"
|
||||
break
|
||||
|
||||
if permission_type is None:
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from collections.abc import Sequence
|
||||
from enum import Enum
|
||||
from typing import Any, Final, Literal, Optional, Union
|
||||
|
||||
|
|
@ -59,7 +60,7 @@ class RedisPipelineRpushOperation(TypedDict):
|
|||
"""
|
||||
|
||||
key: str
|
||||
values: list[Any]
|
||||
values: Sequence[Any]
|
||||
|
||||
|
||||
class RedisPipelineLpopOperation(TypedDict):
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from collections.abc import Mapping
|
|||
from datetime import datetime, timezone
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
from pydantic import AliasChoices, BaseModel, ConfigDict, Field, computed_field, field_validator
|
||||
from pydantic import AliasChoices, BaseModel, ConfigDict, Field, computed_field, field_validator, model_validator
|
||||
|
||||
from litellm.router_strategy.complexity_router.config import ComplexityRouterConfig
|
||||
from litellm.types.utils import StandardLoggingRoutingDecision
|
||||
|
|
@ -146,11 +146,13 @@ class AutoRouterBenchmarksResponse(BaseModel):
|
|||
|
||||
ShadowEvalStatus: TypeAlias = Literal["running", "completed", "stopped"]
|
||||
|
||||
ShadowEvalDirection: TypeAlias = Literal["forward", "reverse"]
|
||||
|
||||
DEFAULT_SHADOW_EVAL_JUDGE_MODEL: Final[str] = "anthropic/claude-sonnet-5"
|
||||
|
||||
|
||||
class StartShadowEvalRequest(BaseModel):
|
||||
"""Start shadowing a key's traffic through an auto-router for blind comparison."""
|
||||
"""Start duplicating a key's traffic for blind comparison against an auto-router."""
|
||||
|
||||
api_key_id: str = Field(
|
||||
description=(
|
||||
|
|
@ -158,7 +160,23 @@ class StartShadowEvalRequest(BaseModel):
|
|||
"key's traffic; requests made with any other key are not sampled."
|
||||
)
|
||||
)
|
||||
router_name: str = Field(description="The auto-router config to shadow requests through")
|
||||
router_name: str = Field(description="The auto-router under evaluation, in either direction")
|
||||
direction: ShadowEvalDirection = Field(
|
||||
default="forward",
|
||||
description=(
|
||||
"forward answers 'should this key adopt router_name': it samples the requests the key did NOT "
|
||||
"route through the router and duplicates them through it. reverse answers 'is the router still "
|
||||
"worth it for a key already on it': it samples the requests the router did serve and duplicates "
|
||||
"them against baseline_model. The response the caller received is always the real arm"
|
||||
),
|
||||
)
|
||||
baseline_model: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Required when direction is reverse and rejected otherwise: the fixed model the router's own "
|
||||
"responses are judged against. Must be a plain model rather than another auto-router"
|
||||
),
|
||||
)
|
||||
shadow_percentage: float = Field(
|
||||
ge=0.1,
|
||||
le=100.0,
|
||||
|
|
@ -193,15 +211,33 @@ class StartShadowEvalRequest(BaseModel):
|
|||
def _round_percentage(cls, value: float) -> float:
|
||||
return round(value, 2)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _baseline_model_matches_direction(self) -> "StartShadowEvalRequest":
|
||||
if self.direction == "reverse" and self.baseline_model is None:
|
||||
raise ValueError("baseline_model is required when direction is 'reverse'")
|
||||
if self.direction == "forward" and self.baseline_model is not None:
|
||||
raise ValueError("baseline_model is only meaningful when direction is 'reverse'")
|
||||
return self
|
||||
|
||||
|
||||
class ShadowEvalSlice(BaseModel):
|
||||
"""Judge outcomes for one slice of a job's verdicts (a router tier, or one of the
|
||||
models the shadowed key currently uses)."""
|
||||
models that served the real arm)."""
|
||||
|
||||
group: str
|
||||
turn_count: int
|
||||
real_win_rate_pct: float = Field(description="Share of judged turns where the real (control) model won")
|
||||
shadow_win_rate_pct: float = Field(description="Share of judged turns where the shadowed router's pick won")
|
||||
real_win_rate_pct: float = Field(
|
||||
description=(
|
||||
"Share of judged turns the real arm won, meaning the response the caller actually received: "
|
||||
"the key's own model in forward mode, the router's pick in reverse"
|
||||
)
|
||||
)
|
||||
shadow_win_rate_pct: float = Field(
|
||||
description=(
|
||||
"Share of judged turns the shadow arm won, meaning the duplicated response nobody was served: "
|
||||
"the router's pick in forward mode, baseline_model in reverse"
|
||||
)
|
||||
)
|
||||
tie_rate_pct: float
|
||||
avg_judge_confidence: float
|
||||
|
||||
|
|
@ -210,7 +246,12 @@ class ShadowEvalResult(BaseModel):
|
|||
"""Stratified results of a shadow-eval job's verdicts so far."""
|
||||
|
||||
by_tier: tuple[ShadowEvalSlice, ...]
|
||||
by_current_model: tuple[ShadowEvalSlice, ...]
|
||||
by_current_model: tuple[ShadowEvalSlice, ...] = Field(
|
||||
description=(
|
||||
"Sliced by the model that served the real arm: the key's incumbent models in forward mode, "
|
||||
"and in reverse the models the router itself picked"
|
||||
)
|
||||
)
|
||||
overall_shadow_win_rate_pct: float
|
||||
overall_tie_rate_pct: float
|
||||
|
||||
|
|
@ -226,6 +267,8 @@ class ShadowEvalJobResponse(BaseModel):
|
|||
job_id: str = Field(validation_alias=AliasChoices("id", "job_id"))
|
||||
api_key_id: str = Field(description="The hashed virtual key whose traffic this job evaluates, and only that key's")
|
||||
router_name: str
|
||||
direction: ShadowEvalDirection = "forward"
|
||||
baseline_model: str | None = None
|
||||
judge_model: str
|
||||
shadow_percentage: float
|
||||
max_turns: int
|
||||
|
|
|
|||
|
|
@ -3758,6 +3758,7 @@ class SearchProviders(str, Enum):
|
|||
YOU_COM = "you_com"
|
||||
APISERPENT = "apiserpent"
|
||||
TINYFISH = "tinyfish"
|
||||
NIMBLE = "nimble"
|
||||
|
||||
|
||||
# Create a set of all search provider values for quick lookup
|
||||
|
|
|
|||
|
|
@ -9064,6 +9064,7 @@ class ProviderConfigManager:
|
|||
from litellm.llms.firecrawl.search.transformation import FirecrawlSearchConfig
|
||||
from litellm.llms.google_pse.search.transformation import GooglePSESearchConfig
|
||||
from litellm.llms.linkup.search.transformation import LinkupSearchConfig
|
||||
from litellm.llms.nimble.search.transformation import NimbleSearchConfig
|
||||
from litellm.llms.parallel_ai.search.transformation import (
|
||||
ParallelAISearchConfig,
|
||||
)
|
||||
|
|
@ -9093,6 +9094,7 @@ class ProviderConfigManager:
|
|||
SearchProviders.YOU_COM: YouComSearchConfig,
|
||||
SearchProviders.APISERPENT: APISerpentSearchConfig,
|
||||
SearchProviders.TINYFISH: TinyfishSearchConfig,
|
||||
SearchProviders.NIMBLE: NimbleSearchConfig,
|
||||
}
|
||||
config_class: Final = PROVIDER_TO_CONFIG_MAP.get(provider, None)
|
||||
if config_class is None:
|
||||
|
|
|
|||
|
|
@ -16295,6 +16295,14 @@
|
|||
"notes": "TinyFish Search API"
|
||||
}
|
||||
},
|
||||
"nimble/search": {
|
||||
"input_cost_per_query": 0.005,
|
||||
"litellm_provider": "nimble",
|
||||
"mode": "search",
|
||||
"metadata": {
|
||||
"notes": "Nimble Search API pay-as-you-go list price: $5 per 1,000 searches, up to 100 results per search. Volume plans price differently."
|
||||
}
|
||||
},
|
||||
"elevenlabs/scribe_v1": {
|
||||
"input_cost_per_second": 6.11e-05,
|
||||
"litellm_provider": "elevenlabs",
|
||||
|
|
|
|||
|
|
@ -2423,6 +2423,13 @@
|
|||
"search": true
|
||||
}
|
||||
},
|
||||
"nimble": {
|
||||
"display_name": "Nimble (`nimble`)",
|
||||
"url": "https://docs.nimbleway.com/api-reference/search/search",
|
||||
"endpoints": {
|
||||
"search": true
|
||||
}
|
||||
},
|
||||
"triton": {
|
||||
"display_name": "Triton (`triton`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/triton-inference-server",
|
||||
|
|
|
|||
|
|
@ -1450,15 +1450,20 @@ model LiteLLM_AutoRouterSession {
|
|||
@@index([last_turn_at], map: "idx_autorouter_session_last_turn")
|
||||
}
|
||||
|
||||
// Shadow eval: pre-adoption evaluation of an auto-router against a key's live traffic.
|
||||
// A sampled slice of requests is duplicated through the router in a detached task and an
|
||||
// LLM judge compares real vs shadow responses blind. The job row is immutable config plus
|
||||
// Shadow eval: evaluation of an auto-router against a key's live traffic, in either
|
||||
// direction. forward duplicates the requests the key did not route through the router
|
||||
// through it, answering whether the key should adopt it; reverse duplicates the requests
|
||||
// the router did serve against a fixed baseline model, answering whether a key already on
|
||||
// it still benefits. Either way a sampled slice runs in a detached task and an LLM judge
|
||||
// compares real vs shadow responses blind. The job row is immutable config plus
|
||||
// stopped_at; every count, status, and spend figure is derived from the append-only
|
||||
// attempt rows, so nothing can disagree across pods or stop races.
|
||||
model LiteLLM_ShadowEvalJob {
|
||||
id String @id @default(cuid())
|
||||
api_key_id String // hashed virtual key whose traffic is shadowed
|
||||
router_name String
|
||||
router_name String // the auto-router under evaluation, in either direction
|
||||
direction String @default("forward") // forward | reverse
|
||||
baseline_model String? // reverse only: the fixed model the router is judged against
|
||||
judge_model String
|
||||
shadow_percentage Float
|
||||
max_turns Int // sample budget: judge at most this many turns
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ SEARCH_PROVIDERS = [
|
|||
"serper",
|
||||
"apiserpent",
|
||||
"tinyfish",
|
||||
"nimble",
|
||||
]
|
||||
|
||||
ALLOWED_FILES_IN_LLMS_FOLDER = [
|
||||
|
|
|
|||
155
tests/search_tests/test_nimble_search.py
Normal file
155
tests/search_tests/test_nimble_search.py
Normal file
|
|
@ -0,0 +1,155 @@
|
|||
"""
|
||||
Tests for Nimble Search API integration.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import litellm
|
||||
from tests.search_tests.base_search_unit_tests import BaseSearchTest
|
||||
|
||||
MOCK_NIMBLE_RESPONSE = {
|
||||
"request_id": "0f8b3a1c-1d2e-4f5a-9b0c-6d7e8f9a0b1c",
|
||||
"total_results": 2,
|
||||
"results": [
|
||||
{
|
||||
"title": "Nimble Web API",
|
||||
"description": "Short SERP description",
|
||||
"url": "https://nimbleway.com/",
|
||||
"content": "Full markdown content for the first result",
|
||||
"metadata": {"position": 1, "entity_type": "organic", "country": "US", "locale": "en"},
|
||||
"additional_data": {"publish_date": "2026-07-15"},
|
||||
},
|
||||
{
|
||||
"title": "Nimble Docs",
|
||||
"description": "Only a description here",
|
||||
"url": "https://docs.nimbleway.com/",
|
||||
"content": "",
|
||||
"metadata": {"position": 2, "entity_type": "organic"},
|
||||
"additional_data": None,
|
||||
},
|
||||
],
|
||||
"serp_data": None,
|
||||
}
|
||||
|
||||
|
||||
def _mock_response():
|
||||
response = Mock()
|
||||
response.status_code = 200
|
||||
response.headers = {}
|
||||
response.content = json.dumps(MOCK_NIMBLE_RESPONSE).encode()
|
||||
return response
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="Local only tested search providers")
|
||||
class TestNimbleSearch(BaseSearchTest):
|
||||
"""
|
||||
E2E tests for Nimble Search functionality that make real API calls.
|
||||
Inherits from BaseSearchTest to run standard search tests.
|
||||
"""
|
||||
|
||||
def get_search_provider(self) -> str:
|
||||
return "nimble"
|
||||
|
||||
|
||||
class TestNimbleSearchTransformation:
|
||||
"""
|
||||
Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked.
|
||||
Transformation details are unit-tested in tests/test_litellm/llms/nimble/search/.
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _server_key(self, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("NIMBLE_API_KEY", "test-api-key")
|
||||
monkeypatch.delenv("NIMBLE_API_BASE", raising=False)
|
||||
|
||||
def test_nimble_search_request_and_response(self):
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
|
||||
return_value=_mock_response(),
|
||||
) as mock_post:
|
||||
response = litellm.search(
|
||||
query="nimble web scraping",
|
||||
search_provider="nimble",
|
||||
max_results=2,
|
||||
country="us",
|
||||
search_domain_filter=["nimbleway.com", "-spam.example"],
|
||||
)
|
||||
|
||||
assert mock_post.called
|
||||
call_kwargs = mock_post.call_args.kwargs
|
||||
assert call_kwargs["url"] == "https://sdk.nimbleway.com/v2/search"
|
||||
assert call_kwargs["headers"]["Authorization"] == "Bearer test-api-key"
|
||||
assert call_kwargs["headers"]["X-Client-Source"] == "litellm"
|
||||
|
||||
request_body = call_kwargs["json"]
|
||||
assert request_body["query"] == "nimble web scraping"
|
||||
assert request_body["max_results"] == 2
|
||||
assert request_body["country"] == "US"
|
||||
assert request_body["include_domains"] == ("nimbleway.com",)
|
||||
assert request_body["exclude_domains"] == ("spam.example",)
|
||||
|
||||
assert response.object == "search"
|
||||
assert len(response.results) == 2
|
||||
assert response.results[0].title == "Nimble Web API"
|
||||
assert response.results[0].url == "https://nimbleway.com/"
|
||||
assert response.results[0].snippet == "Full markdown content for the first result"
|
||||
assert response.results[0].date == "2026-07-15"
|
||||
# Second result has no `content`, so the SERP description is the snippet.
|
||||
assert response.results[1].snippet == "Only a description here"
|
||||
assert response.results[1].date is None
|
||||
|
||||
def test_provider_specific_params_survive_to_the_wire(self):
|
||||
"""Nimble-native params must not be eaten by `filter_out_litellm_params`."""
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
|
||||
return_value=_mock_response(),
|
||||
) as mock_post:
|
||||
litellm.search(
|
||||
query="test query",
|
||||
search_provider="nimble",
|
||||
focus="news",
|
||||
search_depth="deep",
|
||||
time_range="week",
|
||||
locale="fr",
|
||||
output_format="plain_text",
|
||||
max_subagents=5,
|
||||
)
|
||||
|
||||
request_body = mock_post.call_args.kwargs["json"]
|
||||
assert request_body["focus"] == "news"
|
||||
assert request_body["search_depth"] == "deep"
|
||||
assert request_body["time_range"] == "week"
|
||||
assert request_body["locale"] == "fr"
|
||||
assert request_body["output_format"] == "plain_text"
|
||||
assert request_body["max_subagents"] == 5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_nimble_asearch(self):
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
|
||||
new=AsyncMock(return_value=_mock_response()),
|
||||
) as mock_post:
|
||||
response = await litellm.asearch(
|
||||
query="latest ai developments",
|
||||
search_provider="nimble",
|
||||
focus="news",
|
||||
)
|
||||
|
||||
assert mock_post.call_args.kwargs["json"]["focus"] == "news"
|
||||
assert len(response.results) == 2
|
||||
|
||||
def test_nimble_search_tracks_cost(self):
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post",
|
||||
return_value=_mock_response(),
|
||||
):
|
||||
response = litellm.search(query="pricing check", search_provider="nimble")
|
||||
|
||||
assert response._hidden_params["response_cost"] == pytest.approx(0.005)
|
||||
|
|
@ -6,6 +6,7 @@ from datetime import datetime, timedelta, timezone
|
|||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
|
||||
|
|
@ -19,7 +20,7 @@ from litellm.integrations.shadow_eval_logger import (
|
|||
_sample_hits,
|
||||
_unmask_preference,
|
||||
)
|
||||
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN
|
||||
from litellm.types.utils import SHADOW_EVAL_JUDGE_CALL_ORIGIN, SHADOW_EVAL_ROUTER_CALL_ORIGIN, ModelResponse
|
||||
|
||||
|
||||
def _job(**overrides) -> ActiveShadowEvalJob:
|
||||
|
|
@ -51,6 +52,8 @@ def _job_record(job: ActiveShadowEvalJob, api_key_id="key-hash") -> MagicMock:
|
|||
id=job.id,
|
||||
api_key_id=api_key_id,
|
||||
router_name=job.router_name,
|
||||
direction=job.direction,
|
||||
baseline_model=job.baseline_model,
|
||||
shadow_percentage=job.shadow_percentage,
|
||||
judge_model=job.judge_model,
|
||||
max_turns=job.max_turns,
|
||||
|
|
@ -61,40 +64,55 @@ def _job_record(job: ActiveShadowEvalJob, api_key_id="key-hash") -> MagicMock:
|
|||
|
||||
|
||||
def _router(shadow_text="shadow answer", judge_json='{"preference": "A", "confidence": 0.9, "reasoning": "x"}'):
|
||||
"""One mock router serving the shadow call first, the judge call second. The shadow
|
||||
call's metadata receives the routing decision write-back, like the real router."""
|
||||
"""One mock router serving the shadow call first, the judge call second, told apart by
|
||||
the internal-origin stamp rather than the model, since a reverse job's shadow arm names
|
||||
a plain model. Only the auto-router writes a routing decision back, and only a plain
|
||||
model reports the model it served on the response, which is how each direction learns
|
||||
which model answered."""
|
||||
router = MagicMock()
|
||||
router.model_group_alias = {}
|
||||
router.get_model_list = MagicMock(return_value=[{"litellm_params": {"model": "openai/gpt-4o-mini"}}])
|
||||
|
||||
async def acompletion(**kwargs):
|
||||
if kwargs["metadata"].get(INTERNAL_CALL_ORIGIN_METADATA_KEY) != SHADOW_EVAL_ROUTER_CALL_ORIGIN:
|
||||
return {"choices": [{"message": {"content": judge_json}}]}
|
||||
if kwargs["model"] == "my-router":
|
||||
kwargs["metadata"]["routing_decision"] = {"tier_label": "SIMPLE", "routed_model": "cheap-model"}
|
||||
return {"choices": [{"message": {"content": shadow_text}}], "usage": {"completion_tokens": 5}}
|
||||
return {"choices": [{"message": {"content": judge_json}}]}
|
||||
return ModelResponse(
|
||||
model=kwargs["model"],
|
||||
choices=[{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": shadow_text}}],
|
||||
)
|
||||
|
||||
router.acompletion = MagicMock(side_effect=acompletion)
|
||||
return router
|
||||
|
||||
|
||||
def _logger(router=None, prisma=None, job=None) -> ShadowEvalLogger:
|
||||
def _logger(router=None, prisma=None, jobs=()) -> ShadowEvalLogger:
|
||||
cache = InMemoryCache(max_size_in_memory=4, default_ttl=60)
|
||||
logger = ShadowEvalLogger(
|
||||
router_provider=lambda: router,
|
||||
prisma_provider=lambda: prisma,
|
||||
jobs_cache=cache,
|
||||
)
|
||||
if job is not None:
|
||||
cache.set_cache("shadow_eval:active_jobs", {"key-hash": job})
|
||||
if jobs:
|
||||
cache.set_cache("shadow_eval:active_jobs", {"key-hash": tuple(jobs)})
|
||||
return logger
|
||||
|
||||
|
||||
def _success_kwargs(request_id="req-1", api_key_hash="key-hash", request_metadata=None, call_type="acompletion"):
|
||||
def _routed_by(router_name="my-router", tier="COMPLEX"):
|
||||
"""Metadata as a pre-routing strategy leaves it on the request it served."""
|
||||
return {"routing_decision": {"router_model_name": router_name, "tier_label": tier, "routed_model": "router-pick"}}
|
||||
|
||||
|
||||
def _success_kwargs(
|
||||
request_id="req-1", api_key_hash="key-hash", request_metadata=None, call_type="acompletion", model="claude-opus"
|
||||
):
|
||||
return {
|
||||
"standard_logging_object": {
|
||||
"id": request_id,
|
||||
"call_type": call_type,
|
||||
"model": "claude-opus",
|
||||
"model": model,
|
||||
"metadata": {"user_api_key_hash": api_key_hash},
|
||||
"model_parameters": {"temperature": 0.5, "stream": True},
|
||||
},
|
||||
|
|
@ -164,7 +182,7 @@ class TestSuccessHookSkipChain:
|
|||
monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.005)
|
||||
prisma = _prisma()
|
||||
router = _router()
|
||||
logger = _logger(router=router, prisma=prisma, job=_job())
|
||||
logger = _logger(router=router, prisma=prisma, jobs=(_job(),))
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
|
@ -209,7 +227,7 @@ class TestSuccessHookSkipChain:
|
|||
async def test_skip_paths_store_nothing(self, kwargs_mutation, job_mutation):
|
||||
starts = job_mutation.pop("_starts", 0)
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, job=_job(**job_mutation))
|
||||
logger = _logger(router=_router(), prisma=prisma, jobs=(_job(**job_mutation),))
|
||||
logger._job_starts = {"job-1": starts}
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(**kwargs_mutation), RESPONSE, None, None)
|
||||
|
|
@ -222,7 +240,7 @@ class TestSuccessHookSkipChain:
|
|||
"""A finished pipeline frees its concurrency slot but not its slice of the turn
|
||||
budget; the budget only reopens when a cache refill absorbs the written rows."""
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, job=_job(attempts=199, max_turns=200))
|
||||
logger = _logger(router=_router(), prisma=prisma, jobs=(_job(attempts=199, max_turns=200),))
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(request_id="req-1"), RESPONSE, None, None)
|
||||
await _drain(logger)
|
||||
|
|
@ -237,7 +255,7 @@ class TestSuccessHookSkipChain:
|
|||
identity to the shadow and judge calls."""
|
||||
prisma = _prisma()
|
||||
router = _router()
|
||||
logger = _logger(router=router, prisma=prisma, job=_job())
|
||||
logger = _logger(router=router, prisma=prisma, jobs=(_job(),))
|
||||
|
||||
hook_kwargs = _success_kwargs()
|
||||
hook_kwargs["litellm_params"] = {
|
||||
|
|
@ -256,7 +274,7 @@ class TestSuccessHookSkipChain:
|
|||
predicate, so every redaction source counts."""
|
||||
prisma = _prisma()
|
||||
router = _router()
|
||||
logger = _logger(router=router, prisma=prisma, job=_job())
|
||||
logger = _logger(router=router, prisma=prisma, jobs=(_job(),))
|
||||
|
||||
hook_kwargs = _success_kwargs()
|
||||
hook_kwargs["standard_callback_dynamic_params"] = {"turn_off_message_logging": True}
|
||||
|
|
@ -268,7 +286,7 @@ class TestSuccessHookSkipChain:
|
|||
|
||||
async def test_inflight_cap_sheds_instead_of_queueing(self):
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, job=_job())
|
||||
logger = _logger(router=_router(), prisma=prisma, jobs=(_job(),))
|
||||
logger._inflight_shadow_tasks = _MAX_CONCURRENT_SHADOW_TASKS
|
||||
|
||||
await logger.async_log_success_event(_success_kwargs(), RESPONSE, None, None)
|
||||
|
|
@ -291,8 +309,8 @@ class TestActiveJobsCache:
|
|||
first = await logger._active_jobs()
|
||||
second = await logger._active_jobs()
|
||||
|
||||
assert first["key-hash"].id == "job-1"
|
||||
assert second["key-hash"].attempts == 7
|
||||
assert [job.id for job in first["key-hash"]] == ["job-1"]
|
||||
assert second["key-hash"][0].attempts == 7
|
||||
assert prisma.db.litellm_shadowevaljob.find_many.await_count == 1
|
||||
where = prisma.db.litellm_shadowevaljob.find_many.call_args.kwargs["where"]
|
||||
assert where["stopped_at"] is None
|
||||
|
|
@ -353,6 +371,7 @@ class TestShadowPipeline:
|
|||
messages=({"role": "user", "content": "hi"},),
|
||||
response_obj=RESPONSE,
|
||||
real_model="claude-opus",
|
||||
control_tier=None,
|
||||
model_parameters={},
|
||||
parent_metadata={},
|
||||
)
|
||||
|
|
@ -381,6 +400,7 @@ class TestShadowPipeline:
|
|||
messages=({"role": "user", "content": "hi"},),
|
||||
response_obj=RESPONSE,
|
||||
real_model="claude-opus",
|
||||
control_tier=None,
|
||||
model_parameters={},
|
||||
parent_metadata={"user_api_key_auth": UserAPIKeyAuth(api_key="sk-abc", max_budget=10.0)},
|
||||
)
|
||||
|
|
@ -411,6 +431,7 @@ class TestShadowPipeline:
|
|||
messages=({"role": "user", "content": "hi"},),
|
||||
response_obj=RESPONSE,
|
||||
real_model="claude-opus",
|
||||
control_tier=None,
|
||||
model_parameters={},
|
||||
parent_metadata={},
|
||||
)
|
||||
|
|
@ -438,6 +459,7 @@ class TestShadowPipeline:
|
|||
messages=({"role": "user", "content": "hi"},),
|
||||
response_obj=RESPONSE,
|
||||
real_model="claude-opus",
|
||||
control_tier=None,
|
||||
model_parameters={"stream": True, "temperature": 0.2, "metadata": {"x": 1}},
|
||||
parent_metadata=parent_metadata,
|
||||
)
|
||||
|
|
@ -458,6 +480,164 @@ class TestShadowPipeline:
|
|||
assert judge_call["max_tokens"] == JUDGE_MAX_OUTPUT_TOKENS
|
||||
|
||||
|
||||
def _reverse_job(**overrides) -> ActiveShadowEvalJob:
|
||||
return _job(**{"direction": "reverse", "baseline_model": "baseline-model", **overrides})
|
||||
|
||||
|
||||
class TestJobValidation:
|
||||
@pytest.mark.parametrize(
|
||||
"overrides",
|
||||
[
|
||||
{"direction": "reverse"},
|
||||
{"baseline_model": "baseline-model"},
|
||||
{"direction": "sideways", "baseline_model": "baseline-model"},
|
||||
],
|
||||
ids=["reverse-without-baseline", "forward-with-baseline", "unknown-direction"],
|
||||
)
|
||||
def test_unsamplable_shapes_are_rejected(self, overrides):
|
||||
with pytest.raises(ValidationError):
|
||||
_job(**overrides)
|
||||
|
||||
def test_shadow_target_follows_direction(self):
|
||||
assert _job().shadow_target == "my-router"
|
||||
assert _reverse_job().shadow_target == "baseline-model"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestDirection:
|
||||
@pytest.mark.parametrize(
|
||||
"job,routed_by,sampled",
|
||||
[
|
||||
(_job(), None, True),
|
||||
(_job(), "my-router", False),
|
||||
(_job(), "other-router", True),
|
||||
(_reverse_job(), "my-router", True),
|
||||
(_reverse_job(), None, False),
|
||||
(_reverse_job(), "other-router", False),
|
||||
],
|
||||
ids=[
|
||||
"forward-samples-unrouted",
|
||||
"forward-skips-its-own-router",
|
||||
"forward-samples-another-router",
|
||||
"reverse-samples-its-own-router",
|
||||
"reverse-skips-unrouted",
|
||||
"reverse-skips-another-router",
|
||||
],
|
||||
)
|
||||
async def test_direction_decides_which_traffic_is_sampled(self, job, routed_by, sampled):
|
||||
"""The two directions partition the key's traffic: whatever one samples, the other
|
||||
skips, so a key running both never judges the same turn twice for the same reason."""
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, jobs=(job,))
|
||||
|
||||
await logger.async_log_success_event(
|
||||
_success_kwargs(request_metadata=_routed_by(routed_by) if routed_by else {}), RESPONSE, None, None
|
||||
)
|
||||
await _drain(logger)
|
||||
|
||||
assert prisma.db.litellm_shadowevalattempt.create.await_count == int(sampled)
|
||||
|
||||
async def test_reverse_duplicates_against_the_baseline_model(self):
|
||||
prisma = _prisma()
|
||||
router = _router()
|
||||
logger = _logger(router=router, prisma=prisma, jobs=(_reverse_job(),))
|
||||
|
||||
await logger.async_log_success_event(
|
||||
_success_kwargs(request_metadata=_routed_by()), RESPONSE, None, None
|
||||
)
|
||||
await _drain(logger)
|
||||
|
||||
assert router.acompletion.call_args_list[0].kwargs["model"] == "baseline-model"
|
||||
|
||||
async def test_reverse_row_orients_arms_and_reads_tier_off_the_served_request(self):
|
||||
"""real is what the caller received, so in reverse it is the router's own pick and
|
||||
the tier that produced it; only the shadow arm moves to the baseline."""
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, jobs=(_reverse_job(),))
|
||||
|
||||
await logger.async_log_success_event(
|
||||
_success_kwargs(request_metadata=_routed_by(tier="COMPLEX"), model="router-pick"), RESPONSE, None, None
|
||||
)
|
||||
await _drain(logger)
|
||||
|
||||
row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]
|
||||
assert row["real_model"] == "router-pick"
|
||||
assert row["shadow_model"] == "baseline-model"
|
||||
assert row["tier"] == "COMPLEX"
|
||||
|
||||
async def test_forward_row_still_reads_tier_off_the_shadow_call(self):
|
||||
"""A forward job's tier describes the arm being evaluated, which is the shadow one,
|
||||
so a routing decision on the incumbent request must not leak into it."""
|
||||
prisma = _prisma()
|
||||
logger = _logger(router=_router(), prisma=prisma, jobs=(_job(),))
|
||||
|
||||
await logger.async_log_success_event(
|
||||
_success_kwargs(request_metadata=_routed_by("other-router", tier="CONTROL_TIER")), RESPONSE, None, None
|
||||
)
|
||||
await _drain(logger)
|
||||
|
||||
row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"]
|
||||
assert row["tier"] == "SIMPLE"
|
||||
assert row["shadow_model"] == "cheap-model"
|
||||
|
||||
async def test_a_key_running_both_directions_dispatches_both(self):
|
||||
"""One request can qualify for a forward job on a router that did not serve it and a
|
||||
reverse job on the router that did. The two are separately budgeted experiments, so
|
||||
both fire rather than one silently losing the turn."""
|
||||
prisma = _prisma()
|
||||
logger = _logger(
|
||||
router=_router(),
|
||||
prisma=prisma,
|
||||
jobs=(_job(id="forward-job", router_name="other-router"), _reverse_job(id="reverse-job")),
|
||||
)
|
||||
|
||||
await logger.async_log_success_event(
|
||||
_success_kwargs(request_metadata=_routed_by()), RESPONSE, None, None
|
||||
)
|
||||
await _drain(logger)
|
||||
|
||||
rows = [call.kwargs["data"] for call in prisma.db.litellm_shadowevalattempt.create.call_args_list]
|
||||
assert sorted(row["job_id"] for row in rows) == ["forward-job", "reverse-job"]
|
||||
assert logger._job_starts == {"forward-job": 1, "reverse-job": 1}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestActiveJobsFailClosed:
|
||||
async def test_a_row_the_sampler_cannot_read_is_dropped_not_guessed(self):
|
||||
"""A reverse row with no baseline model has no second arm to call, so it is skipped
|
||||
rather than silently dispatched at the router it is supposed to be judging."""
|
||||
broken = _job_record(_job(id="job-broken"))
|
||||
broken.direction = "reverse"
|
||||
broken.baseline_model = None
|
||||
prisma = _prisma(jobs=[broken, _job_record(_job(id="job-ok"))], attempt_counts=[("job-ok", 1)])
|
||||
logger = ShadowEvalLogger(
|
||||
router_provider=lambda: None,
|
||||
prisma_provider=lambda: prisma,
|
||||
jobs_cache=InMemoryCache(max_size_in_memory=4, default_ttl=60),
|
||||
)
|
||||
|
||||
assert [job.id for job in (await logger._active_jobs())["key-hash"]] == ["job-ok"]
|
||||
|
||||
async def test_both_of_a_key_s_jobs_survive_the_lookup(self):
|
||||
records = [
|
||||
_job_record(_job(id="job-forward")),
|
||||
_job_record(_reverse_job(id="job-reverse")),
|
||||
_job_record(_job(id="job-other"), api_key_id="other-key"),
|
||||
]
|
||||
prisma = _prisma(jobs=records, attempt_counts=[("job-reverse", 3)])
|
||||
logger = ShadowEvalLogger(
|
||||
router_provider=lambda: None,
|
||||
prisma_provider=lambda: prisma,
|
||||
jobs_cache=InMemoryCache(max_size_in_memory=4, default_ttl=60),
|
||||
)
|
||||
|
||||
jobs = await logger._active_jobs()
|
||||
|
||||
assert sorted(job.id for job in jobs["key-hash"]) == ["job-forward", "job-reverse"]
|
||||
assert [job.id for job in jobs["other-key"]] == ["job-other"]
|
||||
assert {job.id: job.attempts for job in jobs["key-hash"]}["job-reverse"] == 3
|
||||
|
||||
|
||||
def _failing_router():
|
||||
router = MagicMock()
|
||||
router.model_group_alias = {}
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig
|
|||
from litellm.llms.firecrawl.search.transformation import FirecrawlSearchConfig
|
||||
from litellm.llms.google_pse.search.transformation import GooglePSESearchConfig
|
||||
from litellm.llms.linkup.search.transformation import LinkupSearchConfig
|
||||
from litellm.llms.nimble.search.transformation import NimbleSearchConfig
|
||||
from litellm.llms.parallel_ai.search.transformation import ParallelAISearchConfig
|
||||
from litellm.llms.perplexity.search.transformation import PerplexitySearchConfig
|
||||
from litellm.llms.searchapi.search.transformation import SearchAPIConfig
|
||||
|
|
@ -57,6 +58,7 @@ _BASE_ENV_VARS = (
|
|||
"DATAFORSEO_API_BASE",
|
||||
"TINYFISH_API_BASE",
|
||||
"CRW_API_BASE",
|
||||
"NIMBLE_API_BASE",
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -96,6 +98,7 @@ PROVIDERS: Tuple[ProviderSpec, ...] = (
|
|||
),
|
||||
(TinyfishSearchConfig, {"TINYFISH_API_KEY": "srv"}, "caller-key", {}),
|
||||
(FastCRWSearchConfig, {"CRW_API_KEY": "srv"}, "caller-key", {}),
|
||||
(NimbleSearchConfig, {"NIMBLE_API_KEY": "srv"}, "caller-key", {}),
|
||||
)
|
||||
|
||||
_IDS = tuple(spec[0].__name__ for spec in PROVIDERS)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,251 @@
|
|||
import json
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.nimble.search.transformation import NimbleSearchConfig
|
||||
|
||||
|
||||
def _config() -> NimbleSearchConfig:
|
||||
return NimbleSearchConfig()
|
||||
|
||||
|
||||
def _resp(payload, status_code: int = 200):
|
||||
r = Mock()
|
||||
r.status_code = status_code
|
||||
r.headers = {}
|
||||
r.content = (payload if isinstance(payload, str) else json.dumps(payload)).encode()
|
||||
return r
|
||||
|
||||
|
||||
def _result(**overrides):
|
||||
base = {
|
||||
"title": "Test Title",
|
||||
"description": "Test description",
|
||||
"url": "https://example.com",
|
||||
"content": "Test content",
|
||||
"metadata": {"position": 1, "entity_type": "organic"},
|
||||
"additional_data": None,
|
||||
}
|
||||
return {**base, **overrides}
|
||||
|
||||
|
||||
def test_ui_friendly_name():
|
||||
assert _config().ui_friendly_name() == "Nimble"
|
||||
|
||||
|
||||
def test_validate_environment_with_explicit_key():
|
||||
headers = _config().validate_environment({}, api_key="explicit-key")
|
||||
assert headers["Authorization"] == "Bearer explicit-key"
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
assert headers["X-Client-Source"] == "litellm"
|
||||
|
||||
|
||||
def test_validate_environment_reads_env_key(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("NIMBLE_API_KEY", "env-key")
|
||||
assert _config().validate_environment({})["Authorization"] == "Bearer env-key"
|
||||
|
||||
|
||||
def test_validate_environment_missing_key_raises(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.delenv("NIMBLE_API_KEY", raising=False)
|
||||
with pytest.raises(ValueError, match="NIMBLE_API_KEY"):
|
||||
_config().validate_environment({})
|
||||
|
||||
|
||||
def test_validate_environment_does_not_mutate_and_is_idempotent():
|
||||
"""The http handler re-runs validate_environment after search/main.py already did."""
|
||||
config = _config()
|
||||
caller_headers = {"X-Custom": "keep-me"}
|
||||
|
||||
once = config.validate_environment(caller_headers, api_key="k")
|
||||
twice = config.validate_environment(once, api_key="k")
|
||||
|
||||
assert caller_headers == {"X-Custom": "keep-me"}
|
||||
assert once == twice
|
||||
assert once["X-Custom"] == "keep-me"
|
||||
|
||||
|
||||
def test_get_complete_url_default_base(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.delenv("NIMBLE_API_BASE", raising=False)
|
||||
assert _config().get_complete_url(None, {}) == "https://sdk.nimbleway.com/v2/search"
|
||||
|
||||
|
||||
def test_get_complete_url_reads_env_base(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setenv("NIMBLE_API_BASE", "https://env-base.local/v2")
|
||||
assert _config().get_complete_url(None, {}) == "https://env-base.local/v2/search"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base",
|
||||
[
|
||||
"https://self-hosted.local/v2",
|
||||
"https://self-hosted.local/v2/",
|
||||
"https://self-hosted.local/v2/search",
|
||||
"https://self-hosted.local/v2/search/",
|
||||
],
|
||||
)
|
||||
def test_get_complete_url_appends_search_exactly_once(api_base: str):
|
||||
assert _config().get_complete_url(api_base, {}) == "https://self-hosted.local/v2/search"
|
||||
|
||||
|
||||
def test_transform_search_request_joins_list_query():
|
||||
assert _config().transform_search_request(["foo", "bar"], {})["query"] == "foo bar"
|
||||
|
||||
|
||||
def test_transform_search_request_max_results_is_not_clamped():
|
||||
"""Nimble validates 1-100 itself; a clearer error beats silently rewriting the request."""
|
||||
assert _config().transform_search_request("q", {"max_results": 500})["max_results"] == 500
|
||||
|
||||
|
||||
def test_transform_search_request_uppercases_country():
|
||||
assert _config().transform_search_request("q", {"country": "us"})["country"] == "US"
|
||||
|
||||
|
||||
def test_transform_search_request_drops_max_tokens_per_page():
|
||||
assert "max_tokens_per_page" not in _config().transform_search_request("q", {"max_tokens_per_page": 1024})
|
||||
|
||||
|
||||
def test_transform_search_request_splits_domain_filter():
|
||||
data = _config().transform_search_request("q", {"search_domain_filter": ["arxiv.org", "-spam.com", "nature.com"]})
|
||||
assert data["include_domains"] == ("arxiv.org", "nature.com")
|
||||
assert data["exclude_domains"] == ("spam.com",)
|
||||
|
||||
|
||||
def test_transform_search_request_omits_empty_domain_lists():
|
||||
data = _config().transform_search_request("q", {"search_domain_filter": ["arxiv.org"]})
|
||||
assert data["include_domains"] == ("arxiv.org",)
|
||||
assert "exclude_domains" not in data
|
||||
|
||||
|
||||
def test_transform_search_request_ignores_non_list_domain_filter():
|
||||
assert "include_domains" not in _config().transform_search_request("q", {"search_domain_filter": "arxiv.org"})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("native_key", ["include_domains", "exclude_domains"])
|
||||
def test_transform_search_request_native_domains_win(native_key: str):
|
||||
"""An explicit provider-native value must not be silently clobbered by the unified param."""
|
||||
data = _config().transform_search_request(
|
||||
"q",
|
||||
{"search_domain_filter": ["derived.com", "-derived-ex.com"], native_key: ["native.com"]},
|
||||
)
|
||||
assert data[native_key] == ["native.com"]
|
||||
|
||||
|
||||
def test_transform_search_response_prefers_content():
|
||||
resp = _config().transform_search_response(_resp({"results": [_result()]}), logging_obj=Mock())
|
||||
assert resp.results[0].snippet == "Test content"
|
||||
|
||||
|
||||
def test_transform_search_response_falls_back_to_description():
|
||||
resp = _config().transform_search_response(_resp({"results": [_result(content="")]}), logging_obj=Mock())
|
||||
assert resp.results[0].snippet == "Test description"
|
||||
|
||||
|
||||
def test_transform_search_response_reads_publish_date():
|
||||
resp = _config().transform_search_response(
|
||||
_resp({"results": [_result(additional_data={"publish_date": "2026-08-01"})]}),
|
||||
logging_obj=Mock(),
|
||||
)
|
||||
assert resp.results[0].date == "2026-08-01"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("additional_data", [{}, "not-a-dict"])
|
||||
def test_transform_search_response_date_is_none_without_usable_publish_date(additional_data):
|
||||
resp = _config().transform_search_response(
|
||||
_resp({"results": [_result(additional_data=additional_data)]}), logging_obj=Mock()
|
||||
)
|
||||
assert resp.results[0].date is None
|
||||
|
||||
|
||||
def test_transform_search_response_keeps_additional_data():
|
||||
"""News results often carry only a relative `publish_date_raw`, which is not a date;
|
||||
it must still reach the caller rather than being dropped on the floor."""
|
||||
resp = _config().transform_search_response(
|
||||
_resp({"results": [_result(additional_data={"publish_date_raw": "1 day ago"})]}),
|
||||
logging_obj=Mock(),
|
||||
)
|
||||
assert resp.results[0].date is None
|
||||
assert resp.results[0].additional_data == {"publish_date_raw": "1 day ago"}
|
||||
|
||||
|
||||
def test_transform_search_response_omits_additional_data_when_absent():
|
||||
resp = _config().transform_search_response(_resp({"results": [_result()]}), logging_obj=Mock())
|
||||
assert not hasattr(resp.results[0], "additional_data")
|
||||
|
||||
|
||||
def test_transform_search_response_preserves_order():
|
||||
resp = _config().transform_search_response(
|
||||
_resp({"results": [_result(title=t) for t in ("first", "second", "third")]}),
|
||||
logging_obj=Mock(),
|
||||
)
|
||||
assert [r.title for r in resp.results] == ["first", "second", "third"]
|
||||
|
||||
|
||||
def test_transform_search_response_degraded_result_does_not_fail_the_call():
|
||||
resp = _config().transform_search_response(
|
||||
_resp({"results": [{"url": "https://example.com"}, _result()]}), logging_obj=Mock()
|
||||
)
|
||||
assert len(resp.results) == 2
|
||||
assert resp.results[0].title == ""
|
||||
assert resp.results[0].snippet == ""
|
||||
assert resp.results[1].title == "Test Title"
|
||||
|
||||
|
||||
def test_transform_search_response_zero_hits():
|
||||
"""A search with no hits really does come back as `"results": []`."""
|
||||
payload = {"request_id": "abc", "total_results": 0, "results": []}
|
||||
assert _config().transform_search_response(_resp(payload), logging_obj=Mock()).results == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"body",
|
||||
[
|
||||
"<html>502 Bad Gateway</html>", # non-JSON body
|
||||
'{"results": ["garbage"]}', # right key, wrong element shape
|
||||
'{"results": {"unexpected": "shape"}}',
|
||||
'{"results": null}', # must not degrade to a successful empty search
|
||||
"{}", # ditto for an absent key
|
||||
],
|
||||
)
|
||||
def test_transform_search_response_malformed_body_raises_instead_of_reporting_empty(body: str):
|
||||
"""A body LiteLLM cannot parse must not be reported as a successful zero-result search."""
|
||||
with pytest.raises(Exception, match="Nimble Search"):
|
||||
_config().transform_search_response(_resp(body, status_code=502), logging_obj=Mock())
|
||||
|
||||
|
||||
def test_get_error_class_attributes_the_provider():
|
||||
error = _config().get_error_class(error_message="quota exceeded", status_code=429, headers={})
|
||||
assert error.status_code == 429
|
||||
assert "Nimble Search: quota exceeded" in str(error)
|
||||
assert "docs.nimbleway.com" in str(error)
|
||||
|
||||
|
||||
def test_get_error_class_unwraps_nimble_detail_envelope():
|
||||
"""Verbatim body from a live 422; the raw JSON envelope should not reach the user."""
|
||||
error = _config().get_error_class(
|
||||
error_message='{"detail":"search_depth=\'fast\' is only supported with focus=\'general\'."}',
|
||||
status_code=422,
|
||||
headers={},
|
||||
)
|
||||
assert (
|
||||
str(error) == "Nimble Search: search_depth='fast' is only supported with focus='general'. "
|
||||
"See https://docs.nimbleway.com/api-reference/search/search for details."
|
||||
)
|
||||
|
||||
|
||||
def test_get_error_class_unwraps_nimble_message_envelope():
|
||||
"""Verbatim body from a live collection failure, which uses a different envelope."""
|
||||
error = _config().get_error_class(
|
||||
error_message='{"success":"false","task_id":"4f74af04","message":"can\'t download the query response"}',
|
||||
status_code=500,
|
||||
headers={},
|
||||
)
|
||||
assert (
|
||||
str(error) == "Nimble Search: can't download the query response. "
|
||||
"See https://docs.nimbleway.com/api-reference/search/search for details."
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("body", ["<html>502 Bad Gateway</html>", '{"detail": null}'])
|
||||
def test_get_error_class_falls_back_to_the_raw_body(body: str):
|
||||
assert f"Nimble Search: {body}." in str(_config().get_error_class(body, status_code=500, headers={}))
|
||||
|
|
@ -4,12 +4,32 @@ import pytest
|
|||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
_upstream_credential_headers,
|
||||
build_synthetic_mcp_request,
|
||||
logging_safe_mcp_headers,
|
||||
validate_and_normalize_mcp_server_payload,
|
||||
validate_tool_display_names,
|
||||
)
|
||||
from litellm.proxy._types import NewMCPServerRequest
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def _server_forwarding(*header_names: str) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id="srv-1",
|
||||
name="deepwiki",
|
||||
transport="http",
|
||||
url="https://mcp.example.com/mcp",
|
||||
extra_headers=list(header_names),
|
||||
)
|
||||
|
||||
|
||||
def _configured_servers(*servers: MCPServer):
|
||||
return patch.dict(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.config_mcp_servers",
|
||||
{server.server_id: server for server in servers},
|
||||
clear=False,
|
||||
)
|
||||
|
||||
|
||||
class TestValidateToolDisplayNames:
|
||||
|
|
@ -114,6 +134,70 @@ class TestLoggingSafeMcpHeaders:
|
|||
|
||||
assert safe == {"x-nuid": "nuid-1"}
|
||||
|
||||
def test_strips_headers_a_server_forwards_upstream(self):
|
||||
"""mcp_servers.<name>.extra_headers names the headers the proxy relays upstream, so a
|
||||
caller supplied value under one of them is an upstream credential no prefix rule can spot.
|
||||
Config is written in canonical casing while the wire header arrives lowercased."""
|
||||
with _configured_servers(_server_forwarding("X-GitHub-Token", "X-Tenant")):
|
||||
safe = logging_safe_mcp_headers({"x-github-token": "ghp_secret", "x-tenant": "acct-1", "x-nuid": "nuid-1"})
|
||||
|
||||
assert safe == {"x-nuid": "nuid-1"}
|
||||
|
||||
def test_strips_caller_asserted_host(self):
|
||||
"""This mapping reaches the guardrail payload and the list_tools spend row, so a caller
|
||||
must not be able to name the deployment there either."""
|
||||
safe = logging_safe_mcp_headers({"host": "evil.attacker.example", "x-nuid": "nuid-1"})
|
||||
|
||||
assert safe == {"x-nuid": "nuid-1"}
|
||||
|
||||
def test_keeps_identity_header_a_server_also_forwards(self):
|
||||
"""get_user_from_headers resolves end user attribution off this same request, so a header
|
||||
the deployment reads identity from stays even when a server forwards it upstream."""
|
||||
with patch.dict(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"user_header_name": "x-user-email"},
|
||||
clear=False,
|
||||
):
|
||||
with _configured_servers(_server_forwarding("x-user-email", "x-github-token")):
|
||||
safe = logging_safe_mcp_headers({"x-user-email": "alice@corp.example", "x-github-token": "ghp_secret"})
|
||||
|
||||
assert safe == {"x-user-email": "alice@corp.example"}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"configured",
|
||||
[
|
||||
[{"header_name": "X-User", "litellm_user_role": "customer"}],
|
||||
{"header_name": "X-User", "litellm_user_role": "customer"},
|
||||
],
|
||||
ids=["list-of-mappings", "bare-mapping"],
|
||||
)
|
||||
def test_keeps_identity_header_from_user_header_mappings(self, configured):
|
||||
"""get_internal_user_header_from_mapping and get_customer_user_header_from_mapping both
|
||||
accept a bare mapping as well as a list, and config_settings.md documents the key as a
|
||||
dict, so the exemption has to read both shapes."""
|
||||
with patch.dict(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"user_header_mappings": configured},
|
||||
clear=False,
|
||||
):
|
||||
with _configured_servers(_server_forwarding("X-User", "X-GitHub-Token")):
|
||||
safe = logging_safe_mcp_headers({"x-user": "alice", "x-github-token": "ghp_secret"})
|
||||
|
||||
assert safe == {"x-user": "alice"}
|
||||
|
||||
def test_keeps_authorization_classification_for_oauth_passthrough(self):
|
||||
"""clean_headers already strips authorization, and claiming it here would change which
|
||||
header authenticated_with_header resolves to on a config that lists it by design."""
|
||||
with _configured_servers(_server_forwarding("Authorization", "X-GitHub-Token")):
|
||||
assert "authorization" not in _upstream_credential_headers(["authorization", "x-github-token"])
|
||||
assert "x-github-token" in _upstream_credential_headers(["authorization", "x-github-token"])
|
||||
|
||||
def test_keeps_headers_when_no_server_forwards_them(self):
|
||||
with _configured_servers(_server_forwarding("x-github-token")):
|
||||
safe = logging_safe_mcp_headers({"x-other-token": "not-forwarded", "x-nuid": "nuid-1"})
|
||||
|
||||
assert safe == {"x-other-token": "not-forwarded", "x-nuid": "nuid-1"}
|
||||
|
||||
|
||||
class TestBuildSyntheticMcpRequest:
|
||||
def test_forwards_client_headers_without_upstream_credentials(self):
|
||||
|
|
@ -147,3 +231,41 @@ class TestBuildSyntheticMcpRequest:
|
|||
|
||||
assert request.headers.get("x-nuid") == "nuid-1"
|
||||
assert "x-company-key" not in request.headers
|
||||
|
||||
def test_drops_caller_host_so_the_logged_url_is_not_client_steerable(self):
|
||||
"""add_litellm_data_to_request records str(request.url) as proxy_server_request.url, and
|
||||
Request.url is built from the host header, so forwarding it hands the caller that value."""
|
||||
request = build_synthetic_mcp_request(
|
||||
path="/mcp/tools/call",
|
||||
raw_headers={"host": "evil.attacker.example", "x-nuid": "nuid-1"},
|
||||
)
|
||||
|
||||
assert "evil.attacker.example" not in str(request.url)
|
||||
assert "host" not in request.headers
|
||||
assert request.headers.get("x-nuid") == "nuid-1"
|
||||
|
||||
def test_drops_headers_a_server_forwards_upstream(self):
|
||||
with _configured_servers(_server_forwarding("x-github-token")):
|
||||
request = build_synthetic_mcp_request(
|
||||
path="/mcp/tools/call",
|
||||
raw_headers={"x-github-token": "ghp_secret", "x-nuid": "nuid-1"},
|
||||
)
|
||||
|
||||
assert "x-github-token" not in request.headers
|
||||
assert request.headers.get("x-nuid") == "nuid-1"
|
||||
|
||||
def test_keeps_identity_header_so_end_user_attribution_survives(self):
|
||||
"""add_litellm_data_to_request reads user_header_name off this request to fill
|
||||
end_user_id, so forwarding that header upstream must not remove it here."""
|
||||
with patch.dict(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"user_header_name": "x-user-email"},
|
||||
clear=False,
|
||||
):
|
||||
with _configured_servers(_server_forwarding("x-user-email")):
|
||||
request = build_synthetic_mcp_request(
|
||||
path="/mcp/tools/call",
|
||||
raw_headers={"x-user-email": "alice@corp.example"},
|
||||
)
|
||||
|
||||
assert request.headers.get("x-user-email") == "alice@corp.example"
|
||||
|
|
|
|||
|
|
@ -270,6 +270,70 @@ async def test_get_all_transactions_from_redis_buffer_pipeline_no_redis():
|
|||
assert result == (None, None, None, None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_restore_transactions_to_redis_pushes_only_provided(
|
||||
redis_update_buffer, mock_redis_cache
|
||||
):
|
||||
"""
|
||||
restore_transactions_to_redis re-pushes only the transaction sets it was
|
||||
given, to their matching buffer keys, so uncommitted spend can be retried.
|
||||
"""
|
||||
from litellm.constants import (
|
||||
REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY,
|
||||
REDIS_UPDATE_BUFFER_KEY,
|
||||
)
|
||||
|
||||
mock_redis_cache.async_rpush_pipeline = AsyncMock(return_value=[1, 1])
|
||||
|
||||
db_spend = {"key_list_transactions": {"key1": 1.0}}
|
||||
daily_user = {"user_key1": {"spend": 1.0}}
|
||||
|
||||
await redis_update_buffer.restore_transactions_to_redis(
|
||||
db_spend_update_transactions=db_spend,
|
||||
daily_spend_update_transactions=daily_user,
|
||||
)
|
||||
|
||||
mock_redis_cache.async_rpush_pipeline.assert_called_once()
|
||||
rpush_list = mock_redis_cache.async_rpush_pipeline.call_args.kwargs["rpush_list"]
|
||||
pushed_keys = {op["key"] for op in rpush_list}
|
||||
assert pushed_keys == {
|
||||
REDIS_UPDATE_BUFFER_KEY,
|
||||
REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY,
|
||||
}
|
||||
# Payloads round-trip through the same JSON encoding used on the store path
|
||||
payloads = {op["key"]: json.loads(op["values"][0]) for op in rpush_list}
|
||||
assert payloads[REDIS_UPDATE_BUFFER_KEY] == db_spend
|
||||
assert payloads[REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY] == daily_user
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_restore_transactions_to_redis_noop_when_empty(
|
||||
redis_update_buffer, mock_redis_cache
|
||||
):
|
||||
"""Nothing to restore -> no Redis call."""
|
||||
mock_redis_cache.async_rpush_pipeline = AsyncMock()
|
||||
await redis_update_buffer.restore_transactions_to_redis()
|
||||
mock_redis_cache.async_rpush_pipeline.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_restore_transactions_to_redis_swallows_redis_error(
|
||||
redis_update_buffer, mock_redis_cache
|
||||
):
|
||||
"""A Redis failure during restore must not propagate to the caller's finally block."""
|
||||
from redis.exceptions import RedisError
|
||||
|
||||
mock_redis_cache.async_rpush_pipeline = AsyncMock(
|
||||
side_effect=RedisError("redis down")
|
||||
)
|
||||
|
||||
await redis_update_buffer.restore_transactions_to_redis(
|
||||
db_spend_update_transactions={"key_list_transactions": {"key1": 1.0}},
|
||||
)
|
||||
|
||||
mock_redis_cache.async_rpush_pipeline.assert_called_once()
|
||||
|
||||
|
||||
def test_validate_redis_transaction_buffer_raises_without_redis():
|
||||
"""
|
||||
When use_redis_transaction_buffer=true but no Redis cache is configured,
|
||||
|
|
|
|||
|
|
@ -1425,6 +1425,52 @@ async def test_update_daily_spend_re_raises_exception_after_logging():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_daily_spend_keeps_failed_transactions_for_retry():
|
||||
"""
|
||||
A failed batch must stay in the caller's transaction dict, otherwise the
|
||||
Redis re-queue in _commit_spend_updates_to_db_with_redis has nothing left to
|
||||
push back and the spend is lost permanently.
|
||||
"""
|
||||
|
||||
def raise_outage():
|
||||
raise ValueError("simulated database outage")
|
||||
|
||||
prisma_client = _RecordingPrisma(execute_raw=raise_outage)
|
||||
|
||||
daily_spend_transactions = {
|
||||
"test_key": {
|
||||
"user_id": "test-user",
|
||||
"date": "2024-01-01",
|
||||
"api_key": "test-api-key",
|
||||
"model": "gpt-4",
|
||||
"custom_llm_provider": "openai",
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 20,
|
||||
"spend": 0.1,
|
||||
"api_requests": 1,
|
||||
"successful_requests": 1,
|
||||
"failed_requests": 0,
|
||||
}
|
||||
}
|
||||
expected = dict(daily_spend_transactions)
|
||||
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.failure_handler = AsyncMock()
|
||||
|
||||
with pytest.raises(ValueError, match="simulated database outage"):
|
||||
await DBSpendUpdateWriter._update_daily_spend(
|
||||
n_retry_times=0,
|
||||
prisma_client=prisma_client,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
daily_spend_transactions=daily_spend_transactions,
|
||||
entity_type="user",
|
||||
entity_id_field="user_id",
|
||||
)
|
||||
|
||||
assert daily_spend_transactions == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_key_spend_updates_includes_last_active():
|
||||
"""
|
||||
|
|
@ -1685,9 +1731,9 @@ async def test_commit_spend_updates_uses_pipeline():
|
|||
|
||||
mock_redis_update_buffer = AsyncMock()
|
||||
mock_redis_update_buffer.store_in_memory_spend_updates_in_redis = AsyncMock()
|
||||
# Return all-None tuple (no data to commit)
|
||||
# Return all-None tuple (no data to commit); the pipeline yields 6 slots
|
||||
mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = (
|
||||
AsyncMock(return_value=(None, None, None, None, None, None, None))
|
||||
AsyncMock(return_value=(None, None, None, None, None, None))
|
||||
)
|
||||
db_writer.redis_update_buffer = mock_redis_update_buffer
|
||||
|
||||
|
|
@ -1718,6 +1764,225 @@ async def test_commit_spend_updates_uses_pipeline():
|
|||
mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_with_redis_requeues_all_on_db_failure():
|
||||
"""
|
||||
Regression for #33872: if the DB commit fails after the leader has already
|
||||
popped transactions from Redis, the popped transactions must be re-queued to
|
||||
Redis so a later tick can retry them, instead of being silently lost.
|
||||
"""
|
||||
db_writer = DBSpendUpdateWriter()
|
||||
|
||||
db_spend = {
|
||||
"user_list_transactions": {"user1": 1.5},
|
||||
"end_user_list_transactions": {},
|
||||
"key_list_transactions": {"key1": 1.5},
|
||||
"team_list_transactions": {},
|
||||
"team_member_list_transactions": {},
|
||||
"org_list_transactions": {},
|
||||
"tag_list_transactions": {},
|
||||
"agent_list_transactions": {},
|
||||
}
|
||||
daily_user = {"user_key1": {"spend": 1.5, "api_requests": 1}}
|
||||
|
||||
mock_redis_update_buffer = AsyncMock()
|
||||
mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = AsyncMock(
|
||||
return_value=(db_spend, daily_user, None, None, None, None)
|
||||
)
|
||||
mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock()
|
||||
db_writer.redis_update_buffer = mock_redis_update_buffer
|
||||
|
||||
mock_pod_lock_manager = AsyncMock()
|
||||
mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True)
|
||||
mock_pod_lock_manager.release_lock = AsyncMock()
|
||||
db_writer.pod_lock_manager = mock_pod_lock_manager
|
||||
|
||||
# Every DB write raises -> simulates a full database outage
|
||||
db_writer._commit_spend_updates_to_db = AsyncMock(side_effect=Exception("db down"))
|
||||
|
||||
with patch.object(
|
||||
DBSpendUpdateWriter,
|
||||
"update_daily_user_spend",
|
||||
new=AsyncMock(side_effect=Exception("db down")),
|
||||
):
|
||||
await db_writer._commit_spend_updates_to_db_with_redis(
|
||||
prisma_client=MagicMock(),
|
||||
n_retry_times=0,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
# Both failed categories must be re-queued to Redis, nothing lost
|
||||
mock_redis_update_buffer.restore_transactions_to_redis.assert_awaited_once()
|
||||
_, kwargs = mock_redis_update_buffer.restore_transactions_to_redis.call_args
|
||||
assert kwargs["db_spend_update_transactions"] == db_spend
|
||||
assert kwargs["daily_spend_update_transactions"] == daily_user
|
||||
# The lock must still be released
|
||||
mock_pod_lock_manager.release_lock.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_with_redis_only_requeues_failed_category():
|
||||
"""
|
||||
A partial DB failure must not re-queue categories that already committed,
|
||||
otherwise their spend would be double-counted on the next tick.
|
||||
"""
|
||||
db_writer = DBSpendUpdateWriter()
|
||||
|
||||
db_spend = {
|
||||
"user_list_transactions": {"user1": 1.5},
|
||||
"end_user_list_transactions": {},
|
||||
"key_list_transactions": {},
|
||||
"team_list_transactions": {},
|
||||
"team_member_list_transactions": {},
|
||||
"org_list_transactions": {},
|
||||
"tag_list_transactions": {},
|
||||
"agent_list_transactions": {},
|
||||
}
|
||||
daily_user = {"user_key1": {"spend": 1.5, "api_requests": 1}}
|
||||
|
||||
mock_redis_update_buffer = AsyncMock()
|
||||
mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = AsyncMock(
|
||||
return_value=(db_spend, daily_user, None, None, None, None)
|
||||
)
|
||||
mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock()
|
||||
db_writer.redis_update_buffer = mock_redis_update_buffer
|
||||
|
||||
mock_pod_lock_manager = AsyncMock()
|
||||
mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True)
|
||||
mock_pod_lock_manager.release_lock = AsyncMock()
|
||||
db_writer.pod_lock_manager = mock_pod_lock_manager
|
||||
|
||||
# db_spend commits fine; only the daily user commit fails
|
||||
db_writer._commit_spend_updates_to_db = AsyncMock()
|
||||
|
||||
with patch.object(
|
||||
DBSpendUpdateWriter,
|
||||
"update_daily_user_spend",
|
||||
new=AsyncMock(side_effect=Exception("db down")),
|
||||
):
|
||||
await db_writer._commit_spend_updates_to_db_with_redis(
|
||||
prisma_client=MagicMock(),
|
||||
n_retry_times=0,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
mock_redis_update_buffer.restore_transactions_to_redis.assert_awaited_once()
|
||||
_, kwargs = mock_redis_update_buffer.restore_transactions_to_redis.call_args
|
||||
# Only the failed daily category is requeued; the committed db_spend is not
|
||||
assert kwargs == {"daily_spend_update_transactions": daily_user}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_with_redis_no_requeue_on_success():
|
||||
"""When all commits succeed, nothing should be re-queued to Redis."""
|
||||
db_writer = DBSpendUpdateWriter()
|
||||
|
||||
db_spend = {
|
||||
"user_list_transactions": {"user1": 1.5},
|
||||
"end_user_list_transactions": {},
|
||||
"key_list_transactions": {},
|
||||
"team_list_transactions": {},
|
||||
"team_member_list_transactions": {},
|
||||
"org_list_transactions": {},
|
||||
"tag_list_transactions": {},
|
||||
"agent_list_transactions": {},
|
||||
}
|
||||
|
||||
mock_redis_update_buffer = AsyncMock()
|
||||
mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = AsyncMock(
|
||||
return_value=(db_spend, None, None, None, None, None)
|
||||
)
|
||||
mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock()
|
||||
db_writer.redis_update_buffer = mock_redis_update_buffer
|
||||
|
||||
mock_pod_lock_manager = AsyncMock()
|
||||
mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True)
|
||||
mock_pod_lock_manager.release_lock = AsyncMock()
|
||||
db_writer.pod_lock_manager = mock_pod_lock_manager
|
||||
|
||||
db_writer._commit_spend_updates_to_db = AsyncMock()
|
||||
|
||||
await db_writer._commit_spend_updates_to_db_with_redis(
|
||||
prisma_client=MagicMock(),
|
||||
n_retry_times=0,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
mock_redis_update_buffer.restore_transactions_to_redis.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_daily_tag_spend_requeues_on_db_failure():
|
||||
"""A failed daily tag commit must re-queue the popped tag transactions and release the lock."""
|
||||
db_writer = DBSpendUpdateWriter()
|
||||
|
||||
daily_tag = {"tag_key1": {"spend": 1.5, "api_requests": 1}}
|
||||
|
||||
mock_redis_update_buffer = AsyncMock()
|
||||
mock_redis_update_buffer.store_in_memory_daily_tag_spend_updates_in_redis = AsyncMock()
|
||||
mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer = AsyncMock(
|
||||
return_value=daily_tag
|
||||
)
|
||||
mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock()
|
||||
db_writer.redis_update_buffer = mock_redis_update_buffer
|
||||
|
||||
mock_pod_lock_manager = AsyncMock()
|
||||
mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True)
|
||||
mock_pod_lock_manager.release_lock = AsyncMock()
|
||||
db_writer.pod_lock_manager = mock_pod_lock_manager
|
||||
|
||||
with patch.object(
|
||||
DBSpendUpdateWriter,
|
||||
"update_daily_tag_spend",
|
||||
new=AsyncMock(side_effect=Exception("db down")),
|
||||
):
|
||||
await db_writer._commit_daily_tag_spend_to_db_with_redis(
|
||||
prisma_client=MagicMock(),
|
||||
n_retry_times=0,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
mock_redis_update_buffer.restore_transactions_to_redis.assert_awaited_once_with(
|
||||
daily_tag_spend_update_transactions=daily_tag,
|
||||
)
|
||||
mock_pod_lock_manager.release_lock.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_commit_daily_tag_spend_no_requeue_on_success():
|
||||
"""A successful daily tag commit must not re-queue anything."""
|
||||
db_writer = DBSpendUpdateWriter()
|
||||
|
||||
daily_tag = {"tag_key1": {"spend": 1.5, "api_requests": 1}}
|
||||
|
||||
mock_redis_update_buffer = AsyncMock()
|
||||
mock_redis_update_buffer.store_in_memory_daily_tag_spend_updates_in_redis = AsyncMock()
|
||||
mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer = AsyncMock(
|
||||
return_value=daily_tag
|
||||
)
|
||||
mock_redis_update_buffer.restore_transactions_to_redis = AsyncMock()
|
||||
db_writer.redis_update_buffer = mock_redis_update_buffer
|
||||
|
||||
mock_pod_lock_manager = AsyncMock()
|
||||
mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True)
|
||||
mock_pod_lock_manager.release_lock = AsyncMock()
|
||||
db_writer.pod_lock_manager = mock_pod_lock_manager
|
||||
|
||||
with patch.object(
|
||||
DBSpendUpdateWriter,
|
||||
"update_daily_tag_spend",
|
||||
new=AsyncMock(),
|
||||
):
|
||||
await db_writer._commit_daily_tag_spend_to_db_with_redis(
|
||||
prisma_client=MagicMock(),
|
||||
n_retry_times=0,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
mock_redis_update_buffer.restore_transactions_to_redis.assert_not_awaited()
|
||||
mock_pod_lock_manager.release_lock.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bucket_name,input_dict,table_attr,method_name,where_key,expected_order",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -559,7 +559,7 @@ def _start_request(**overrides: object) -> StartShadowEvalRequest:
|
|||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_creates_job_and_frees_expired_or_exhausted_ones(monkeypatch: pytest.MonkeyPatch):
|
||||
"""Expiry and turn-budget exhaustion both end sampling on their own; either must
|
||||
release the one-active-per-key index so a new eval can start."""
|
||||
release the key's slot in the active-job index so a new eval can start."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
prisma = _shadow_prisma()
|
||||
|
|
@ -592,8 +592,21 @@ async def test_start_shadow_eval_creates_job_and_frees_expired_or_exhausted_ones
|
|||
(ADMIN, {"judge_model": "not/a real model!"}, None, 400),
|
||||
(ADMIN, {"judge_model": "my-router"}, None, 400),
|
||||
(ADMIN, {}, "active", 409),
|
||||
(ADMIN, {"direction": "reverse", "baseline_model": "my-router"}, None, 400),
|
||||
(ADMIN, {"direction": "reverse", "baseline_model": "not/a real model!"}, None, 400),
|
||||
(ADMIN, {"direction": "reverse", "baseline_model": "openai/gpt-4o", "router_name": "not-a-router"}, None, 400),
|
||||
],
|
||||
ids=[
|
||||
"non-admin",
|
||||
"view-only",
|
||||
"unknown-router",
|
||||
"unresolvable-judge",
|
||||
"router-as-judge",
|
||||
"already-active",
|
||||
"router-as-baseline",
|
||||
"unresolvable-baseline",
|
||||
"reverse-still-needs-an-auto-router",
|
||||
],
|
||||
ids=["non-admin", "view-only", "unknown-router", "unresolvable-judge", "router-as-judge", "already-active"],
|
||||
)
|
||||
async def test_start_shadow_eval_rejections(
|
||||
monkeypatch: pytest.MonkeyPatch, caller, request_overrides, active, expected_status
|
||||
|
|
@ -609,6 +622,68 @@ async def test_start_shadow_eval_rejections(
|
|||
assert exc.value.status_code == expected_status
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"overrides",
|
||||
[
|
||||
{"direction": "reverse"},
|
||||
{"baseline_model": "openai/gpt-4o"},
|
||||
{"direction": "sideways", "baseline_model": "openai/gpt-4o"},
|
||||
],
|
||||
ids=["reverse-without-baseline", "forward-with-baseline", "unknown-direction"],
|
||||
)
|
||||
def test_start_request_pins_baseline_model_to_reverse(overrides):
|
||||
"""A forward job has no second arm to name and a reverse job cannot run without one,
|
||||
so neither shape reaches the endpoint to be half-validated there."""
|
||||
with pytest.raises(ValidationError):
|
||||
_start_request(**overrides)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_reverse_records_its_arms_and_holds_its_own_slot(monkeypatch: pytest.MonkeyPatch):
|
||||
"""The two directions ask opposite questions of the same key, so a forward job holding
|
||||
the slot must not block a reverse one. The second reverse start still 409s."""
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
prisma = _shadow_prisma()
|
||||
active = {"forward": _job_record()}
|
||||
prisma.db.litellm_shadowevaljob.find_first = AsyncMock(
|
||||
side_effect=lambda where, **_: active.get(str(where.get("direction")))
|
||||
)
|
||||
prisma.db.litellm_shadowevaljob.create = AsyncMock(
|
||||
return_value=_job_record(direction="reverse", baseline_model="openai/gpt-4o")
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
||||
reverse = _start_request(direction="reverse", baseline_model="openai/gpt-4o")
|
||||
response = await start_shadow_eval(reverse, ADMIN)
|
||||
|
||||
assert (response.direction, response.baseline_model) == ("reverse", "openai/gpt-4o")
|
||||
create_data = prisma.db.litellm_shadowevaljob.create.call_args.kwargs["data"]
|
||||
assert create_data["direction"] == "reverse"
|
||||
assert create_data["baseline_model"] == "openai/gpt-4o"
|
||||
|
||||
active["reverse"] = _job_record(id="job-2", direction="reverse")
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await start_shadow_eval(reverse, ADMIN)
|
||||
assert exc.value.status_code == 409
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_forward_leaves_the_baseline_column_empty(monkeypatch: pytest.MonkeyPatch):
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
|
||||
prisma = _shadow_prisma()
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", _shadow_router())
|
||||
|
||||
await start_shadow_eval(_start_request(), ADMIN)
|
||||
|
||||
create_data = prisma.db.litellm_shadowevaljob.create.call_args.kwargs["data"]
|
||||
assert create_data["direction"] == "forward"
|
||||
assert create_data["baseline_model"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_start_shadow_eval_rejects_a_key_this_proxy_does_not_know(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A typo'd api_key_id would otherwise create a job no traffic can ever match."""
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
|||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import Request, Response
|
||||
from fastapi import HTTPException, Request, Response
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
sys.path.insert(
|
||||
|
|
@ -19,10 +19,13 @@ import litellm
|
|||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
BaseOpenAIPassThroughHandler,
|
||||
RouteChecks,
|
||||
azure_proxy_route,
|
||||
bedrock_llm_proxy_route,
|
||||
create_pass_through_route,
|
||||
cursor_proxy_route,
|
||||
get_azure_ai_search_index_from_endpoint,
|
||||
get_vertex_base_url,
|
||||
is_azure_ai_search_service_level_index_create,
|
||||
llm_passthrough_factory_proxy_route,
|
||||
milvus_proxy_route,
|
||||
mistral_proxy_route,
|
||||
|
|
@ -31,7 +34,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
|||
vertex_proxy_route,
|
||||
vllm_proxy_route,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
|
||||
|
||||
|
||||
|
|
@ -3249,3 +3252,221 @@ def test_is_passthrough_request_streaming_tolerates_non_object_bodies(request_bo
|
|||
)
|
||||
|
||||
assert is_passthrough_request_streaming(request_body) is expected
|
||||
|
||||
|
||||
class TestGetAzureAISearchIndexFromEndpoint:
|
||||
"""The operable index is only the segment right after ``indexes``.
|
||||
|
||||
A doc-write path ends in ``.../docs/index``; the trailing ``index`` must not
|
||||
be mistaken for the target, otherwise a caller could be authorized on one
|
||||
index while Azure applies the write to another.
|
||||
"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"endpoint, expected",
|
||||
[
|
||||
("indexes/my-index/docs/index", "my-index"),
|
||||
("indexes/my-index/docs/search", "my-index"),
|
||||
("indexes/my-index", "my-index"),
|
||||
("indexes/my-index?api-version=2024-07-01", "my-index"),
|
||||
("/indexes/my-index/docs/index", "my-index"),
|
||||
("indexes/victim/docs/index", "victim"),
|
||||
("openai/deployments/gpt-4o/chat/completions", None),
|
||||
("indexes", None),
|
||||
("indexes/", None),
|
||||
],
|
||||
)
|
||||
def test_extracts_positional_index_only(self, endpoint, expected):
|
||||
assert get_azure_ai_search_index_from_endpoint(endpoint) == expected
|
||||
|
||||
|
||||
class TestAzureProxyRouteCrossIndexAuthorization:
|
||||
"""Regression tests: the passthrough must authorize the index that the request
|
||||
actually targets (the ``/indexes/{name}`` segment), never a different segment
|
||||
that merely happens to match a managed index the caller can access.
|
||||
"""
|
||||
|
||||
def _request(self, method: str, path: str) -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = method
|
||||
request.headers = {"content-type": "application/json"}
|
||||
request.url = MagicMock()
|
||||
request.url.path = path
|
||||
return request
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authorizes_the_targeted_index(self):
|
||||
index_object = MagicMock()
|
||||
index_object.litellm_params.vector_store_name = "my-store"
|
||||
vector_store = {"litellm_params": {"api_base": "https://svc.search.windows.net"}}
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config"
|
||||
) as mock_get_config,
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint"
|
||||
) as mock_is_allowed,
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.assert_user_can_access_vector_store",
|
||||
new=AsyncMock(),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler",
|
||||
new=AsyncMock(return_value=Response()),
|
||||
),
|
||||
patch.object(litellm, "vector_store_index_registry") as mock_index_registry,
|
||||
patch.object(litellm, "vector_store_registry") as mock_vector_registry,
|
||||
):
|
||||
mock_get_config.return_value.get_auth_credentials.return_value = {"headers": {"api-key": "k"}}
|
||||
mock_index_registry.is_vector_store_index.side_effect = lambda vector_store_index_name: (
|
||||
vector_store_index_name == "my-index"
|
||||
)
|
||||
mock_index_registry.get_vector_store_index_by_name.return_value = index_object
|
||||
mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = vector_store
|
||||
|
||||
await azure_proxy_route(
|
||||
endpoint="indexes/my-index/docs/index",
|
||||
request=self._request("POST", "/azure_ai/indexes/my-index/docs/index"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
|
||||
)
|
||||
|
||||
mock_is_allowed.assert_called_once()
|
||||
assert mock_is_allowed.call_args.kwargs["index_name"] == "my-index"
|
||||
mock_index_registry.get_vector_store_index_by_name.assert_called_once_with(
|
||||
vector_store_index_name="my-index"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_trailing_index_segment_does_not_authorize_a_different_index(self):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
|
||||
return_value=False,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint"
|
||||
) as mock_is_allowed,
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str",
|
||||
return_value="https://azure-openai.example.com",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="azure-key",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler",
|
||||
new=AsyncMock(return_value=Response()),
|
||||
) as mock_handler,
|
||||
patch.object(litellm, "vector_store_index_registry") as mock_index_registry,
|
||||
):
|
||||
mock_index_registry.is_vector_store_index.side_effect = lambda vector_store_index_name: (
|
||||
vector_store_index_name == "index"
|
||||
)
|
||||
|
||||
await azure_proxy_route(
|
||||
endpoint="indexes/victim/docs/index",
|
||||
request=self._request("POST", "/azure_ai/indexes/victim/docs/index"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
|
||||
)
|
||||
|
||||
mock_is_allowed.assert_not_called()
|
||||
mock_handler.assert_awaited_once()
|
||||
assert mock_handler.await_args.kwargs["custom_llm_provider"] == litellm.LlmProviders.AZURE
|
||||
|
||||
|
||||
class TestAzureProxyRouteServiceLevelIndexCreate:
|
||||
"""``POST /indexes`` carries no index name, so the managed-index branch cannot
|
||||
claim it and it would otherwise reach the generic Azure passthrough on the
|
||||
proxy's own credential. The admin-only index management guard has to be
|
||||
enforced on the route itself, not just on the permission gate the route skips.
|
||||
"""
|
||||
|
||||
def _request(self, method: str, path: str) -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = method
|
||||
request.headers = {"content-type": "application/json"}
|
||||
request.url = MagicMock()
|
||||
request.url.path = path
|
||||
return request
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"method, endpoint, expected",
|
||||
[
|
||||
("POST", "indexes", True),
|
||||
("POST", "indexes?api-version=2024-07-01", True),
|
||||
("POST", "/indexes/", True),
|
||||
("POST", "indexes/my-index", False),
|
||||
("POST", "indexes/my-index/docs/index", False),
|
||||
("GET", "indexes", False),
|
||||
("POST", "openai/deployments/gpt-4o/chat/completions", False),
|
||||
],
|
||||
)
|
||||
def test_recognizes_service_level_create(self, method, endpoint, expected):
|
||||
assert is_azure_ai_search_service_level_index_create(method=method, endpoint=endpoint) is expected
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_admin_cannot_create_an_index(self):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str",
|
||||
return_value="https://svc.search.windows.net",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler",
|
||||
new=AsyncMock(return_value=Response()),
|
||||
) as mock_handler,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await azure_proxy_route(
|
||||
endpoint="indexes?api-version=2024-07-01",
|
||||
request=self._request("POST", "/azure_ai/indexes"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
token="sk-team-token",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Only proxy admins can create" in exc_info.value.detail
|
||||
mock_handler.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_can_still_create_an_index(self):
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_secret_str",
|
||||
return_value="https://svc.search.windows.net",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="azure-key",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.BaseOpenAIPassThroughHandler._base_openai_pass_through_handler",
|
||||
new=AsyncMock(return_value=Response()),
|
||||
) as mock_handler,
|
||||
):
|
||||
await azure_proxy_route(
|
||||
endpoint="indexes?api-version=2024-07-01",
|
||||
request=self._request("POST", "/azure_ai/indexes"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
token="sk-admin-token",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
),
|
||||
)
|
||||
|
||||
mock_handler.assert_awaited_once()
|
||||
|
|
|
|||
|
|
@ -2928,3 +2928,187 @@ class TestUpdateVectorStoreAccessControlAndRedaction:
|
|||
params = response["vector_store"]["litellm_params"]
|
||||
assert params["api_key"] == REDACTED_BY_LITELM_STRING
|
||||
assert params["api_base"] == "https://api.openai.com/v1"
|
||||
|
||||
|
||||
class TestAzureAIDocumentWritePassthroughPermission:
|
||||
"""Regression tests for the Azure AI Search passthrough write mapping.
|
||||
|
||||
Azure's batch document write/merge/delete endpoint is
|
||||
``POST /indexes/{name}/docs/index``. A non-admin team holding a ``write``
|
||||
grant on the index must be allowed to call it, while index lifecycle
|
||||
(create / update / delete the index itself) stays proxy-admin only.
|
||||
|
||||
These exercise the real ``AzureAIVectorStoreConfig`` endpoint map on
|
||||
purpose (no mocked provider config), so reverting the map to the old
|
||||
``("PUT", "/docs")`` entry makes ``test_team_with_write_grant_can_upload``
|
||||
fail.
|
||||
"""
|
||||
|
||||
INDEX = "my-index"
|
||||
|
||||
READ_ROUTES = [
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/stats"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs/$count"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs/seed-doc-1"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs/suggest"),
|
||||
("GET", f"/azure_ai/indexes/{INDEX}/docs/autocomplete"),
|
||||
("POST", f"/azure_ai/indexes/{INDEX}/docs/suggest"),
|
||||
("POST", f"/azure_ai/indexes/{INDEX}/docs/autocomplete"),
|
||||
("POST", f"/azure_ai/indexes/{INDEX}/analyze"),
|
||||
]
|
||||
|
||||
def _request(self, method: str, path: str) -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = method
|
||||
request.url.path = path
|
||||
return request
|
||||
|
||||
def _team_member(self, permissions: list) -> MagicMock:
|
||||
user = MagicMock(spec=UserAPIKeyAuth)
|
||||
user.user_role = None
|
||||
user.metadata = {"allowed_vector_store_indexes": [{"index_name": self.INDEX, "index_permissions": permissions}]}
|
||||
user.team_metadata = None
|
||||
return user
|
||||
|
||||
def test_team_with_write_grant_can_upload(self):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{self.INDEX}/docs/index"),
|
||||
user_api_key_dict=self._team_member(["read", "write"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_team_without_write_grant_cannot_upload(self):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{self.INDEX}/docs/index"),
|
||||
user_api_key_dict=self._team_member(["read"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
def test_team_with_read_grant_can_search(self):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{self.INDEX}/docs/search"),
|
||||
user_api_key_dict=self._team_member(["read"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_team_with_read_grant_can_get_index_details(self):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("GET", f"/azure_ai/indexes/{self.INDEX}"),
|
||||
user_api_key_dict=self._team_member(["read"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_team_without_read_grant_cannot_get_index_details(self):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request("GET", f"/azure_ai/indexes/{self.INDEX}"),
|
||||
user_api_key_dict=self._team_member(["write"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.parametrize("method, path", READ_ROUTES)
|
||||
def test_team_with_read_grant_can_call_every_read_route(self, method, path):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request(method, path),
|
||||
user_api_key_dict=self._team_member(["read"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
@pytest.mark.parametrize("method, path", READ_ROUTES)
|
||||
def test_team_without_read_grant_cannot_call_read_routes(self, method, path):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request(method, path),
|
||||
user_api_key_dict=self._team_member(["write"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"method, operation, path",
|
||||
[
|
||||
("PUT", "update", f"/azure_ai/indexes/{INDEX}?api-version=2024-07-01"),
|
||||
("DELETE", "delete", f"/azure_ai/indexes/{INDEX}?api-version=2024-07-01"),
|
||||
("POST", "create", "/azure_ai/indexes?api-version=2024-07-01"),
|
||||
],
|
||||
)
|
||||
def test_team_cannot_manage_index_lifecycle_even_with_write_grant(self, method, operation, path):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=self.INDEX,
|
||||
request=self._request(method, path),
|
||||
user_api_key_dict=self._team_member(["read", "write"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
assert f"Only proxy admins can {operation}" in exc_info.value.detail
|
||||
|
||||
|
||||
class TestAzureAIAnalyzeNamedIndexClassification:
|
||||
"""Regression tests for write-before-read endpoint classification.
|
||||
|
||||
The endpoint matcher is substring-based, so the batch-write path of an
|
||||
index named ``analyze*`` contains the ``("POST", "/analyze")`` read
|
||||
fragment. Reads-first classification labeled that write a read, letting a
|
||||
read-only grant upload, merge, and delete documents (and refusing
|
||||
legitimate write-only grants). Writes are classified first now, so an
|
||||
ambiguous path demands the stronger grant.
|
||||
"""
|
||||
|
||||
def _request(self, method: str, path: str) -> MagicMock:
|
||||
request = MagicMock(spec=Request)
|
||||
request.method = method
|
||||
request.url.path = path
|
||||
return request
|
||||
|
||||
def _team_member(self, index: str, permissions: list) -> MagicMock:
|
||||
user = MagicMock(spec=UserAPIKeyAuth)
|
||||
user.user_role = None
|
||||
user.metadata = {"allowed_vector_store_indexes": [{"index_name": index, "index_permissions": permissions}]}
|
||||
user.team_metadata = None
|
||||
return user
|
||||
|
||||
@pytest.mark.parametrize("index", ["analyze", "analyzer-reports"])
|
||||
def test_read_only_grant_cannot_upload_to_analyze_named_index(self, index):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=index,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{index}/docs/index"),
|
||||
user_api_key_dict=self._team_member(index, ["read"]),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.parametrize("index", ["analyze", "analyzer-reports"])
|
||||
def test_write_grant_can_upload_to_analyze_named_index(self, index):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name=index,
|
||||
request=self._request("POST", f"/azure_ai/indexes/{index}/docs/index"),
|
||||
user_api_key_dict=self._team_member(index, ["write"]),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_read_only_grant_can_still_analyze_on_analyze_named_index(self):
|
||||
result = is_allowed_to_call_vector_store_endpoint(
|
||||
provider=LlmProviders.AZURE_AI,
|
||||
index_name="analyze",
|
||||
request=self._request("POST", "/azure_ai/indexes/analyze/analyze"),
|
||||
user_api_key_dict=self._team_member("analyze", ["read"]),
|
||||
)
|
||||
assert result is True
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ def _setup_mcp_call_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
|
|||
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_module)
|
||||
|
||||
fake_manager = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=AsyncMock(return_value=_DummyMCPResult()),
|
||||
# Newer logging path calls this to enrich spend logs metadata
|
||||
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
|
||||
|
|
@ -373,6 +374,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
|
|||
post_call_failure_hook = _setup_proxy_logging(monkeypatch)
|
||||
|
||||
fake_manager = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom"))
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -464,6 +466,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch
|
|||
|
||||
# Patch manager methods used by _get_mcp_tools_from_manager to avoid needing full UserAPIKeyAuth fields.
|
||||
fake_manager = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
get_mcp_servers_from_ids=MagicMock(return_value=[]),
|
||||
get_mcp_server_by_name=MagicMock(return_value=None),
|
||||
|
|
@ -516,6 +519,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch):
|
|||
mock_get_tools,
|
||||
)
|
||||
fake_manager = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
get_allowed_mcp_servers=AsyncMock(return_value=[]),
|
||||
get_mcp_servers_from_ids=MagicMock(return_value=[]),
|
||||
get_mcp_server_by_name=MagicMock(return_value=None),
|
||||
|
|
|
|||
|
|
@ -69,6 +69,7 @@ def _mock_mcp_environment(monkeypatch: pytest.MonkeyPatch) -> AsyncMock:
|
|||
"""Patch the MCP tool-call plumbing so _execute_tool_calls can run in tests."""
|
||||
call_tool = AsyncMock(return_value=CallToolResult(content=[TextContent(type="text", text="ok")], isError=False))
|
||||
fake_manager = types.SimpleNamespace(
|
||||
get_registry=MagicMock(return_value={}),
|
||||
call_tool=call_tool,
|
||||
_get_mcp_server_from_tool_name=MagicMock(return_value=None),
|
||||
get_mcp_server_by_name=MagicMock(return_value=None),
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22941
|
||||
"limit": 22938
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 27139
|
||||
|
|
|
|||
BIN
ui/litellm-dashboard/public/assets/logos/nimble.png
Normal file
BIN
ui/litellm-dashboard/public/assets/logos/nimble.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 6.4 KiB |
|
|
@ -358,6 +358,7 @@ describe("ShadowEvalSection", () => {
|
|||
const expectedBody = {
|
||||
api_key_id: "hash-alpha",
|
||||
router_name: "gpt-auto",
|
||||
direction: "forward",
|
||||
shadow_percentage: 10,
|
||||
duration_days: 7,
|
||||
max_turns: 200,
|
||||
|
|
|
|||
|
|
@ -308,6 +308,7 @@ const StartForm: React.FC = () => {
|
|||
const startBody = {
|
||||
api_key_id: apiKeyId,
|
||||
router_name: routerName,
|
||||
direction: "forward" as const,
|
||||
shadow_percentage: parsedPct,
|
||||
duration_days: Number.parseInt(durationDays, 10),
|
||||
max_turns: parsedMaxTurns,
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import { AvailableSearchProvider, SearchTool } from "./types";
|
|||
import dataforseoLogo from "../../../../../public/assets/logos/dataforseo.png";
|
||||
import exaAiLogo from "../../../../../public/assets/logos/exa_ai.png";
|
||||
import googlePseLogo from "../../../../../public/assets/logos/google_pse.png";
|
||||
import nimbleLogo from "../../../../../public/assets/logos/nimble.png";
|
||||
import parallelAiLogo from "../../../../../public/assets/logos/parallel_ai.png";
|
||||
import perplexityLogo from "../../../../../public/assets/logos/perplexity.png";
|
||||
import tavilyLogo from "../../../../../public/assets/logos/tavily.png";
|
||||
|
|
@ -25,6 +26,7 @@ const searchProviderLogoMap: Record<string, string> = {
|
|||
exa_ai: exaAiLogo.src,
|
||||
google_pse: googlePseLogo.src,
|
||||
dataforseo: dataforseoLogo.src,
|
||||
nimble: nimbleLogo.src,
|
||||
};
|
||||
|
||||
interface SearchProviderLabelProps {
|
||||
|
|
|
|||
47
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
47
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -838,9 +838,15 @@ export interface paths {
|
|||
put?: never;
|
||||
/**
|
||||
* Start Shadow Eval
|
||||
* @description Start a pre-adoption shadow eval: duplicate a sampled slice of a key's live traffic
|
||||
* through an auto-router, judge real vs. shadow responses blind, and stratify win rates
|
||||
* by the router's tier classification and by the incumbent model.
|
||||
* @description Start a shadow eval: duplicate a sampled slice of a key's live traffic against a second
|
||||
* arm, judge the two responses blind, and stratify win rates by tier and by the model that
|
||||
* served the real arm.
|
||||
*
|
||||
* A forward job answers whether the key should adopt router_name: it samples the requests
|
||||
* the router did not serve and duplicates them through it. A reverse job answers whether a
|
||||
* key already on the router still gains from it: it samples the requests the router did
|
||||
* serve and duplicates them against baseline_model. A key can hold one active job per
|
||||
* direction, so both questions can run at once.
|
||||
*
|
||||
* Shadow responses are never served to users. The job samples until it has judged
|
||||
* max_turns turns, reaches the end of its window, or is stopped; sampling changes
|
||||
|
|
@ -32737,11 +32743,19 @@ export interface components {
|
|||
* @description The hashed virtual key whose traffic this job evaluates, and only that key's
|
||||
*/
|
||||
api_key_id: string;
|
||||
/** Baseline Model */
|
||||
baseline_model?: string | null;
|
||||
/**
|
||||
* Created At
|
||||
* Format: date-time
|
||||
*/
|
||||
created_at: string;
|
||||
/**
|
||||
* Direction
|
||||
* @default forward
|
||||
* @enum {string}
|
||||
*/
|
||||
direction: "forward" | "reverse";
|
||||
/**
|
||||
* Ends At
|
||||
* Format: date-time
|
||||
|
|
@ -32794,7 +32808,10 @@ export interface components {
|
|||
* @description Stratified results of a shadow-eval job's verdicts so far.
|
||||
*/
|
||||
ShadowEvalResult: {
|
||||
/** By Current Model */
|
||||
/**
|
||||
* By Current Model
|
||||
* @description Sliced by the model that served the real arm: the key's incumbent models in forward mode, and in reverse the models the router itself picked
|
||||
*/
|
||||
by_current_model: components["schemas"]["ShadowEvalSlice"][];
|
||||
/** By Tier */
|
||||
by_tier: components["schemas"]["ShadowEvalSlice"][];
|
||||
|
|
@ -32806,7 +32823,7 @@ export interface components {
|
|||
/**
|
||||
* ShadowEvalSlice
|
||||
* @description Judge outcomes for one slice of a job's verdicts (a router tier, or one of the
|
||||
* models the shadowed key currently uses).
|
||||
* models that served the real arm).
|
||||
*/
|
||||
ShadowEvalSlice: {
|
||||
/** Avg Judge Confidence */
|
||||
|
|
@ -32815,12 +32832,12 @@ export interface components {
|
|||
group: string;
|
||||
/**
|
||||
* Real Win Rate Pct
|
||||
* @description Share of judged turns where the real (control) model won
|
||||
* @description Share of judged turns the real arm won, meaning the response the caller actually received: the key's own model in forward mode, the router's pick in reverse
|
||||
*/
|
||||
real_win_rate_pct: number;
|
||||
/**
|
||||
* Shadow Win Rate Pct
|
||||
* @description Share of judged turns where the shadowed router's pick won
|
||||
* @description Share of judged turns the shadow arm won, meaning the duplicated response nobody was served: the router's pick in forward mode, baseline_model in reverse
|
||||
*/
|
||||
shadow_win_rate_pct: number;
|
||||
/** Tie Rate Pct */
|
||||
|
|
@ -33003,7 +33020,7 @@ export interface components {
|
|||
};
|
||||
/**
|
||||
* StartShadowEvalRequest
|
||||
* @description Start shadowing a key's traffic through an auto-router for blind comparison.
|
||||
* @description Start duplicating a key's traffic for blind comparison against an auto-router.
|
||||
*/
|
||||
StartShadowEvalRequest: {
|
||||
/**
|
||||
|
|
@ -33011,6 +33028,18 @@ export interface components {
|
|||
* @description The hashed virtual key whose traffic will be shadowed. Shadow evaluation runs ONLY on this key's traffic; requests made with any other key are not sampled.
|
||||
*/
|
||||
api_key_id: string;
|
||||
/**
|
||||
* Baseline Model
|
||||
* @description Required when direction is reverse and rejected otherwise: the fixed model the router's own responses are judged against. Must be a plain model rather than another auto-router
|
||||
*/
|
||||
baseline_model?: string | null;
|
||||
/**
|
||||
* Direction
|
||||
* @description forward answers 'should this key adopt router_name': it samples the requests the key did NOT route through the router and duplicates them through it. reverse answers 'is the router still worth it for a key already on it': it samples the requests the router did serve and duplicates them against baseline_model. The response the caller received is always the real arm
|
||||
* @default forward
|
||||
* @enum {string}
|
||||
*/
|
||||
direction: "forward" | "reverse";
|
||||
/**
|
||||
* Duration Days
|
||||
* @description How many days the job samples traffic before completing on its own
|
||||
|
|
@ -33031,7 +33060,7 @@ export interface components {
|
|||
max_turns: number;
|
||||
/**
|
||||
* Router Name
|
||||
* @description The auto-router config to shadow requests through
|
||||
* @description The auto-router under evaluation, in either direction
|
||||
*/
|
||||
router_name: string;
|
||||
/**
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue