mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_fix_agent_mcp_grants
# Conflicts: # type-discipline-budget.json
This commit is contained in:
commit
45988143bc
80 changed files with 3886 additions and 1200 deletions
|
|
@ -54,7 +54,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5599
|
||||
"limit": 5597
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15290
|
||||
|
|
@ -99,7 +99,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44358
|
||||
"limit": 44355
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
|
|
@ -108,10 +108,10 @@
|
|||
"limit": 38332
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19623
|
||||
"limit": 19621
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 29858
|
||||
"limit": 29855
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 111
|
||||
|
|
|
|||
|
|
@ -424,6 +424,10 @@ anthropic_beta_headers_url: str = os.getenv(
|
|||
"LITELLM_ANTHROPIC_BETA_HEADERS_URL",
|
||||
"https://raw.githubusercontent.com/BerriAI/litellm/main/litellm/anthropic_beta_headers_config.json",
|
||||
)
|
||||
autorouter_presets_url: str = os.getenv(
|
||||
"LITELLM_AUTOROUTER_PRESETS_URL",
|
||||
"https://raw.githubusercontent.com/BerriAI/litellm/main/litellm/proxy/public_endpoints/autorouter_presets.json",
|
||||
)
|
||||
suppress_debug_info: bool = False
|
||||
dynamodb_table_name: Optional[str] = None
|
||||
s3_callback_params: Optional[Dict] = None
|
||||
|
|
|
|||
|
|
@ -17,7 +17,11 @@ from litellm.constants import (
|
|||
from litellm.litellm_core_utils.env_utils import get_env_int
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.litellm_core_utils.secret_redaction import redact_string, redact_structured_value
|
||||
from litellm.litellm_core_utils.secret_redaction import (
|
||||
redact_internal_details,
|
||||
redact_string,
|
||||
redact_structured_value,
|
||||
)
|
||||
|
||||
set_verbose = False
|
||||
|
||||
|
|
@ -89,6 +93,14 @@ def redact_secrets(value: str) -> str:
|
|||
return _redact_string(value)
|
||||
|
||||
|
||||
def redact_internal_details_from_client_message(value: str) -> str:
|
||||
"""Public API: redact_secrets() plus filesystem paths, internal hostnames, and an
|
||||
embedded traceback, for a string about to leave the process in an HTTP response."""
|
||||
if not _ENABLE_SECRET_REDACTION:
|
||||
return value
|
||||
return redact_internal_details(value)
|
||||
|
||||
|
||||
def _substituted_color_message(record: logging.LogRecord) -> str | None:
|
||||
"""Render a record's ``color_message`` against its args, or None if absent.
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ This hook is called before making an LLM request when a vector store is configur
|
|||
It searches the vector store for relevant context and appends it to the messages.
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -80,10 +81,17 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
|
||||
# Get prisma_client for database fallback
|
||||
prisma_client = None
|
||||
llm_router = None
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client as _prisma_client
|
||||
from litellm.proxy.proxy_server import (
|
||||
llm_router as _llm_router,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client as _prisma_client,
|
||||
)
|
||||
|
||||
prisma_client = _prisma_client
|
||||
llm_router = _llm_router
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
|
@ -114,12 +122,26 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
vector_store_id = vector_store_to_run.get("vector_store_id", "")
|
||||
custom_llm_provider = vector_store_to_run.get("custom_llm_provider")
|
||||
litellm_params_for_vector_store = vector_store_to_run.get("litellm_params", {}) or {}
|
||||
# Call litellm.vector_stores.search() with the required parameters
|
||||
search_response = await litellm.vector_stores.asearch(
|
||||
request_litellm_params = litellm_logging_obj.model_call_details.get("litellm_params", {})
|
||||
request_metadata = (
|
||||
request_litellm_params.get("metadata", {}) if isinstance(request_litellm_params, dict) else {}
|
||||
)
|
||||
if llm_router is not None:
|
||||
search_function = cast( # cast-ok: normalize router search callable
|
||||
Callable[..., Awaitable[VectorStoreSearchResponse]],
|
||||
llm_router.avector_store_search,
|
||||
)
|
||||
else:
|
||||
search_function = cast( # cast-ok: normalize SDK search callable
|
||||
Callable[..., Awaitable[VectorStoreSearchResponse]],
|
||||
litellm.vector_stores.asearch,
|
||||
)
|
||||
search_response = await search_function(
|
||||
**{
|
||||
"vector_store_id": vector_store_id,
|
||||
"query": query,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"metadata": request_metadata,
|
||||
**litellm_params_for_vector_store,
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -92,6 +92,27 @@ def redact_string(value: str) -> str:
|
|||
return _SECRET_RE.sub(_REDACTED, value)
|
||||
|
||||
|
||||
_UNIX_SYSTEM_PATH: Final = r"/(?:etc|var|opt|usr|home|root|private|Users|tmp|mnt|srv)/[^\s'\"\)\]}>,]+"
|
||||
_WINDOWS_DRIVE_PATH: Final = r"[A-Za-z]:\\[^\s'\"\)\]}>,]+"
|
||||
_PRIVATE_OR_LOOPBACK_IPV4: Final = (
|
||||
r"\b(?:10(?:\.\d{1,3}){3}|172\.(?:1[6-9]|2\d|3[01])(?:\.\d{1,3}){2}|192\.168(?:\.\d{1,3}){2}|127(?:\.\d{1,3}){3})\b"
|
||||
)
|
||||
_INTERNAL_SUFFIX_HOSTNAME: Final = r"\b[A-Za-z0-9-]+(?:\.[A-Za-z0-9-]+)*\.(?:internal|local|corp|lan|intra|private)\b"
|
||||
_INTERNAL_DETAIL_RE: Final = re.compile(
|
||||
"|".join((_UNIX_SYSTEM_PATH, _WINDOWS_DRIVE_PATH, _PRIVATE_OR_LOOPBACK_IPV4, _INTERNAL_SUFFIX_HOSTNAME)),
|
||||
re.IGNORECASE,
|
||||
)
|
||||
_TRACEBACK_MARKER: Final = "Traceback (most recent call last):"
|
||||
|
||||
|
||||
def redact_internal_details(value: str) -> str:
|
||||
"""Drop an embedded traceback and scrub filesystem paths and internal hostnames,
|
||||
on top of redact_string(). For client-facing messages only: server logs keep this detail."""
|
||||
marker_index: Final = value.find(_TRACEBACK_MARKER)
|
||||
without_traceback: Final = value[:marker_index].rstrip() if marker_index != -1 else value
|
||||
return _INTERNAL_DETAIL_RE.sub(_REDACTED, redact_string(without_traceback))
|
||||
|
||||
|
||||
def redact_structured_value(key: str | None, value: str) -> str:
|
||||
"""Scrub *value* as it appeared under *key* inside a structured record.
|
||||
|
||||
|
|
|
|||
|
|
@ -1,10 +1,15 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure.common_utils import BaseAzureLLM
|
||||
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseQueryEmbeddingVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.vector_stores import (
|
||||
BaseVectorStoreAuthCredentials,
|
||||
|
|
@ -26,7 +31,7 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM):
|
||||
class AzureAIVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig, BaseAzureLLM):
|
||||
"""
|
||||
Configuration for Azure AI Search Vector Store
|
||||
|
||||
|
|
@ -110,83 +115,73 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM):
|
|||
def transform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | list[str],
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
router: "Router | None" = None,
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
"""
|
||||
Transform search request for Azure AI Search API
|
||||
litellm_params: Mapping[str, object],
|
||||
extra_body: Mapping[str, object] | None = None,
|
||||
router: Router | None = None,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
query_text: Final = self.query_text(query)
|
||||
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor, router)
|
||||
return self._search_request(
|
||||
vector_store_id,
|
||||
query_text,
|
||||
query_vector,
|
||||
vector_store_search_optional_params,
|
||||
api_base,
|
||||
litellm_logging_obj,
|
||||
litellm_params,
|
||||
)
|
||||
|
||||
Generates embeddings using litellm.embeddings and constructs Azure AI Search request
|
||||
"""
|
||||
# Convert query to string if it's a list
|
||||
if isinstance(query, list):
|
||||
query = " ".join(query)
|
||||
async def atransform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
extra_body: Mapping[str, object] | None = None,
|
||||
router: Router | None = None,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
query_text: Final = self.query_text(query)
|
||||
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor, router)
|
||||
return self._search_request(
|
||||
vector_store_id,
|
||||
query_text,
|
||||
query_vector,
|
||||
vector_store_search_optional_params,
|
||||
api_base,
|
||||
litellm_logging_obj,
|
||||
litellm_params,
|
||||
)
|
||||
|
||||
# Get embedding model from litellm_params (required)
|
||||
embedding_model: Final = litellm_params.get("litellm_embedding_model")
|
||||
if not embedding_model:
|
||||
raise ValueError(
|
||||
"embedding_model is required in litellm_params for Azure AI Search. "
|
||||
"Example: litellm_params['embedding_model'] = 'azure/text-embedding-3-large'"
|
||||
)
|
||||
|
||||
embedding_config: Final = litellm_params.get("litellm_embedding_config", {})
|
||||
if not embedding_config:
|
||||
raise ValueError(
|
||||
"embedding_config is required in litellm_params for Azure AI Search. "
|
||||
"Example: litellm_params['embedding_config'] = {'api_base': 'https://krris-mh44uf7y-eastus2.cognitiveservices.azure.com/', 'api_key': 'os.environ/AZURE_API_KEY', 'api_version': '2025-09-01'}"
|
||||
)
|
||||
|
||||
# Get vector field name (defaults to contentVector)
|
||||
@staticmethod
|
||||
def _search_request(
|
||||
vector_store_id: str,
|
||||
query_text: str,
|
||||
query_vector: Sequence[float],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
vector_field: Final = litellm_params.get("azure_search_vector_field", "contentVector")
|
||||
|
||||
# Get top_k (number of results to return)
|
||||
top_k: Final = vector_store_search_optional_params.get("top_k", 10)
|
||||
|
||||
# Generate embedding for the query using litellm.embeddings
|
||||
try:
|
||||
embedding_response: Final = litellm.embedding(
|
||||
model=embedding_model,
|
||||
input=[query],
|
||||
**embedding_config,
|
||||
)
|
||||
query_vector: Final = embedding_response.data[0]["embedding"]
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to generate embedding for query: {e}")
|
||||
|
||||
# Azure AI Search endpoint for search
|
||||
index_name: Final = vector_store_id # vector_store_id is the index name
|
||||
url: Final = f"{api_base}/indexes/{index_name}/docs/search?api-version=2024-07-01"
|
||||
|
||||
# Build the request body for Azure AI Search with vector search
|
||||
request_body: Final = {
|
||||
"search": "*", # Get all documents (filtered by vector similarity)
|
||||
"vectorQueries": [
|
||||
{
|
||||
"vector": query_vector,
|
||||
"fields": vector_field,
|
||||
"kind": "vector",
|
||||
"k": top_k, # Number of nearest neighbors to return
|
||||
}
|
||||
],
|
||||
"select": "id,content", # Fields to return (customize based on schema)
|
||||
litellm_logging_obj.model_call_details["input"] = query_text
|
||||
litellm_logging_obj.model_call_details["embedding_model"] = litellm_params.get("litellm_embedding_model")
|
||||
litellm_logging_obj.model_call_details["top_k"] = top_k
|
||||
return f"{api_base}/indexes/{vector_store_id}/docs/search?api-version=2024-07-01", {
|
||||
"search": "*",
|
||||
"vectorQueries": [{"vector": query_vector, "fields": vector_field, "kind": "vector", "k": top_k}],
|
||||
"select": "id,content",
|
||||
"top": top_k,
|
||||
}
|
||||
|
||||
#########################################################
|
||||
# Update logging object with details of the request
|
||||
#########################################################
|
||||
litellm_logging_obj.model_call_details["input"] = query
|
||||
litellm_logging_obj.model_call_details["embedding_model"] = embedding_model
|
||||
litellm_logging_obj.model_call_details["top_k"] = top_k
|
||||
|
||||
return url, request_body
|
||||
|
||||
def transform_search_vector_store_response(
|
||||
self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj
|
||||
) -> VectorStoreSearchResponse:
|
||||
|
|
|
|||
|
|
@ -1,10 +1,16 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from abc import abstractmethod
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, NoReturn
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol, runtime_checkable
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.types.vector_stores import (
|
||||
VECTOR_STORE_OPENAI_PARAMS,
|
||||
BaseVectorStoreAuthCredentials,
|
||||
|
|
@ -28,6 +34,95 @@ else:
|
|||
BaseLLMException = Any
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class VectorStoreEmbeddingExecutor(Protocol):
|
||||
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse: ...
|
||||
|
||||
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LiteLLMVectorStoreEmbeddingExecutor:
|
||||
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
|
||||
import litellm
|
||||
|
||||
return litellm.embedding( # pyright: ignore[reportCallIssue, reportUnknownMemberType, reportUnknownVariableType] # provider kwargs are intentionally dynamic
|
||||
model=model,
|
||||
input=[query], # mutable-ok: LiteLLM embedding requires a mutable input list
|
||||
**dict(configuration), # pyright: ignore[reportArgumentType] # provider-specific embedding config is validated downstream # mutable-ok: kwargs require a concrete dict
|
||||
)
|
||||
|
||||
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
|
||||
import litellm
|
||||
|
||||
return await litellm.aembedding( # pyright: ignore[reportUnknownMemberType] # provider kwargs are intentionally dynamic
|
||||
model=model,
|
||||
input=[query], # mutable-ok: LiteLLM embedding requires a mutable input list
|
||||
**dict(configuration), # pyright: ignore[reportArgumentType] # provider-specific embedding config is validated downstream # mutable-ok: kwargs require a concrete dict
|
||||
)
|
||||
|
||||
|
||||
_REQUEST_METADATA: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def vector_store_request_metadata(kwargs: Mapping[str, object]) -> Mapping[str, object]:
|
||||
litellm_metadata: Final = kwargs.get("litellm_metadata")
|
||||
if isinstance(litellm_metadata, dict):
|
||||
return _REQUEST_METADATA.validate_python(litellm_metadata)
|
||||
metadata: Final = kwargs.get("metadata")
|
||||
if isinstance(metadata, dict):
|
||||
return _REQUEST_METADATA.validate_python(metadata)
|
||||
return MappingProxyType({})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RouterVectorStoreEmbeddingExecutor:
|
||||
router: Router
|
||||
metadata: Mapping[str, object]
|
||||
|
||||
def _embedding_kwargs(self, configuration: Mapping[str, object]) -> Mapping[str, object]:
|
||||
configured_metadata: Final = configuration.get("metadata")
|
||||
metadata: Final = {
|
||||
**(configured_metadata if isinstance(configured_metadata, Mapping) else {}),
|
||||
**self.metadata,
|
||||
}
|
||||
return {
|
||||
**{key: value for key, value in configuration.items() if key not in ("input", "metadata", "model")},
|
||||
"metadata": metadata,
|
||||
}
|
||||
|
||||
def _router_serves(self, model: str) -> bool:
|
||||
team_id: Final = self.metadata.get("user_api_key_team_id")
|
||||
resolved: Final = self.router.resolved_litellm_models(model, team_id if isinstance(team_id, str) else None)
|
||||
deployment_models: Final = (
|
||||
deployment.get("litellm_params", {}).get("model") for deployment in self.router.get_model_list() or ()
|
||||
)
|
||||
return bool(resolved) or model in deployment_models
|
||||
|
||||
def _embeds_through_sdk(self, model: str, configuration: Mapping[str, object]) -> bool:
|
||||
return bool(configuration) and not self._router_serves(model)
|
||||
|
||||
def embed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
|
||||
embedding_kwargs: Final = self._embedding_kwargs(configuration)
|
||||
if self._embeds_through_sdk(model, configuration):
|
||||
return LiteLLMVectorStoreEmbeddingExecutor().embed(model, query, embedding_kwargs)
|
||||
return self.router.embedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list
|
||||
model=model,
|
||||
input=[query], # mutable-ok: Router embedding requires a mutable input list
|
||||
**embedding_kwargs, # pyright: ignore[reportArgumentType] # provider kwargs are intentionally dynamic
|
||||
)
|
||||
|
||||
async def aembed(self, model: str, query: str, configuration: Mapping[str, object]) -> EmbeddingResponse:
|
||||
embedding_kwargs: Final = self._embedding_kwargs(configuration)
|
||||
if self._embeds_through_sdk(model, configuration):
|
||||
return await LiteLLMVectorStoreEmbeddingExecutor().aembed(model, query, embedding_kwargs)
|
||||
return await self.router.aembedding( # pyright: ignore[reportUnknownMemberType] # Router embedding input retains a legacy untyped list
|
||||
model=model,
|
||||
input=[query], # mutable-ok: Router embedding requires a mutable input list
|
||||
**embedding_kwargs, # pyright: ignore[reportArgumentType] # provider kwargs are intentionally dynamic
|
||||
)
|
||||
|
||||
|
||||
class BaseVectorStoreConfig:
|
||||
def get_supported_openai_params(self, model: str) -> list[VECTOR_STORE_OPENAI_PARAMS]:
|
||||
return []
|
||||
|
|
@ -58,7 +153,7 @@ class BaseVectorStoreConfig:
|
|||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
router: "Router | None" = None,
|
||||
router: Router | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
pass
|
||||
|
||||
|
|
@ -71,7 +166,7 @@ class BaseVectorStoreConfig:
|
|||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
router: "Router | None" = None,
|
||||
router: Router | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""
|
||||
Optional async version of transform_search_vector_store_request.
|
||||
|
|
@ -161,6 +256,116 @@ class BaseVectorStoreConfig:
|
|||
return 0.0, 0.0
|
||||
|
||||
|
||||
_EMPTY_EMBEDDING_CONFIGURATION: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
_QUERY_VECTOR: Final = TypeAdapter(list[float])
|
||||
|
||||
|
||||
class BaseQueryEmbeddingVectorStoreConfig(BaseVectorStoreConfig):
|
||||
@abstractmethod
|
||||
def transform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
extra_body: Mapping[str, object] | None = None,
|
||||
router: Router | None = None,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
pass
|
||||
|
||||
async def atransform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
extra_body: Mapping[str, object] | None = None,
|
||||
router: Router | None = None,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
return self.transform_search_vector_store_request(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
litellm_params=litellm_params,
|
||||
extra_body=extra_body,
|
||||
router=router,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def query_text(query: str | Sequence[str]) -> str:
|
||||
return query if isinstance(query, str) else " ".join(query)
|
||||
|
||||
@staticmethod
|
||||
def query_embedding_model(litellm_params: Mapping[str, object]) -> str:
|
||||
embedding_model: Final = litellm_params.get("litellm_embedding_model")
|
||||
if isinstance(embedding_model, str) and embedding_model:
|
||||
return embedding_model
|
||||
raise ValueError(
|
||||
"litellm_embedding_model is required in litellm_params for this vector store. "
|
||||
"Example: litellm_params['litellm_embedding_model'] = 'openai/text-embedding-3-small'"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def query_embedding_configuration(litellm_params: Mapping[str, object]) -> Mapping[str, object]:
|
||||
configuration: Final = litellm_params.get("litellm_embedding_config")
|
||||
if isinstance(configuration, Mapping):
|
||||
return {str(key): value for key, value in configuration.items()} # pyright: ignore[reportUnknownVariableType, reportUnknownArgumentType] # litellm_params is an untyped dict, keys are re-validated as str here
|
||||
return _EMPTY_EMBEDDING_CONFIGURATION
|
||||
|
||||
@staticmethod
|
||||
def query_embedding_executor(
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None,
|
||||
router: Router | None,
|
||||
request_metadata: Mapping[str, object] = MappingProxyType({}),
|
||||
) -> VectorStoreEmbeddingExecutor:
|
||||
if embedding_executor is not None:
|
||||
return embedding_executor
|
||||
if router is not None:
|
||||
return RouterVectorStoreEmbeddingExecutor(router=router, metadata=request_metadata)
|
||||
return LiteLLMVectorStoreEmbeddingExecutor()
|
||||
|
||||
def embed_query(
|
||||
self,
|
||||
query_text: str,
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None,
|
||||
router: Router | None = None,
|
||||
) -> Sequence[float]:
|
||||
model: Final = self.query_embedding_model(litellm_params)
|
||||
configuration: Final = self.query_embedding_configuration(litellm_params)
|
||||
executor: Final = self.query_embedding_executor(embedding_executor, router)
|
||||
try:
|
||||
response: Final = executor.embed(model, query_text, configuration)
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to generate embedding for query: {e}")
|
||||
return _QUERY_VECTOR.validate_python(response.data[0]["embedding"]) # pyright: ignore[reportUnknownMemberType] # EmbeddingResponse.data is an untyped list, the vector is validated here
|
||||
|
||||
async def aembed_query(
|
||||
self,
|
||||
query_text: str,
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None,
|
||||
router: Router | None = None,
|
||||
) -> Sequence[float]:
|
||||
model: Final = self.query_embedding_model(litellm_params)
|
||||
configuration: Final = self.query_embedding_configuration(litellm_params)
|
||||
executor: Final = self.query_embedding_executor(embedding_executor, router)
|
||||
try:
|
||||
response: Final = await executor.aembed(model, query_text, configuration)
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to generate embedding for query: {e}")
|
||||
return _QUERY_VECTOR.validate_python(response.data[0]["embedding"]) # pyright: ignore[reportUnknownMemberType] # EmbeddingResponse.data is an untyped list, the vector is validated here
|
||||
|
||||
|
||||
class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
|
||||
"""
|
||||
Base config for vector store providers whose datastore has no HTTP API
|
||||
|
|
@ -176,6 +381,7 @@ class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
|
|||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
pass
|
||||
|
|
@ -188,6 +394,7 @@ class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
|
|||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
pass
|
||||
|
|
@ -201,7 +408,7 @@ class BaseDirectVectorStoreConfig(BaseVectorStoreConfig):
|
|||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
extra_body: Mapping[str, object] | None = None,
|
||||
router: "Router | None" = None,
|
||||
router: Router | None = None,
|
||||
) -> NoReturn:
|
||||
raise NotImplementedError("Direct vector store providers execute the search themselves; no HTTP request shape")
|
||||
|
||||
|
|
|
|||
|
|
@ -9,13 +9,15 @@ import threading
|
|||
import time
|
||||
from collections.abc import AsyncIterable, Callable, Iterable, Mapping
|
||||
from http.cookiejar import CookieJar, DefaultCookiePolicy
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, Optional, TypeAlias, TypedDict
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, Optional, TypeAlias, TypedDict, TypeVar
|
||||
|
||||
import certifi
|
||||
import httpx
|
||||
from aiohttp import ClientSession, DummyCookieJar, TCPConnector
|
||||
from httpx import USE_CLIENT_DEFAULT, AsyncHTTPTransport, HTTPTransport
|
||||
from httpx._types import RequestFiles
|
||||
from httpx._types import CertTypes, RequestFiles
|
||||
from httpx._utils import get_environment_proxies
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -66,6 +68,22 @@ _AddrInfo: TypeAlias = tuple[int | socket.AddressFamily, int | socket.SocketKind
|
|||
|
||||
_RequestContent: TypeAlias = str | bytes | Iterable[bytes] | AsyncIterable[bytes]
|
||||
|
||||
_IPV4_LOCAL_ADDRESS: Final = "0.0.0.0"
|
||||
|
||||
_HttpxTransportT = TypeVar("_HttpxTransportT", HTTPTransport, AsyncHTTPTransport)
|
||||
|
||||
|
||||
def _environment_proxy_mounts(
|
||||
build_proxy_transport: Callable[[str], _HttpxTransportT],
|
||||
) -> Mapping[str, _HttpxTransportT | None]:
|
||||
"""httpx skips its own HTTP(S)_PROXY / NO_PROXY mounts whenever an explicit `transport=` is passed."""
|
||||
return MappingProxyType(
|
||||
{
|
||||
pattern: None if proxy_url is None else build_proxy_transport(proxy_url)
|
||||
for pattern, proxy_url in get_environment_proxies().items()
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class _TCPConnectorKwargs(TypedDict, total=False):
|
||||
local_addr: tuple[str, int] | None
|
||||
|
|
@ -607,6 +625,7 @@ class AsyncHTTPHandler:
|
|||
|
||||
return httpx.AsyncClient(
|
||||
transport=transport,
|
||||
mounts=AsyncHTTPHandler._create_httpx_proxy_mounts(transport, verify=ssl_config, cert=cert),
|
||||
event_hooks=event_hooks,
|
||||
timeout=timeout,
|
||||
verify=ssl_config,
|
||||
|
|
@ -1191,10 +1210,22 @@ class AsyncHTTPHandler:
|
|||
- [Default] If force_ipv4 is False, it will return None
|
||||
"""
|
||||
if litellm.force_ipv4:
|
||||
return AsyncHTTPTransport(local_address="0.0.0.0")
|
||||
return AsyncHTTPTransport(local_address=_IPV4_LOCAL_ADDRESS)
|
||||
else:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _create_httpx_proxy_mounts(
|
||||
transport: LiteLLMAiohttpTransport | AsyncHTTPTransport | None,
|
||||
verify: VerifyTypes,
|
||||
cert: CertTypes | None,
|
||||
) -> Mapping[str, AsyncHTTPTransport | None] | None:
|
||||
if not isinstance(transport, AsyncHTTPTransport):
|
||||
return None
|
||||
return _environment_proxy_mounts(
|
||||
lambda proxy_url: AsyncHTTPTransport(proxy=proxy_url, verify=verify, cert=cert)
|
||||
)
|
||||
|
||||
|
||||
class HTTPHandler:
|
||||
def __init__(
|
||||
|
|
@ -1227,6 +1258,7 @@ class HTTPHandler:
|
|||
# Create a client with a connection pool
|
||||
return httpx.Client(
|
||||
transport=self._create_sync_transport(),
|
||||
mounts=self._create_sync_proxy_mounts(verify=ssl_config, cert=cert),
|
||||
timeout=self.timeout if self.timeout is not None else _DEFAULT_TIMEOUT,
|
||||
verify=ssl_config,
|
||||
cert=cert,
|
||||
|
|
@ -1507,10 +1539,19 @@ class HTTPHandler:
|
|||
Some users have seen httpx ConnectionError when using ipv6 - forcing ipv4 resolves the issue for them
|
||||
"""
|
||||
if litellm.force_ipv4:
|
||||
return HTTPTransport(local_address="0.0.0.0")
|
||||
return HTTPTransport(local_address=_IPV4_LOCAL_ADDRESS)
|
||||
else:
|
||||
return getattr(litellm, "sync_transport", None)
|
||||
|
||||
@staticmethod
|
||||
def _create_sync_proxy_mounts(
|
||||
verify: VerifyTypes,
|
||||
cert: CertTypes | None,
|
||||
) -> Mapping[str, HTTPTransport | None] | None:
|
||||
if not litellm.force_ipv4:
|
||||
return None
|
||||
return _environment_proxy_mounts(lambda proxy_url: HTTPTransport(proxy=proxy_url, verify=verify, cert=cert))
|
||||
|
||||
|
||||
def get_async_httpx_client(
|
||||
llm_provider: LlmProviders | httpxSpecialProvider,
|
||||
|
|
|
|||
|
|
@ -69,7 +69,9 @@ from litellm.llms.base_llm.skills.transformation import BaseSkillsAPIConfig
|
|||
from litellm.llms.base_llm.text_to_speech.transformation import BaseTextToSpeechConfig
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseDirectVectorStoreConfig,
|
||||
BaseQueryEmbeddingVectorStoreConfig,
|
||||
BaseVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.llms.base_llm.vector_store_files.transformation import (
|
||||
BaseVectorStoreFilesConfig,
|
||||
|
|
@ -9701,6 +9703,7 @@ class BaseLLMHTTPHandler:
|
|||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
|
|
@ -9721,6 +9724,7 @@ class BaseLLMHTTPHandler:
|
|||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params), # mutable-ok: snapshot GenericLiteLLMParams into the Mapping shape
|
||||
embedding_executor=embedding_executor,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
|
|
@ -9744,8 +9748,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
# Check if provider has async transform method
|
||||
if hasattr(vector_store_provider_config, "atransform_search_vector_store_request"):
|
||||
if isinstance(vector_store_provider_config, BaseQueryEmbeddingVectorStoreConfig):
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
|
|
@ -9758,12 +9761,13 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
router=router,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
else:
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
) = vector_store_provider_config.transform_search_vector_store_request(
|
||||
) = await vector_store_provider_config.atransform_search_vector_store_request(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
|
|
@ -9818,6 +9822,7 @@ class BaseLLMHTTPHandler:
|
|||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_body: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
|
|
@ -9834,6 +9839,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
embedding_executor=embedding_executor,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
|
|
@ -9854,6 +9860,7 @@ class BaseLLMHTTPHandler:
|
|||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params), # mutable-ok: snapshot GenericLiteLLMParams into the Mapping shape
|
||||
embedding_executor=embedding_executor,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
|
|
@ -9874,19 +9881,35 @@ class BaseLLMHTTPHandler:
|
|||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
) = vector_store_provider_config.transform_search_vector_store_request(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
router=router,
|
||||
)
|
||||
if isinstance(vector_store_provider_config, BaseQueryEmbeddingVectorStoreConfig):
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
) = vector_store_provider_config.transform_search_vector_store_request(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
router=router,
|
||||
embedding_executor=embedding_executor,
|
||||
)
|
||||
else:
|
||||
(
|
||||
url,
|
||||
request_body,
|
||||
) = vector_store_provider_config.transform_search_vector_store_request(
|
||||
vector_store_id=vector_store_id,
|
||||
query=query,
|
||||
vector_store_search_optional_params=vector_store_search_optional_params,
|
||||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
router=router,
|
||||
)
|
||||
|
||||
all_optional_params: Final[dict[str, object]] = dict(litellm_params)
|
||||
all_optional_params.update(vector_store_search_optional_params or {})
|
||||
|
|
|
|||
|
|
@ -1,9 +1,14 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseQueryEmbeddingVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.vector_stores import (
|
||||
|
|
@ -37,7 +42,7 @@ MILVUS_OPTIONAL_PARAMS: Final = {
|
|||
}
|
||||
|
||||
|
||||
class MilvusVectorStoreConfig(BaseVectorStoreConfig):
|
||||
class MilvusVectorStoreConfig(BaseQueryEmbeddingVectorStoreConfig):
|
||||
"""
|
||||
Configuration for Milvus Vector Store
|
||||
|
||||
|
|
@ -118,78 +123,79 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig):
|
|||
def transform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | list[str],
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: dict[str, Any] | None = None,
|
||||
router: "Router | None" = None,
|
||||
) -> tuple[str, dict[str, Any]]:
|
||||
"""
|
||||
Transform search request for Azure AI Search API
|
||||
litellm_params: Mapping[str, object],
|
||||
extra_body: Mapping[str, object] | None = None,
|
||||
router: Router | None = None,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
query_text: Final = self.query_text(query)
|
||||
query_vector: Final = self.embed_query(query_text, litellm_params, embedding_executor, router)
|
||||
return self._search_request(
|
||||
vector_store_id,
|
||||
query_text,
|
||||
query_vector,
|
||||
vector_store_search_optional_params,
|
||||
api_base,
|
||||
litellm_logging_obj,
|
||||
litellm_params,
|
||||
)
|
||||
|
||||
Generates embeddings using litellm.embeddings and constructs Azure AI Search request
|
||||
"""
|
||||
# Convert query to string if it's a list
|
||||
if isinstance(query, list):
|
||||
query = " ".join(query)
|
||||
async def atransform_search_vector_store_request(
|
||||
self,
|
||||
vector_store_id: str,
|
||||
query: str | Sequence[str],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
extra_body: Mapping[str, object] | None = None,
|
||||
router: Router | None = None,
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
query_text: Final = self.query_text(query)
|
||||
query_vector: Final = await self.aembed_query(query_text, litellm_params, embedding_executor, router)
|
||||
return self._search_request(
|
||||
vector_store_id,
|
||||
query_text,
|
||||
query_vector,
|
||||
vector_store_search_optional_params,
|
||||
api_base,
|
||||
litellm_logging_obj,
|
||||
litellm_params,
|
||||
)
|
||||
|
||||
# Get embedding model from litellm_params (required)
|
||||
embedding_model: Final = litellm_params.get("litellm_embedding_model")
|
||||
if not embedding_model:
|
||||
raise ValueError(
|
||||
"embedding_model is required in litellm_params for Milvus. You can call any litellm embedding model."
|
||||
"Example: litellm_params['embedding_model'] = 'azure/text-embedding-3-large'"
|
||||
@staticmethod
|
||||
def _search_request(
|
||||
vector_store_id: str,
|
||||
query_text: str,
|
||||
query_vector: Sequence[float],
|
||||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, object]]:
|
||||
scope: Final = {
|
||||
key: value
|
||||
for key, value in (
|
||||
("dbName", litellm_params.get("milvus_db_name")),
|
||||
("partitionNames", litellm_params.get("milvus_partition_names")),
|
||||
)
|
||||
|
||||
embedding_config: Final = litellm_params.get("litellm_embedding_config", {})
|
||||
if not embedding_config:
|
||||
raise ValueError(
|
||||
"embedding_config is required in litellm_params for Milvus. You can call any litellm embedding model."
|
||||
"Example: litellm_params['embedding_config'] = {'api_base': 'https://krris-mh44uf7y-eastus2.cognitiveservices.azure.com/', 'api_key': 'os.environ/AZURE_API_KEY', 'api_version': '2025-09-01'}"
|
||||
)
|
||||
|
||||
# Get top_k (number of results to return)
|
||||
# Generate embedding for the query using litellm.embeddings
|
||||
try:
|
||||
embedding_response: Final = litellm.embedding(
|
||||
model=embedding_model,
|
||||
input=[query],
|
||||
**embedding_config,
|
||||
)
|
||||
query_vector: Final = embedding_response.data[0]["embedding"]
|
||||
except Exception as e:
|
||||
raise Exception(f"Failed to generate embedding for query: {e}")
|
||||
|
||||
# Azure AI Search endpoint for search
|
||||
index_name: Final = vector_store_id # vector_store_id is the index name
|
||||
url: Final = f"{api_base}/v2/vectordb/entities/search"
|
||||
|
||||
# Build the request body for Azure AI Search with vector search
|
||||
request_body: Final[dict[str, Any]] = {
|
||||
"collectionName": index_name,
|
||||
if value
|
||||
}
|
||||
litellm_logging_obj.model_call_details["input"] = query_text
|
||||
litellm_logging_obj.model_call_details["embedding_model"] = litellm_params.get("litellm_embedding_model")
|
||||
return f"{api_base}/v2/vectordb/entities/search", {
|
||||
"collectionName": vector_store_id,
|
||||
"data": [query_vector],
|
||||
"annsField": "book_intro_vector",
|
||||
**vector_store_search_optional_params,
|
||||
**scope,
|
||||
}
|
||||
|
||||
db_name: Final = litellm_params.get("milvus_db_name")
|
||||
if db_name:
|
||||
request_body["dbName"] = db_name
|
||||
|
||||
partition_names: Final = litellm_params.get("milvus_partition_names")
|
||||
if partition_names:
|
||||
request_body["partitionNames"] = partition_names
|
||||
|
||||
#########################################################
|
||||
# Update logging object with details of the request
|
||||
#########################################################
|
||||
litellm_logging_obj.model_call_details["input"] = query
|
||||
litellm_logging_obj.model_call_details["embedding_model"] = embedding_model
|
||||
|
||||
return url, request_body
|
||||
|
||||
def transform_search_vector_store_response(
|
||||
self, response: httpx.Response, litellm_logging_obj: LiteLLMLoggingObj
|
||||
) -> VectorStoreSearchResponse:
|
||||
|
|
|
|||
|
|
@ -305,14 +305,16 @@ class BaseOpenAILLM:
|
|||
|
||||
# Get unified SSL configuration
|
||||
ssl_config: Final = get_ssl_configuration()
|
||||
transport: Final = AsyncHTTPHandler._create_async_transport(
|
||||
ssl_context=(ssl_config if isinstance(ssl_config, ssl.SSLContext) else None),
|
||||
ssl_verify=ssl_config if isinstance(ssl_config, bool) else None,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
|
||||
return httpx.AsyncClient(
|
||||
verify=ssl_config,
|
||||
transport=AsyncHTTPHandler._create_async_transport(
|
||||
ssl_context=(ssl_config if isinstance(ssl_config, ssl.SSLContext) else None),
|
||||
ssl_verify=ssl_config if isinstance(ssl_config, bool) else None,
|
||||
shared_session=shared_session,
|
||||
),
|
||||
transport=transport,
|
||||
mounts=AsyncHTTPHandler._create_httpx_proxy_mounts(transport, verify=ssl_config, cert=None),
|
||||
follow_redirects=True,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -15,7 +15,10 @@ import httpx
|
|||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.vector_store.transformation import BaseDirectVectorStoreConfig
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseDirectVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.llms.valkey.common_utils import build_valkey_url, pack_vector
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.types.vector_stores import (
|
||||
|
|
@ -213,6 +216,7 @@ class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
params: Final = _ValkeySearchParams.model_validate(litellm_params)
|
||||
|
|
@ -222,10 +226,18 @@ class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
embedding_field=params.embedding_field,
|
||||
text_field=params.text_field,
|
||||
)
|
||||
embedding_response: Final = self.embedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: litellm.embedding's input contract is a list
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
embedding_response: Final = (
|
||||
embedding_executor.embed(
|
||||
params.require_embedding_model(),
|
||||
query_text,
|
||||
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
|
||||
)
|
||||
if embedding_executor is not None
|
||||
else self.embedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: the injected embedding callable requires list input
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
)
|
||||
)
|
||||
vec_params: Final = {"vec": pack_vector(embedding_response.data[0]["embedding"])} # mutable-ok: redis-py API
|
||||
|
||||
|
|
@ -252,6 +264,7 @@ class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
vector_store_search_optional_params: VectorStoreSearchOptionalRequestParams,
|
||||
litellm_logging_obj: "LiteLLMLoggingObj",
|
||||
litellm_params: Mapping[str, object],
|
||||
embedding_executor: VectorStoreEmbeddingExecutor | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
) -> VectorStoreSearchResponse:
|
||||
params: Final = _ValkeySearchParams.model_validate(litellm_params)
|
||||
|
|
@ -261,10 +274,18 @@ class ValkeyVectorStoreConfig(BaseDirectVectorStoreConfig):
|
|||
embedding_field=params.embedding_field,
|
||||
text_field=params.text_field,
|
||||
)
|
||||
embedding_response: Final = await self.aembedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: litellm.embedding's input contract is a list
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
embedding_response: Final = (
|
||||
await embedding_executor.aembed(
|
||||
params.require_embedding_model(),
|
||||
query_text,
|
||||
params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG,
|
||||
)
|
||||
if embedding_executor is not None
|
||||
else await self.aembedding_fn(
|
||||
model=params.require_embedding_model(),
|
||||
input=[query_text], # mutable-ok: the injected embedding callable requires list input
|
||||
**(params.litellm_embedding_config or _EMPTY_EMBEDDING_CONFIG),
|
||||
)
|
||||
)
|
||||
vec_params: Final = {"vec": pack_vector(embedding_response.data[0]["embedding"])} # mutable-ok: redis-py API
|
||||
|
||||
|
|
|
|||
|
|
@ -121,6 +121,23 @@ class TokenEndpointClient:
|
|||
return Ok(ExchangedToken(access_token=parsed.access_token, expires_in=parsed.expires_in))
|
||||
|
||||
|
||||
class _KeyGuard:
|
||||
"""The per-key single-flight lock plus the invalidation generation that lock protects.
|
||||
|
||||
Both live on one object so their lifetimes cannot diverge. `get_or_compute` binds the guard to
|
||||
a local for its whole critical section, which keeps the weak map's entry alive for as long as
|
||||
that compute could still write; an `invalidate` overlapping the compute therefore reaches the
|
||||
very same object and its bump is guaranteed to be observed. Conversely a guard nobody holds is
|
||||
collectible precisely because no write is outstanding for it to fence.
|
||||
"""
|
||||
|
||||
__slots__ = ("__weakref__", "generation", "lock")
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.lock = asyncio.Lock()
|
||||
self.generation = 0
|
||||
|
||||
|
||||
class ExchangedTokenCache:
|
||||
"""Memoizes the final token string per key, single-flighting concurrent misses on one lock."""
|
||||
|
||||
|
|
@ -129,7 +146,7 @@ class ExchangedTokenCache:
|
|||
max_size_in_memory=MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE,
|
||||
default_ttl=MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL,
|
||||
)
|
||||
self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = weakref.WeakValueDictionary()
|
||||
self._guards: weakref.WeakValueDictionary[str, _KeyGuard] = weakref.WeakValueDictionary()
|
||||
|
||||
async def get_or_compute(
|
||||
self,
|
||||
|
|
@ -144,28 +161,50 @@ class ExchangedTokenCache:
|
|||
guaranteeing the token it gets back was minted for the *current* inputs: a stored entry
|
||||
whose fingerprint differs reads as a miss and is re-minted over. That keeps eviction
|
||||
addressable without the key having to encode the credential material it protects.
|
||||
|
||||
An `invalidate` landing while `compute` is in flight wins over that compute's write. The
|
||||
token is still returned to the caller it was minted for, but it is not stored, so the next
|
||||
resolution re-mints rather than serving a bearer that predates the invalidation for the
|
||||
rest of its TTL.
|
||||
"""
|
||||
cached = self._get(cache_key, fingerprint)
|
||||
if cached is not None:
|
||||
return Ok(cached)
|
||||
async with self._lock(cache_key):
|
||||
guard = self._guard(cache_key)
|
||||
async with guard.lock:
|
||||
cached = self._get(cache_key, fingerprint)
|
||||
if cached is not None:
|
||||
return Ok(cached)
|
||||
generation = guard.generation
|
||||
match await compute():
|
||||
case Ok(token):
|
||||
self._cache.set_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped
|
||||
cache_key,
|
||||
(fingerprint, token.access_token),
|
||||
ttl=_cache_ttl_seconds(token.expires_in),
|
||||
)
|
||||
if guard.generation == generation:
|
||||
self._store(cache_key, fingerprint, token)
|
||||
return Ok(token.access_token)
|
||||
case Error(err):
|
||||
return Error(err)
|
||||
|
||||
def invalidate(self, cache_key: str) -> None:
|
||||
"""Evict one cached token so the next `get_or_compute` re-mints (e.g. after an upstream 401)."""
|
||||
"""Evict one cached token so the next `get_or_compute` re-mints (e.g. after an upstream 401).
|
||||
|
||||
Bumping the guard's generation is what makes the eviction stick against a compute already
|
||||
awaiting the token endpoint: that compute snapshotted the old generation and so skips its
|
||||
write. No guard means no compute is in flight, since an in-flight one pins its own.
|
||||
|
||||
Stays synchronous: callers invalidate from plain `def`s.
|
||||
"""
|
||||
self._cache.delete_cache(cache_key) # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped
|
||||
guard = self._guards.get(cache_key)
|
||||
if guard is None:
|
||||
return
|
||||
guard.generation += 1
|
||||
|
||||
def _store(self, cache_key: str, fingerprint: str, token: ExchangedToken) -> None:
|
||||
self._cache.set_cache( # pyright: ignore[reportUnknownMemberType] # InMemoryCache is untyped
|
||||
cache_key,
|
||||
(fingerprint, token.access_token),
|
||||
ttl=_cache_ttl_seconds(token.expires_in),
|
||||
)
|
||||
|
||||
def _get(self, cache_key: str, fingerprint: str) -> str | None:
|
||||
"""The stored token, or None when absent or minted for different inputs.
|
||||
|
|
@ -180,12 +219,12 @@ class ExchangedTokenCache:
|
|||
return None
|
||||
return token if stored_fingerprint == fingerprint else None
|
||||
|
||||
def _lock(self, cache_key: str) -> asyncio.Lock:
|
||||
lock = self._locks.get(cache_key)
|
||||
if lock is None:
|
||||
lock = asyncio.Lock()
|
||||
self._locks[cache_key] = lock
|
||||
return lock
|
||||
def _guard(self, cache_key: str) -> _KeyGuard:
|
||||
guard = self._guards.get(cache_key)
|
||||
if guard is None:
|
||||
guard = _KeyGuard()
|
||||
self._guards[cache_key] = guard
|
||||
return guard
|
||||
|
||||
|
||||
def _cache_ttl_seconds(expires_in: int | None) -> int:
|
||||
|
|
|
|||
|
|
@ -3940,6 +3940,13 @@ if MCP_AVAILABLE:
|
|||
and server.auth_type == MCPAuth.oauth2_token_exchange
|
||||
and oauth2_headers
|
||||
and len(mcp_servers or []) == 1
|
||||
and server.server_id
|
||||
in frozenset(
|
||||
allowed.server_id
|
||||
for allowed in await _get_allowed_mcp_servers(
|
||||
user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip
|
||||
)
|
||||
)
|
||||
):
|
||||
await global_mcp_server_manager.preflight_token_exchange(
|
||||
server=server,
|
||||
|
|
|
|||
|
|
@ -100,9 +100,10 @@ class CliPollData(TypedDict, total=False):
|
|||
|
||||
|
||||
class CliSsoStartData(TypedDict):
|
||||
login_id: str
|
||||
poll_secret: str
|
||||
user_code: str
|
||||
login_id: ReadOnly[str]
|
||||
poll_secret: ReadOnly[str]
|
||||
user_code: ReadOnly[str]
|
||||
verification_uri_complete: ReadOnly[NotRequired[str]]
|
||||
|
||||
|
||||
class CliAuthResult(TypedDict):
|
||||
|
|
@ -860,11 +861,22 @@ def login(ctx: click.Context, config_claude: bool, pkce: bool) -> None:
|
|||
poll_secret: Final = cli_sso_flow["poll_secret"]
|
||||
user_code: Final = cli_sso_flow["user_code"]
|
||||
|
||||
sso_url = f"{base_url}/sso/key/generate?" + urlencode({"source": LITELLM_CLI_SOURCE_IDENTIFIER, "key": key_id})
|
||||
browser_prefills_code: Final = isinstance(cli_sso_flow.get("verification_uri_complete"), str)
|
||||
sso_url: Final = f"{base_url}/sso/key/generate?" + urlencode(
|
||||
(
|
||||
("source", LITELLM_CLI_SOURCE_IDENTIFIER),
|
||||
("key", key_id),
|
||||
*((("user_code", user_code),) if browser_prefills_code else ()),
|
||||
)
|
||||
)
|
||||
|
||||
click.echo(f"Opening browser to: {sso_url}")
|
||||
click.echo("Please complete the SSO authentication in your browser...")
|
||||
click.echo(f"Verification code: {user_code}")
|
||||
click.echo(
|
||||
f"Verification code: {user_code} (pre-filled in the browser, check it matches)"
|
||||
if browser_prefills_code
|
||||
else f"Verification code: {user_code}"
|
||||
)
|
||||
click.echo(f"Session ID: {key_id}")
|
||||
|
||||
# Open browser
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ import contextlib
|
|||
import json
|
||||
import logging
|
||||
import math
|
||||
import traceback
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
|
|
@ -18,7 +17,7 @@ from fastapi.responses import JSONResponse, Response, StreamingResponse
|
|||
from starlette.types import Receive, Scope, Send
|
||||
|
||||
import litellm
|
||||
from litellm._logging import _redact_string, verbose_proxy_logger
|
||||
from litellm._logging import redact_internal_details_from_client_message, verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import (
|
||||
DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE,
|
||||
|
|
@ -3417,7 +3416,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
else:
|
||||
_code = status.HTTP_500_INTERNAL_SERVER_ERROR
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
message=redact_internal_details_from_client_message(getattr(e, "message", error_msg)),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
openai_code=getattr(e, "code", None),
|
||||
|
|
@ -3629,10 +3628,8 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
if isinstance(e, HTTPException):
|
||||
raise e
|
||||
error_traceback: Final = _redact_string(traceback.format_exc())
|
||||
error_msg: Final = f"{e}\n\n{error_traceback}"
|
||||
proxy_exception: Final = ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
message=redact_internal_details_from_client_message(getattr(e, "message", str(e))),
|
||||
type=getattr(e, "type", "None"),
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", 500),
|
||||
|
|
|
|||
|
|
@ -514,12 +514,26 @@ async def update_guardrail(
|
|||
guardrail_name: Final = result.get("guardrail_name", "Unknown")
|
||||
|
||||
try:
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.update_in_memory_guardrail(
|
||||
guardrail_id=guardrail_id, guardrail=cast(Guardrail, result)
|
||||
)
|
||||
IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=cast(Guardrail, result))
|
||||
verbose_proxy_logger.info(
|
||||
"Immediate sync: Successfully updated guardrail '%s' (ID: %s)", guardrail_name, guardrail_id
|
||||
)
|
||||
except (ValueError, TypeError) as update_error:
|
||||
# The new config is invalid (a raising guardrail __init__):
|
||||
# reinitialize_guardrail already restored the previous live instance, but
|
||||
# update_guardrail_in_db above already persisted the rejected config to
|
||||
# the DB. Roll that back too, so the DB and the live guardrail never
|
||||
# disagree about what's actually enforcing, and surface the rejection to
|
||||
# the caller instead of a misleading 200.
|
||||
await GUARDRAIL_REGISTRY.update_guardrail_in_db(
|
||||
guardrail_id=guardrail_id,
|
||||
guardrail=existing_guardrail,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail=f"Invalid guardrail configuration, update rejected: {update_error}",
|
||||
) from update_error
|
||||
except Exception as update_error:
|
||||
verbose_proxy_logger.warning(
|
||||
"Immediate sync: Failed to update '%s' (ID: %s) in memory: %s",
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.types.guardrails import GuardrailEventHooks, SupportedGuardrailIntegrations
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
||||
from .crowdstrike_aidr import CrowdStrikeAIDRHandler
|
||||
from .crowdstrike_aidr import CrowdStrikeAIDRHandler, streaming_params_from_litellm_params
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
|
@ -15,17 +15,16 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
if not guardrail_name:
|
||||
raise ValueError("CrowdStrike AIDR guardrail name is required")
|
||||
|
||||
streaming_params: Final = streaming_params_from_litellm_params(litellm_params)
|
||||
_crowdstrike_aidr_callback: Final = CrowdStrikeAIDRHandler(
|
||||
guardrail_name=guardrail_name,
|
||||
api_base=litellm_params.api_base,
|
||||
api_key=litellm_params.api_key,
|
||||
# Exclude during_call to prevent duplicate input events
|
||||
event_hook=[
|
||||
GuardrailEventHooks.pre_call.value,
|
||||
GuardrailEventHooks.post_call.value,
|
||||
],
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
fail_on_error=litellm_params.fail_on_error,
|
||||
streaming_end_of_stream_only=streaming_params.streaming_end_of_stream_only,
|
||||
streaming_sampling_rate=streaming_params.streaming_sampling_rate,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_crowdstrike_aidr_callback)
|
||||
|
||||
|
|
|
|||
|
|
@ -24,8 +24,11 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
|
||||
from litellm.types.llms.openai import AllMessageValues, OpenAIChatCompletionToolParam
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.crowdstrike_aidr import (
|
||||
CrowdStrikeAIDRGuardrailConfigModelOptionalParams,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -153,6 +156,21 @@ def _merge_metadata_bags(request_data: Mapping[str, Any]) -> Mapping[str, Any] |
|
|||
return merged if present else None
|
||||
|
||||
|
||||
def streaming_params_from_litellm_params(
|
||||
litellm_params: LitellmParams,
|
||||
) -> CrowdStrikeAIDRGuardrailConfigModelOptionalParams:
|
||||
extras: Final[Mapping[str, object]] = litellm_params.model_extra or {}
|
||||
nested: Final = litellm_params.optional_params
|
||||
optional_params: Final[Mapping[str, object]] = {} if nested is None else nested.model_dump()
|
||||
return CrowdStrikeAIDRGuardrailConfigModelOptionalParams.model_validate(
|
||||
{
|
||||
name: value
|
||||
for name in CrowdStrikeAIDRGuardrailConfigModelOptionalParams.model_fields
|
||||
if (value := optional_params.get(name, extras.get(name))) is not None
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _messages_since_last_assistant(
|
||||
messages: Sequence[AllMessageValues],
|
||||
) -> _FilteredMessages:
|
||||
|
|
@ -241,6 +259,8 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
fail_on_error: bool | None = True,
|
||||
streaming_end_of_stream_only: bool | None = None,
|
||||
streaming_sampling_rate: int | None = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
"""
|
||||
|
|
@ -250,10 +270,19 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
guardrail_name (str): The name of the guardrail instance.
|
||||
api_key (str | None): The CrowdStrike AIDR API key. Reads from CS_AIDR_TOKEN env var if None.
|
||||
api_base (str | None): The CrowdStrike AIDR API base URL. Reads from CS_AIDR_BASE_URL env var if None.
|
||||
streaming_end_of_stream_only (bool | None): Scan streamed output once at end of stream instead of
|
||||
every streaming_sampling_rate chunks. Defaults to False.
|
||||
streaming_sampling_rate (int | None): Scan the accumulated streamed output every Nth chunk. Defaults to 5.
|
||||
**kwargs: Additional arguments passed to the CustomGuardrail base class.
|
||||
"""
|
||||
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
|
||||
self.fail_on_error = True if fail_on_error is None else fail_on_error
|
||||
self._set_streaming_params(
|
||||
CrowdStrikeAIDRGuardrailConfigModelOptionalParams(
|
||||
streaming_end_of_stream_only=streaming_end_of_stream_only,
|
||||
streaming_sampling_rate=streaming_sampling_rate,
|
||||
)
|
||||
)
|
||||
|
||||
self.api_key = api_key or os.environ.get("CS_AIDR_TOKEN")
|
||||
if not self.api_key:
|
||||
|
|
@ -274,6 +303,15 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
|
|||
"Initialized CrowdStrike AIDR Guardrail: name=%s, api_base=%s", guardrail_name, self.api_base
|
||||
)
|
||||
|
||||
def _set_streaming_params(self, streaming_params: CrowdStrikeAIDRGuardrailConfigModelOptionalParams) -> None:
|
||||
self.streaming_end_of_stream_only: bool = streaming_params.streaming_end_of_stream_only or False
|
||||
self.streaming_sampling_rate: int = streaming_params.streaming_sampling_rate or 5
|
||||
|
||||
@override
|
||||
def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None:
|
||||
super().update_in_memory_litellm_params(litellm_params)
|
||||
self._set_streaming_params(streaming_params_from_litellm_params(litellm_params))
|
||||
|
||||
async def _call_crowdstrike_aidr_guard(
|
||||
self, payload: dict[str, Any], hook_name: str
|
||||
) -> _GuardChatCompletionsResult:
|
||||
|
|
|
|||
|
|
@ -826,11 +826,12 @@ class InMemoryGuardrailHandler:
|
|||
Removes old callback from litellm.callbacks and creates fresh instance.
|
||||
|
||||
If the new config fails to initialize (e.g. an invalid on_flagged
|
||||
combination), the previous instance is restored rather than left
|
||||
deleted: initialize_guardrail's own ValueError/TypeError propagate
|
||||
uncaught, so a caller reaching this point after already deleting the
|
||||
old instance would otherwise leave the guardrail providing no
|
||||
protection at all, not merely "still enforcing the old config."
|
||||
combination or an invalid regex), the previous instance is restored
|
||||
rather than left deleted, and the failure is re-raised as ValueError so
|
||||
every init failure reaches callers as one exception type: a caller
|
||||
reaching this point after already deleting the old instance would
|
||||
otherwise leave the guardrail providing no protection at all, not
|
||||
merely "still enforcing the old config."
|
||||
"""
|
||||
guardrail_id: Final = guardrail.get("guardrail_id")
|
||||
if not guardrail_id:
|
||||
|
|
@ -849,7 +850,7 @@ class InMemoryGuardrailHandler:
|
|||
# that was enforcing must never fail open because an update was bad.
|
||||
try:
|
||||
return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source)
|
||||
except Exception:
|
||||
except Exception as init_error:
|
||||
if previous_guardrail is not None:
|
||||
verbose_proxy_logger.exception(
|
||||
"Reinitializing guardrail %s with updated params failed; restoring the previous configuration",
|
||||
|
|
@ -861,7 +862,7 @@ class InMemoryGuardrailHandler:
|
|||
)
|
||||
except Exception: # noqa: BLE001 # the original failure must propagate even if the restore breaks
|
||||
verbose_proxy_logger.exception("Restoring previous guardrail %s also failed", guardrail_id)
|
||||
raise
|
||||
raise ValueError(f"Guardrail initialization failed: {init_error}") from init_error
|
||||
|
||||
def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -261,6 +261,7 @@ class ProxyInitializationHelpers:
|
|||
"app": "litellm.proxy.proxy_server:app",
|
||||
"host": host,
|
||||
"port": port,
|
||||
"server_header": False,
|
||||
}
|
||||
if log_config is not None:
|
||||
print(f"Using log_config: {log_config}")
|
||||
|
|
|
|||
|
|
@ -1,11 +1,13 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from importlib.resources import files
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
|
|
@ -28,6 +30,7 @@ from litellm.types.proxy.management_endpoints.model_management_endpoints import
|
|||
)
|
||||
from litellm.types.proxy.public_endpoints.public_endpoints import (
|
||||
AgentCreateInfo,
|
||||
AutoRouterPresetRecord,
|
||||
ComplexityScorerDefaults,
|
||||
ProviderCreateInfo,
|
||||
PublicModelHubInfo,
|
||||
|
|
@ -464,6 +467,86 @@ async def get_litellm_blog_posts():
|
|||
return BlogPostsResponse(posts=posts)
|
||||
|
||||
|
||||
_AUTOROUTER_PRESETS_ADAPTER: Final = TypeAdapter(dict[str, AutoRouterPresetRecord])
|
||||
|
||||
|
||||
def _load_bundled_autorouter_presets() -> Mapping[str, AutoRouterPresetRecord]:
|
||||
raw: Final = json.loads(
|
||||
files("litellm.proxy.public_endpoints").joinpath("autorouter_presets.json").read_text(encoding="utf-8")
|
||||
)
|
||||
return _AUTOROUTER_PRESETS_ADAPTER.validate_python(raw)
|
||||
|
||||
|
||||
async def _fetch_remote_autorouter_presets(url: str) -> Mapping[str, AutoRouterPresetRecord]:
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.UI)
|
||||
response: Final = await client.get(url, timeout=5.0)
|
||||
response.raise_for_status()
|
||||
presets: Final = _AUTOROUTER_PRESETS_ADAPTER.validate_python(response.json())
|
||||
if not presets:
|
||||
raise ValueError("remote auto-router preset catalog is empty")
|
||||
return presets
|
||||
|
||||
|
||||
async def _resolve_autorouter_presets(
|
||||
url: str,
|
||||
fetch: Callable[[str], Awaitable[Mapping[str, AutoRouterPresetRecord]]],
|
||||
) -> Mapping[str, AutoRouterPresetRecord]:
|
||||
if os.getenv("LITELLM_LOCAL_AUTOROUTER_PRESETS", "").lower() == "true":
|
||||
return _load_bundled_autorouter_presets()
|
||||
try:
|
||||
return await fetch(url)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
"LiteLLM: failed to fetch auto-router presets from %s: %s. Serving the bundled catalog for the life of this process.",
|
||||
url,
|
||||
str(e),
|
||||
)
|
||||
return _load_bundled_autorouter_presets()
|
||||
|
||||
|
||||
class _AutoRouterPresetsCache:
|
||||
presets: Mapping[str, AutoRouterPresetRecord] | None = None
|
||||
lock: asyncio.Lock | None = None
|
||||
|
||||
|
||||
async def get_autorouter_presets(
|
||||
url: str,
|
||||
fetch: Callable[[str], Awaitable[Mapping[str, AutoRouterPresetRecord]]] = _fetch_remote_autorouter_presets,
|
||||
) -> Mapping[str, AutoRouterPresetRecord]:
|
||||
cached: Final = _AutoRouterPresetsCache.presets
|
||||
if cached is not None:
|
||||
return cached
|
||||
if _AutoRouterPresetsCache.lock is None:
|
||||
_AutoRouterPresetsCache.lock = asyncio.Lock()
|
||||
async with _AutoRouterPresetsCache.lock:
|
||||
held: Final = _AutoRouterPresetsCache.presets
|
||||
if held is not None:
|
||||
return held
|
||||
resolved: Final = await _resolve_autorouter_presets(url=url, fetch=fetch)
|
||||
_AutoRouterPresetsCache.presets = resolved
|
||||
return resolved
|
||||
|
||||
|
||||
@router.get(
|
||||
"/public/autorouter_presets",
|
||||
tags=["public", "auto router"], # mutable-ok: FastAPI route tags take a list
|
||||
response_model=dict[str, AutoRouterPresetRecord],
|
||||
)
|
||||
async def get_public_autorouter_presets() -> Mapping[str, AutoRouterPresetRecord]:
|
||||
"""
|
||||
Return the auto-router preset catalog the dashboard's template picker renders.
|
||||
|
||||
Resolved once per process, like the model cost map: fetched from ``litellm.autorouter_presets_url``
|
||||
(override with ``LITELLM_AUTOROUTER_PRESETS_URL``) on the first request, falling back to the
|
||||
catalog bundled with the package on any failure. Set ``LITELLM_LOCAL_AUTOROUTER_PRESETS=True``
|
||||
to serve the bundled catalog only. A restart picks up a newly published catalog.
|
||||
"""
|
||||
return await get_autorouter_presets(url=litellm.autorouter_presets_url)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/public/endpoints",
|
||||
tags=["public"],
|
||||
|
|
|
|||
|
|
@ -729,7 +729,7 @@ async def rag_query(
|
|||
# conflict so callers cannot override the store's provider or credentials.
|
||||
managed_store: Final = resolved_stores.get(retrieval_config["vector_store_id"])
|
||||
store_data: Final = (
|
||||
await build_request_data_from_managed_vector_store(managed_store)
|
||||
build_request_data_from_managed_vector_store(managed_store)
|
||||
if managed_store is not None
|
||||
else MappingProxyType({})
|
||||
)
|
||||
|
|
|
|||
|
|
@ -55,6 +55,10 @@ router: Final = APIRouter()
|
|||
|
||||
SPEND_LOGS_PAGINATION_COUNT_CAP: Final = 10000
|
||||
|
||||
_SESSION_GROUP_KEY_SQL: Final = "COALESCE(NULLIF(session_id, ''), request_id), api_key"
|
||||
_MCP_CALL_TYPES_SQL: Final = "('call_mcp_tool', 'list_mcp_tools')"
|
||||
_AGENT_CALL_TYPE_SQL: Final = "'asend_message'"
|
||||
|
||||
_INTERNAL_HEALTH_CHECK_API_KEYS: Final = (
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
hash_token(token=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME),
|
||||
|
|
@ -144,21 +148,16 @@ class _DailyTagSpendRow(TypedDict):
|
|||
total_spend: float
|
||||
|
||||
|
||||
class _SessionCountAggregate(TypedDict):
|
||||
session_id: int
|
||||
|
||||
|
||||
class _SessionCountRow(TypedDict):
|
||||
session_id: str
|
||||
_count: _SessionCountAggregate
|
||||
|
||||
|
||||
class _SessionSpendRow(TypedDict):
|
||||
session_id: str
|
||||
api_key: ReadOnly[str]
|
||||
session_total_count: ReadOnly[int]
|
||||
session_total_spend: float
|
||||
mcp_tool_call_count: int
|
||||
mcp_tool_call_spend: float
|
||||
session_cache_hit_count: ReadOnly[int]
|
||||
session_llm_count: ReadOnly[int]
|
||||
session_agent_count: ReadOnly[int]
|
||||
|
||||
|
||||
class _SpendSumAggregate(TypedDict, total=False):
|
||||
|
|
@ -242,18 +241,6 @@ async def _count_spend_logs(prisma_client: PrismaClient, where: Mapping[str, obj
|
|||
return await _spend_logs_table(prisma_client).count(where=where)
|
||||
|
||||
|
||||
async def _count_logs_per_session(
|
||||
prisma_client: PrismaClient, session_ids: Sequence[str | None]
|
||||
) -> Sequence[_SessionCountRow]:
|
||||
"""Count spend log rows per session for the given session ids."""
|
||||
rows: Final = await _spend_logs_table(prisma_client).group_by(
|
||||
by=["session_id"],
|
||||
where={"session_id": {"in": session_ids}},
|
||||
count={"session_id": True},
|
||||
)
|
||||
return cast(Sequence[_SessionCountRow], rows) # cast-ok: group_by(count=) shape is fixed by the by/count args
|
||||
|
||||
|
||||
async def _find_team_row(prisma_client: PrismaClient, team_id: str) -> _SupportsModelDump | None:
|
||||
"""Read a single team row as a Prisma model instance."""
|
||||
return await _team_table(prisma_client).find_unique(where={"team_id": team_id})
|
||||
|
|
@ -2290,6 +2277,10 @@ async def ui_view_spend_logs(
|
|||
default=False,
|
||||
description="Exclude LiteLLM internal health check requests from results",
|
||||
),
|
||||
group_by_session: bool = fastapi.Query(
|
||||
default=False,
|
||||
description="Paginate over sessions instead of raw logs: one representative row per session, total counts sessions",
|
||||
),
|
||||
):
|
||||
"""
|
||||
View spend logs with pagination support.
|
||||
|
|
@ -2644,12 +2635,16 @@ async def ui_view_spend_logs(
|
|||
else:
|
||||
_order_expr = order_column
|
||||
|
||||
joined_conditions: Final = " AND ".join(sql_conditions)
|
||||
session_grouping: Final = group_by_session is True
|
||||
count_group_clause: Final = f"GROUP BY {_SESSION_GROUP_KEY_SQL}" if session_grouping else ""
|
||||
count_query: Final = f"""
|
||||
SELECT COUNT(*) AS total_count
|
||||
FROM (
|
||||
SELECT 1
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {" AND ".join(sql_conditions)}
|
||||
WHERE {joined_conditions}
|
||||
{count_group_clause}
|
||||
LIMIT ${p}
|
||||
) AS bounded_matches
|
||||
"""
|
||||
|
|
@ -2660,21 +2655,36 @@ async def ui_view_spend_logs(
|
|||
total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP
|
||||
total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total
|
||||
|
||||
sql_query: Final = f"""
|
||||
SELECT
|
||||
request_id, call_type, api_key, spend, total_tokens,
|
||||
select_columns: Final = """request_id, call_type, api_key, spend, total_tokens,
|
||||
prompt_tokens, completion_tokens, "startTime", "endTime",
|
||||
"completionStartTime", model, model_id, model_group,
|
||||
custom_llm_provider, api_base, "user", metadata,
|
||||
cache_hit, cache_key, request_tags, team_id,
|
||||
organization_id, end_user, requester_ip_address,
|
||||
session_id, status, mcp_namespaced_tool_name, agent_id,
|
||||
COALESCE(request_duration_ms, (EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER) AS request_duration_ms
|
||||
COALESCE(request_duration_ms, (EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER) AS request_duration_ms"""
|
||||
sql_query: Final = (
|
||||
f"""
|
||||
SELECT * FROM (
|
||||
SELECT DISTINCT ON ({_SESSION_GROUP_KEY_SQL})
|
||||
{select_columns}
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {joined_conditions}
|
||||
ORDER BY {_SESSION_GROUP_KEY_SQL}, call_type IN {_MCP_CALL_TYPES_SQL}, "startTime" DESC
|
||||
) AS session_representatives
|
||||
ORDER BY {_order_expr} {_sql_dir}{_nulls_clause}, request_id
|
||||
LIMIT ${p} OFFSET ${p + 1}
|
||||
"""
|
||||
if session_grouping
|
||||
else f"""
|
||||
SELECT
|
||||
{select_columns}
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE {" AND ".join(sql_conditions)}
|
||||
WHERE {joined_conditions}
|
||||
ORDER BY {_order_expr} {_sql_dir}{_nulls_clause}
|
||||
LIMIT ${p} OFFSET ${p + 1}
|
||||
"""
|
||||
)
|
||||
sql_params.extend([page_size, skip])
|
||||
|
||||
data: Final = await prisma_client.db.query_raw(sql_query, *sql_params)
|
||||
|
|
@ -4075,11 +4085,12 @@ async def _build_ui_spend_logs_response(
|
|||
Build the paginated response for the UI spend-logs endpoint.
|
||||
|
||||
When ``enrich_session_counts`` is ``True`` (the default for the v1/UI
|
||||
endpoint), each row is enriched with ``session_total_count`` so the
|
||||
frontend knows which sessions are expandable (multi-call sessions).
|
||||
For every row that carries a ``session_id``, a single ``GROUP BY`` query
|
||||
fetches the total number of logs in each referenced session. Rows without
|
||||
a ``session_id`` default to ``1``.
|
||||
endpoint), each row is enriched with ``session_total_count`` plus spend
|
||||
and call-type aggregates so the frontend knows which sessions are
|
||||
expandable (multi-call sessions). One ``GROUP BY (session_id, api_key)``
|
||||
query serves every referenced session, keyed per api key so two callers
|
||||
reusing a session id never see each other's totals. Rows without a
|
||||
``session_id`` default to ``1``.
|
||||
|
||||
When ``enrich_session_counts`` is ``False`` (v2 endpoint), rows are
|
||||
serialised without the extra query.
|
||||
|
|
@ -4101,7 +4112,6 @@ async def _build_ui_spend_logs_response(
|
|||
A dict with ``data`` (enriched rows), ``total``, ``page``,
|
||||
``page_size``, ``total_pages``, and ``total_is_capped``.
|
||||
"""
|
||||
count_map: dict[str, int] = {}
|
||||
if enrich_session_counts:
|
||||
session_ids: Final[Sequence[str | None]] = list(
|
||||
{
|
||||
|
|
@ -4110,15 +4120,8 @@ async def _build_ui_spend_logs_response(
|
|||
if (row.get("session_id") if isinstance(row, dict) else getattr(row, "session_id", None))
|
||||
}
|
||||
)
|
||||
if session_ids:
|
||||
# NOTE: This GROUP BY runs on every v1/UI page load. The IN clause
|
||||
# is bounded by page_size (typically 25-50 distinct session IDs).
|
||||
# If performance degrades at scale, consider short-lived caching or
|
||||
# folding the count into the main query via a window function.
|
||||
counts: Final = await _count_logs_per_session(prisma_client, session_ids)
|
||||
count_map = {r["session_id"]: r["_count"]["session_id"] for r in counts if r.get("session_id")}
|
||||
|
||||
session_spend_map: dict[str, dict[str, int | float]] = {}
|
||||
session_spend_map: dict[tuple[str, str], dict[str, int | float]] = {}
|
||||
if enrich_session_counts and session_ids:
|
||||
from prisma.errors import PrismaError
|
||||
|
||||
|
|
@ -4130,38 +4133,46 @@ async def _build_ui_spend_logs_response(
|
|||
{
|
||||
(row.get("api_key") if isinstance(row, dict) else getattr(row, "api_key", None))
|
||||
for row in data
|
||||
if (row.get("api_key") if isinstance(row, dict) else getattr(row, "api_key", None))
|
||||
if (row.get("api_key") if isinstance(row, dict) else getattr(row, "api_key", None)) is not None
|
||||
}
|
||||
)
|
||||
rows: Final[Sequence[_SessionSpendRow]] = await _query_raw(
|
||||
prisma_client,
|
||||
"""
|
||||
SELECT session_id,
|
||||
f"""
|
||||
SELECT session_id, api_key,
|
||||
COUNT(*)::int AS session_total_count,
|
||||
COALESCE(SUM(spend), 0)::double precision AS session_total_spend,
|
||||
COUNT(*) FILTER (
|
||||
WHERE call_type IN ('call_mcp_tool', 'list_mcp_tools')
|
||||
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
|
||||
)::int AS mcp_tool_call_count,
|
||||
COALESCE(SUM(spend) FILTER (
|
||||
WHERE call_type IN ('call_mcp_tool', 'list_mcp_tools')
|
||||
WHERE call_type IN {_MCP_CALL_TYPES_SQL}
|
||||
), 0)::double precision AS mcp_tool_call_spend,
|
||||
COUNT(*) FILTER (WHERE LOWER(cache_hit) = 'true')::int AS session_cache_hit_count
|
||||
COUNT(*) FILTER (WHERE LOWER(cache_hit) = 'true')::int AS session_cache_hit_count,
|
||||
COUNT(*) FILTER (
|
||||
WHERE call_type NOT IN {_MCP_CALL_TYPES_SQL} AND call_type != {_AGENT_CALL_TYPE_SQL}
|
||||
)::int AS session_llm_count,
|
||||
COUNT(*) FILTER (WHERE call_type = {_AGENT_CALL_TYPE_SQL})::int AS session_agent_count
|
||||
FROM "LiteLLM_SpendLogs"
|
||||
WHERE session_id = ANY($1::text[])
|
||||
AND api_key = ANY($2::text[])
|
||||
GROUP BY session_id
|
||||
GROUP BY session_id, api_key
|
||||
""",
|
||||
session_ids,
|
||||
authorized_api_keys,
|
||||
)
|
||||
session_spend_map = {
|
||||
row["session_id"]: {
|
||||
(row["session_id"], row["api_key"]): {
|
||||
"session_total_count": int(row.get("session_total_count") or 0),
|
||||
"session_total_spend": float(row.get("session_total_spend") or 0.0),
|
||||
"mcp_tool_call_count": int(row.get("mcp_tool_call_count") or 0),
|
||||
"mcp_tool_call_spend": float(row.get("mcp_tool_call_spend") or 0.0),
|
||||
"session_cache_hit_count": int(row.get("session_cache_hit_count") or 0),
|
||||
"session_llm_count": int(row.get("session_llm_count") or 0),
|
||||
"session_agent_count": int(row.get("session_agent_count") or 0),
|
||||
}
|
||||
for row in rows
|
||||
if row.get("session_id")
|
||||
if row.get("session_id") and row.get("api_key") is not None
|
||||
}
|
||||
except PrismaError:
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -4174,14 +4185,17 @@ async def _build_ui_spend_logs_response(
|
|||
for row in data:
|
||||
row_dict = dict(row) if isinstance(row, dict) else row.model_dump()
|
||||
sid = row_dict.get("session_id")
|
||||
row_dict["session_total_count"] = count_map.get(sid, 1) if sid else 1
|
||||
session_stats = session_spend_map.get(sid) if sid else None
|
||||
row_api_key = row_dict.get("api_key")
|
||||
session_stats = session_spend_map.get((sid, row_api_key)) if sid and row_api_key is not None else None
|
||||
row_dict["session_total_count"] = int(session_stats["session_total_count"]) if session_stats else 1
|
||||
if session_stats:
|
||||
row_dict["session_total_spend"] = session_stats["session_total_spend"]
|
||||
if session_stats["mcp_tool_call_count"]:
|
||||
row_dict["mcp_tool_call_count"] = session_stats["mcp_tool_call_count"]
|
||||
row_dict["mcp_tool_call_spend"] = session_stats["mcp_tool_call_spend"]
|
||||
row_dict["session_cache_hit_count"] = session_stats["session_cache_hit_count"]
|
||||
row_dict["session_llm_count"] = session_stats["session_llm_count"]
|
||||
row_dict["session_agent_count"] = session_stats["session_agent_count"]
|
||||
enriched.append(row_dict)
|
||||
response_data: list = enriched
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -16,9 +16,6 @@ from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.utils import jsonify_object
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import (
|
||||
_resolve_embedding_config,
|
||||
)
|
||||
from litellm.proxy.vector_store_endpoints.utils import (
|
||||
assert_proxy_admin_for_vector_store_index_management,
|
||||
assert_user_can_access_vector_store,
|
||||
|
|
@ -57,19 +54,9 @@ def reject_caller_embedding_selection_params(payload: Mapping[str, object], sour
|
|||
########################################################
|
||||
|
||||
|
||||
async def build_request_data_from_managed_vector_store(
|
||||
def build_request_data_from_managed_vector_store(
|
||||
vector_store: LiteLLM_ManagedVectorStore,
|
||||
) -> Mapping[str, object]:
|
||||
"""
|
||||
Build request params (provider, credential ref, litellm_params) from an
|
||||
already-resolved managed vector store.
|
||||
|
||||
``litellm_embedding_config`` is resolved here, at request-handling time,
|
||||
instead of at row-creation time: the resolved api_key/api_base/api_version
|
||||
lives only in the returned per-request mapping and is never persisted back
|
||||
to the registry cache. Legacy rows that already carry a resolved
|
||||
(cleartext) config skip the lookup and pass through unchanged.
|
||||
"""
|
||||
top_level: Final = MappingProxyType(
|
||||
{
|
||||
key: vector_store.get(key)
|
||||
|
|
@ -78,18 +65,7 @@ async def build_request_data_from_managed_vector_store(
|
|||
}
|
||||
)
|
||||
litellm_params: Final = vector_store.get("litellm_params") or MappingProxyType({})
|
||||
embedding_model: Final = litellm_params.get("litellm_embedding_model")
|
||||
if not embedding_model or litellm_params.get("litellm_embedding_config"):
|
||||
return MappingProxyType({**top_level, **litellm_params})
|
||||
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
resolved_config: Final = await _resolve_embedding_config(
|
||||
embedding_model=embedding_model, prisma_client=prisma_client
|
||||
)
|
||||
if not resolved_config:
|
||||
return MappingProxyType({**top_level, **litellm_params})
|
||||
return MappingProxyType({**top_level, **litellm_params, "litellm_embedding_config": resolved_config})
|
||||
return MappingProxyType({**top_level, **litellm_params})
|
||||
|
||||
|
||||
async def _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
|
|
@ -118,7 +94,7 @@ async def _update_request_data_with_litellm_managed_vector_store_registry(
|
|||
vector_store=vector_store_to_run,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
return {**data, **(await build_request_data_from_managed_vector_store(vector_store_to_run))}
|
||||
return {**data, **build_request_data_from_managed_vector_store(vector_store_to_run)}
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
|
|||
|
|
@ -18,11 +18,8 @@ if TYPE_CHECKING:
|
|||
from prisma.models import LiteLLM_ManagedVectorStoresTable as _VectorStoreRow
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.router import Router
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import REDACTED_BY_LITELM_STRING
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
|
|
@ -32,13 +29,10 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
||||
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
|
||||
from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store
|
||||
from litellm.repositories.model_repository import ModelRepository
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.repositories.table_repositories import ManagedVectorStoresRepository
|
||||
from litellm.secret_managers.main import get_secret
|
||||
from litellm.types.vector_stores import (
|
||||
LiteLLM_ManagedVectorStore,
|
||||
LiteLLM_ManagedVectorStoreListResponse,
|
||||
|
|
@ -64,28 +58,6 @@ _LITELLM_PARAMS_MASKER: Final = SensitiveDataMasker()
|
|||
|
||||
_REDACT_LITELLM_PARAMS_MAX_DEPTH: Final = 10
|
||||
|
||||
# Use-time embedding-config resolution runs on every vector-store request
|
||||
# whose persisted row carries only a model reference (the post-fix shape).
|
||||
# Without a cache, that's one ``litellm_proxymodeltable.find_first`` per
|
||||
# request — the no-DB-in-critical-path rule. Hold the resolved config in
|
||||
# memory for a short TTL so a hot model name pays the DB lookup at most
|
||||
# once per ``_EMBEDDING_CONFIG_CACHE_TTL`` seconds. Cleartext credentials
|
||||
# only ever live in process memory (never persisted, never echoed in
|
||||
# management responses), so the cache doesn't widen the disclosure surface.
|
||||
_EMBEDDING_CONFIG_CACHE_TTL: Final = 60
|
||||
_EMBEDDING_CONFIG_CACHE_MAX_SIZE: Final = 256
|
||||
_embedding_config_cache: InMemoryCache | None = None
|
||||
|
||||
|
||||
def _get_embedding_config_cache() -> InMemoryCache:
|
||||
global _embedding_config_cache
|
||||
if _embedding_config_cache is None:
|
||||
_embedding_config_cache = InMemoryCache(
|
||||
max_size_in_memory=_EMBEDDING_CONFIG_CACHE_MAX_SIZE,
|
||||
default_ttl=_EMBEDDING_CONFIG_CACHE_TTL,
|
||||
)
|
||||
return _embedding_config_cache
|
||||
|
||||
|
||||
def _redact_sensitive_litellm_params(litellm_params: Any, _depth: int = 0) -> Any:
|
||||
"""
|
||||
|
|
@ -155,235 +127,6 @@ async def _fetch_and_authorize_vector_store(
|
|||
return typed
|
||||
|
||||
|
||||
def _resolve_embedding_config_from_router(embedding_model: str, llm_router) -> dict[str, object] | None:
|
||||
"""
|
||||
Resolve embedding config from router's config-defined models.
|
||||
|
||||
Config-defined models (from proxy_config.yaml) are stored in the router's model_list,
|
||||
not in the database. This function looks up the model in the router and extracts
|
||||
api_key, api_base, and api_version from the deployment's litellm_params.
|
||||
|
||||
Args:
|
||||
embedding_model: The embedding model string (e.g., "text-embedding-ada-002" or "azure/text-embedding-3-large")
|
||||
llm_router: The LiteLLM router instance
|
||||
|
||||
Returns:
|
||||
Dictionary with api_key, api_base, and api_version if model found, None otherwise
|
||||
"""
|
||||
if not embedding_model or llm_router is None:
|
||||
return None
|
||||
|
||||
# Extract model name candidates - could be "text-embedding-ada-002" or "azure/text-embedding-3-large"
|
||||
# Try exact match first, then try without provider prefix
|
||||
model_name_candidates: Final = [embedding_model]
|
||||
if "/" in embedding_model:
|
||||
# If it has a provider prefix, also try without it
|
||||
_, model_name = embedding_model.split("/", 1)
|
||||
model_name_candidates.append(model_name)
|
||||
|
||||
# Try to find model in router
|
||||
for model_name in model_name_candidates:
|
||||
try:
|
||||
# Try to get deployment by model group name (model_name in config)
|
||||
deployment = llm_router.get_deployment_by_model_group_name(model_group_name=model_name)
|
||||
|
||||
if deployment is not None and deployment.litellm_params is not None:
|
||||
litellm_params = deployment.litellm_params
|
||||
|
||||
# Build embedding config from model params
|
||||
embedding_config: dict[str, object] = {}
|
||||
|
||||
# Extract api_key
|
||||
api_key = getattr(litellm_params, "api_key", None)
|
||||
if api_key:
|
||||
# Handle os.environ/ prefix
|
||||
if isinstance(api_key, str) and api_key.startswith("os.environ/"):
|
||||
api_key = get_secret(api_key)
|
||||
embedding_config["api_key"] = api_key
|
||||
|
||||
# Extract api_base
|
||||
api_base = getattr(litellm_params, "api_base", None)
|
||||
if api_base:
|
||||
# Handle os.environ/ prefix
|
||||
if isinstance(api_base, str) and api_base.startswith("os.environ/"):
|
||||
api_base = get_secret(api_base)
|
||||
embedding_config["api_base"] = api_base
|
||||
|
||||
# Extract api_version
|
||||
api_version = getattr(litellm_params, "api_version", None)
|
||||
if api_version:
|
||||
embedding_config["api_version"] = api_version
|
||||
|
||||
project_id = getattr(litellm_params, "project_id", None)
|
||||
if project_id:
|
||||
embedding_config["project_id"] = project_id
|
||||
|
||||
# Only return config if we have at least api_key or api_base
|
||||
if embedding_config:
|
||||
verbose_proxy_logger.debug(
|
||||
"Resolved embedding config from router model %s: %s", model_name, list(embedding_config.keys())
|
||||
)
|
||||
return embedding_config
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Error resolving embedding config from router for model %s: %s", model_name, e)
|
||||
continue
|
||||
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_embedding_config_from_db(
|
||||
embedding_model: str, prisma_client: "PrismaClient"
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Resolve embedding config from database model configuration.
|
||||
|
||||
If litellm_embedding_model is provided but litellm_embedding_config is not,
|
||||
this function looks up the model in the database and extracts api_key, api_base,
|
||||
and api_version from the model's litellm_params to build the embedding config.
|
||||
|
||||
Args:
|
||||
embedding_model: The embedding model string (e.g., "text-embedding-ada-002" or "azure/text-embedding-3-large")
|
||||
prisma_client: The Prisma client instance
|
||||
|
||||
Returns:
|
||||
Dictionary with api_key, api_base, and api_version if model found, None otherwise
|
||||
"""
|
||||
if not embedding_model:
|
||||
return None
|
||||
|
||||
# Extract model name - could be "text-embedding-ada-002" or "azure/text-embedding-3-large"
|
||||
# Try to find model by exact match first, then try without provider prefix
|
||||
model_name_candidates: Final = [embedding_model]
|
||||
if "/" in embedding_model:
|
||||
# If it has a provider prefix, also try without it
|
||||
_, model_name = embedding_model.split("/", 1)
|
||||
model_name_candidates.append(model_name)
|
||||
|
||||
# Try to find model in database
|
||||
for model_name in model_name_candidates:
|
||||
try:
|
||||
db_model = await ModelRepository(prisma_client).table.find_first(where={"model_name": model_name})
|
||||
|
||||
if db_model and db_model.litellm_params:
|
||||
# Extract litellm_params (could be dict or JSON string)
|
||||
model_params = db_model.litellm_params
|
||||
if isinstance(model_params, str): # pyright: ignore[reportUnnecessaryIsInstance] # prisma Json is str
|
||||
model_params = json.loads(model_params)
|
||||
|
||||
# Decrypt values from database (similar to how proxy_server.py does it)
|
||||
# Values stored in DB are encrypted, so we need to decrypt them first
|
||||
decrypted_params = {}
|
||||
if isinstance(model_params, dict):
|
||||
for k, v in model_params.items():
|
||||
if isinstance(v, str):
|
||||
# Decrypt value - returns original value if decryption fails or no key is set
|
||||
decrypted_value = decrypt_value_helper(value=v, key=k, return_original_value=True)
|
||||
decrypted_params[k] = decrypted_value
|
||||
else:
|
||||
decrypted_params[k] = v
|
||||
else:
|
||||
decrypted_params = model_params
|
||||
|
||||
# Build embedding config from model params
|
||||
embedding_config = {}
|
||||
|
||||
# Extract api_key
|
||||
api_key = decrypted_params.get("api_key")
|
||||
if api_key:
|
||||
# Handle os.environ/ prefix (after decryption, values may be os.environ/ prefixed)
|
||||
if isinstance(api_key, str) and api_key.startswith("os.environ/"):
|
||||
api_key = get_secret(api_key)
|
||||
embedding_config["api_key"] = api_key
|
||||
|
||||
# Extract api_base
|
||||
api_base = decrypted_params.get("api_base")
|
||||
if api_base:
|
||||
# Handle os.environ/ prefix (after decryption, values may be os.environ/ prefixed)
|
||||
if isinstance(api_base, str) and api_base.startswith("os.environ/"):
|
||||
api_base = get_secret(api_base)
|
||||
embedding_config["api_base"] = api_base
|
||||
|
||||
# Extract api_version
|
||||
api_version = decrypted_params.get("api_version")
|
||||
if api_version:
|
||||
embedding_config["api_version"] = api_version
|
||||
|
||||
# Only return config if we have at least api_key or api_base
|
||||
if embedding_config:
|
||||
verbose_proxy_logger.debug(
|
||||
"Resolved embedding config from database model %s: %s",
|
||||
model_name,
|
||||
list(embedding_config.keys()),
|
||||
)
|
||||
return embedding_config
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Error resolving embedding config for model %s: %s", model_name, e)
|
||||
continue
|
||||
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_embedding_config(
|
||||
embedding_model: str, prisma_client: "PrismaClient | None", llm_router: "Router | None" = None
|
||||
) -> dict[str, object] | None:
|
||||
"""
|
||||
Resolve embedding config from either router (config-defined) or database models.
|
||||
|
||||
This function first checks the router for config-defined models, then falls back
|
||||
to the database. This allows users to use models defined in either location.
|
||||
|
||||
Results are cached in process memory for ``_EMBEDDING_CONFIG_CACHE_TTL``
|
||||
seconds so the request-handling path doesn't hit the database on every
|
||||
vector-store call. Negative results (model not found) are intentionally
|
||||
not cached to avoid blocking a freshly-added model behind the TTL.
|
||||
|
||||
Args:
|
||||
embedding_model: The embedding model string (e.g., "text-embedding-ada-002" or "azure/text-embedding-3-large")
|
||||
prisma_client: The Prisma client instance
|
||||
llm_router: The LiteLLM router instance (optional, will be imported if not provided)
|
||||
|
||||
Returns:
|
||||
Dictionary with api_key, api_base, and api_version if model found, None otherwise
|
||||
"""
|
||||
if not embedding_model:
|
||||
return None
|
||||
|
||||
cache: Final = _get_embedding_config_cache()
|
||||
cached: Final = cache.get_cache(embedding_model)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
# Import llm_router if not provided
|
||||
if llm_router is None:
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
except ImportError:
|
||||
llm_router = None
|
||||
|
||||
# First try to resolve from router (config-defined models)
|
||||
if llm_router is not None:
|
||||
router_config = _resolve_embedding_config_from_router(embedding_model=embedding_model, llm_router=llm_router)
|
||||
if router_config:
|
||||
verbose_proxy_logger.debug("Resolved embedding config from router for model %s", embedding_model)
|
||||
cache.set_cache(embedding_model, router_config)
|
||||
return router_config
|
||||
|
||||
# Fall back to database
|
||||
if prisma_client is not None:
|
||||
db_config: Final = await _resolve_embedding_config_from_db(
|
||||
embedding_model=embedding_model, prisma_client=prisma_client
|
||||
)
|
||||
if db_config:
|
||||
verbose_proxy_logger.debug("Resolved embedding config from database for model %s", embedding_model)
|
||||
cache.set_cache(embedding_model, db_config)
|
||||
return db_config
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"Could not resolve embedding config for model %s from router or database", embedding_model
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
########################################################
|
||||
# Helper Functions
|
||||
########################################################
|
||||
|
|
@ -469,10 +212,9 @@ async def create_vector_store_in_db(
|
|||
# (``api_key``, ``api_base``, ``api_version``) into this row. That
|
||||
# exposed every env-stored embedding-model credential on the
|
||||
# ``/vector_store/{new,info,update,list}`` responses. Keep the user's
|
||||
# raw ``litellm_embedding_model`` reference; resolution now happens in
|
||||
# ``build_request_data_from_managed_vector_store``
|
||||
# at request-handling time so the cleartext config exists only in
|
||||
# per-request memory and never reaches the database.
|
||||
# raw ``litellm_embedding_model`` reference; each search embeds the
|
||||
# query through the router at request time, so the credentials stay
|
||||
# on the deployment and never reach the database.
|
||||
if litellm_params:
|
||||
litellm_params_dict: Final = GenericLiteLLMParams(**litellm_params).model_dump(exclude_none=True)
|
||||
data_to_create["litellm_params"] = safe_dumps(litellm_params_dict)
|
||||
|
|
@ -862,11 +604,9 @@ async def update_vector_store(
|
|||
|
||||
# Handle litellm_params if provided. As with the create path, the
|
||||
# embedding-config auto-resolve previously persisted cleartext
|
||||
# credentials into the row; resolution now happens at request-
|
||||
# handling time in
|
||||
# ``build_request_data_from_managed_vector_store``
|
||||
# so this row only ever stores the user-supplied
|
||||
# ``litellm_embedding_model`` reference.
|
||||
# credentials into the row; each search now embeds the query
|
||||
# through the router at request time, so this row only ever stores
|
||||
# the user-supplied ``litellm_embedding_model`` reference.
|
||||
if "litellm_params" in update_data:
|
||||
_input_litellm_params: Final[dict] = update_data.get("litellm_params", {}) or {}
|
||||
litellm_params_dict: Final = GenericLiteLLMParams(**_input_litellm_params).model_dump(exclude_none=True)
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ from typing_extensions import TypeIs
|
|||
|
||||
import litellm
|
||||
from litellm.constants import (
|
||||
EMPTY_MAPPING,
|
||||
LITELLM_MAX_STREAMING_DURATION_SECONDS,
|
||||
STREAM_SSE_DONE_STRING,
|
||||
)
|
||||
|
|
@ -273,6 +274,9 @@ class BaseResponsesAPIStreamingIterator:
|
|||
self._hidden_params["additional_headers"] = process_response_headers(
|
||||
self.response.headers or {}
|
||||
) # GUARANTEE OPENAI HEADERS IN RESPONSE
|
||||
self._raw_response_headers: Mapping[str, str] = MappingProxyType(
|
||||
dict(self.response.headers or {}) # mutable-ok: immediately frozen by MappingProxyType
|
||||
)
|
||||
|
||||
def _check_max_streaming_duration(self) -> None:
|
||||
"""Raise litellm.Timeout if the stream has exceeded LITELLM_MAX_STREAMING_DURATION_SECONDS."""
|
||||
|
|
@ -446,6 +450,7 @@ class BaseResponsesAPIStreamingIterator:
|
|||
except Exception:
|
||||
# Fallback to original if serialization fails
|
||||
pass
|
||||
self._restore_provider_response_headers(logging_response)
|
||||
|
||||
end_time: Final = datetime.now()
|
||||
if is_async:
|
||||
|
|
@ -480,6 +485,41 @@ class BaseResponsesAPIStreamingIterator:
|
|||
)
|
||||
self._run_post_success_hooks(end_time=end_time)
|
||||
|
||||
def _restore_provider_response_headers(self, logging_response: object) -> None:
|
||||
"""Re-apply the provider's response headers to the copy handed to logging callbacks.
|
||||
|
||||
``model_validate(model_dump())`` above drops pydantic private attributes, so the
|
||||
``_hidden_params`` the provider transform set on the nested response are lost. Returns early
|
||||
when that copy fell back to the original event, so logging-only state never lands on the
|
||||
object the caller is iterating.
|
||||
"""
|
||||
if logging_response is self.completed_response:
|
||||
return
|
||||
target: Final[object] = getattr(logging_response, "response", None)
|
||||
existing_hidden: Final[object] = getattr(target, "_hidden_params", None)
|
||||
if not isinstance(existing_hidden, Mapping):
|
||||
return
|
||||
existing: Final[Mapping[str, object]] = existing_hidden
|
||||
source_hidden: Final[object] = getattr(
|
||||
getattr(self.completed_response, "response", None), "_hidden_params", None
|
||||
)
|
||||
source: Final[Mapping[str, object]] = source_hidden if isinstance(source_hidden, Mapping) else EMPTY_MAPPING
|
||||
processed: Final[object] = source.get("additional_headers") or self._hidden_params.get("additional_headers")
|
||||
raw: Final[object] = source.get("headers") or self._raw_response_headers
|
||||
headers: Final[Mapping[str, object]] = processed if isinstance(processed, Mapping) else EMPTY_MAPPING
|
||||
raw_headers: Final[Mapping[str, object]] = raw if isinstance(raw, Mapping) else EMPTY_MAPPING
|
||||
# rebuild by value and let existing keys win: sharing the source dicts would alias what the proxy
|
||||
# splats into the client's HTTP headers, and copying non-header keys would carry response_cost
|
||||
setattr( # noqa: B010 # target is typed object here, so a plain attribute store does not type check
|
||||
target,
|
||||
"_hidden_params",
|
||||
{ # mutable-ok: the cost calculator writes optional_params into _hidden_params
|
||||
"additional_headers": {**headers}, # mutable-ok: fresh copy, logging callbacks may mutate it
|
||||
"headers": {**raw_headers}, # mutable-ok: fresh copy, logging callbacks may mutate it
|
||||
**existing,
|
||||
},
|
||||
)
|
||||
|
||||
def _handle_logging_completed_response(self):
|
||||
"""Base implementation - should be overridden by subclasses"""
|
||||
|
||||
|
|
|
|||
|
|
@ -85,6 +85,10 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
|
|||
mask_credentials_in_payload,
|
||||
mask_sensitive_structure,
|
||||
)
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
RouterVectorStoreEmbeddingExecutor,
|
||||
vector_store_request_metadata,
|
||||
)
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
|
||||
from litellm.router_strategy.least_busy import LeastBusyLoggingHandler
|
||||
|
|
@ -6481,11 +6485,24 @@ class Router:
|
|||
if custom_llm_provider and "custom_llm_provider" not in kwargs
|
||||
else MappingProxyType(kwargs)
|
||||
)
|
||||
if provider_kwargs.get("model"):
|
||||
return self._generic_api_call_with_fallbacks(original_function=original_function, **provider_kwargs)
|
||||
search_kwargs: Final = (
|
||||
MappingProxyType(
|
||||
{
|
||||
**provider_kwargs,
|
||||
"_direct_vector_store_embedding_executor": RouterVectorStoreEmbeddingExecutor(
|
||||
router=self,
|
||||
metadata=self._vector_store_request_metadata(kwargs),
|
||||
),
|
||||
}
|
||||
)
|
||||
if call_type == "vector_store_search"
|
||||
else provider_kwargs
|
||||
)
|
||||
if search_kwargs.get("model"):
|
||||
return self._generic_api_call_with_fallbacks(original_function=original_function, **search_kwargs)
|
||||
if call_type == "vector_store_search":
|
||||
return original_function(**MappingProxyType({**provider_kwargs, "router": self}))
|
||||
return original_function(**provider_kwargs)
|
||||
return original_function(**MappingProxyType({**search_kwargs, "router": self}))
|
||||
return original_function(**search_kwargs)
|
||||
|
||||
return vector_store_sync_wrapper
|
||||
|
||||
|
|
@ -6658,11 +6675,22 @@ class Router:
|
|||
"avector_store_update",
|
||||
"avector_store_delete",
|
||||
):
|
||||
vector_store_kwargs: Final = (
|
||||
{ # mutable-ok: the async routed request requires dynamic keyword arguments
|
||||
**kwargs,
|
||||
"_direct_vector_store_embedding_executor": RouterVectorStoreEmbeddingExecutor(
|
||||
router=self,
|
||||
metadata=self._vector_store_request_metadata(kwargs),
|
||||
),
|
||||
}
|
||||
if call_type == "avector_store_search"
|
||||
else kwargs
|
||||
)
|
||||
return await self._init_vector_store_api_endpoints(
|
||||
original_function=original_function,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
call_type=call_type,
|
||||
**kwargs,
|
||||
**vector_store_kwargs,
|
||||
)
|
||||
elif call_type in ("afile_delete", "afile_content"):
|
||||
return await self._ageneric_api_call_with_fallbacks(
|
||||
|
|
@ -6698,6 +6726,10 @@ class Router:
|
|||
|
||||
return async_wrapper
|
||||
|
||||
@staticmethod
|
||||
def _vector_store_request_metadata(kwargs: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return vector_store_request_metadata(kwargs)
|
||||
|
||||
async def _init_vector_store_api_endpoints(
|
||||
self,
|
||||
original_function: Callable,
|
||||
|
|
|
|||
|
|
@ -4,7 +4,18 @@ from .base import GuardrailConfigModel
|
|||
|
||||
|
||||
class CrowdStrikeAIDRGuardrailConfigModelOptionalParams(BaseModel):
|
||||
pass
|
||||
streaming_end_of_stream_only: bool | None = Field(
|
||||
default=None,
|
||||
description="If False (default when unset), post_call scans the accumulated streamed response every "
|
||||
"streaming_sampling_rate chunks and an in-flight block stops the stream. If True, the guard runs once "
|
||||
"over the assembled response at end of stream, so flagged content may already have reached the client.",
|
||||
)
|
||||
streaming_sampling_rate: int | None = Field(
|
||||
default=None,
|
||||
ge=1,
|
||||
description="When streaming_end_of_stream_only is False, scan the accumulated streamed response every Nth "
|
||||
"chunk. Defaults to 5 when unset.",
|
||||
)
|
||||
|
||||
|
||||
class CrowdStrikeAIDRGuardrailConfigModel(GuardrailConfigModel[CrowdStrikeAIDRGuardrailConfigModelOptionalParams]):
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
|
||||
class PublicModelHubInfo(BaseModel):
|
||||
|
|
@ -73,6 +73,44 @@ class SupportedEndpointsResponse(BaseModel):
|
|||
endpoints: list[SupportedEndpoint]
|
||||
|
||||
|
||||
class AutoRouterPresetTiers(BaseModel):
|
||||
"""Exactly the four built-in tiers the dashboard's preset prefill can apply.
|
||||
|
||||
extra="forbid" on purpose: a tier name this dashboard cannot apply would grey out or crash the
|
||||
picker, so such a catalog is rejected wholesale and the bundled one serves instead.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
SIMPLE: Sequence[str]
|
||||
MEDIUM: Sequence[str]
|
||||
COMPLEX: Sequence[str]
|
||||
REASONING: Sequence[str]
|
||||
|
||||
|
||||
class AutoRouterPresetConfig(BaseModel):
|
||||
"""The complexity_router_config a preset prefills.
|
||||
|
||||
Only tiers is validated, because every dashboard consumer dereferences it; everything else
|
||||
passes through verbatim with unknown fields kept (extra="allow"), so a catalog published after
|
||||
this proxy shipped still serves its new fields intact.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
tiers: AutoRouterPresetTiers
|
||||
|
||||
|
||||
class AutoRouterPresetRecord(BaseModel):
|
||||
"""One auto-router preset as served to the dashboard's template picker."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
label: str
|
||||
description: str
|
||||
complexity_router_config: AutoRouterPresetConfig
|
||||
|
||||
|
||||
class ComplexityScorerDefaults(BaseModel):
|
||||
"""The complexity router's shipped heuristic scorer defaults.
|
||||
|
||||
|
|
|
|||
|
|
@ -15,6 +15,11 @@ import litellm
|
|||
from litellm.constants import request_timeout
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
BaseQueryEmbeddingVectorStoreConfig,
|
||||
VectorStoreEmbeddingExecutor,
|
||||
vector_store_request_metadata,
|
||||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.vector_stores import (
|
||||
|
|
@ -38,6 +43,16 @@ base_llm_http_handler = BaseLLMHTTPHandler()
|
|||
#################################################
|
||||
|
||||
|
||||
def _direct_vector_store_embedding_executor(
|
||||
value: object, router: "Router | None", request_kwargs: Mapping[str, object]
|
||||
) -> VectorStoreEmbeddingExecutor:
|
||||
if value is not None and not isinstance(value, VectorStoreEmbeddingExecutor):
|
||||
raise TypeError("Invalid direct vector store embedding executor")
|
||||
return BaseQueryEmbeddingVectorStoreConfig.query_embedding_executor(
|
||||
value, router, vector_store_request_metadata(request_kwargs)
|
||||
)
|
||||
|
||||
|
||||
def mock_vector_store_search_response(
|
||||
mock_results: list[VectorStoreSearchResult] | None = None,
|
||||
):
|
||||
|
|
@ -289,7 +304,12 @@ async def asearch(
|
|||
"""
|
||||
Async: Search a vector store for relevant chunks based on a query and file attributes filter.
|
||||
"""
|
||||
local_vars: Final = locals()
|
||||
embedding_executor: Final = _direct_vector_store_embedding_executor(
|
||||
kwargs.pop("_direct_vector_store_embedding_executor", None), router, kwargs
|
||||
)
|
||||
local_vars: Final = { # mutable-ok: exception logging requires a sanitized mutable snapshot
|
||||
key: value for key, value in locals().items() if key != "embedding_executor"
|
||||
}
|
||||
|
||||
try:
|
||||
loop: Final = asyncio.get_event_loop()
|
||||
|
|
@ -312,6 +332,7 @@ async def asearch(
|
|||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
_direct_vector_store_embedding_executor=embedding_executor,
|
||||
router=router,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -369,12 +390,16 @@ def search(
|
|||
Returns:
|
||||
VectorStoreSearchResponse containing the search results.
|
||||
"""
|
||||
local_vars: Final = locals()
|
||||
embedding_executor: Final = _direct_vector_store_embedding_executor(
|
||||
kwargs.pop("_direct_vector_store_embedding_executor", None), router, kwargs
|
||||
)
|
||||
local_vars: Final = { # mutable-ok: exception logging requires a sanitized mutable snapshot
|
||||
key: value for key, value in locals().items() if key != "embedding_executor"
|
||||
}
|
||||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("asearch", False) is True
|
||||
|
||||
# pull credentials from registry if available
|
||||
if litellm.vector_store_registry is not None and vector_store_id is not None:
|
||||
try:
|
||||
|
|
@ -451,6 +476,7 @@ def search(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
embedding_executor=embedding_executor,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout or request_timeout,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 2981
|
||||
"limit": 2979
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 71
|
||||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 3
|
||||
},
|
||||
"BLE001": {
|
||||
"limit": 2915
|
||||
"limit": 2914
|
||||
},
|
||||
"C401": {
|
||||
"limit": 8
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ export const E2E_PROXY_ADMIN_EMAIL = "admin@test.local";
|
|||
export const E2E_INTERNAL_USER_ID = "e2e-internal-user";
|
||||
export const E2E_INTERNAL_USER_EMAIL = "internal@test.local";
|
||||
export const E2E_TEAM_ADMIN_USER_ID = "e2e-team-admin";
|
||||
export const E2E_SEEDED_USER_PASSWORD = "E2e-Test-Pass-2026!";
|
||||
|
||||
// Key aliases for seeded test keys (match seed.sql)
|
||||
export const E2E_UPDATE_LIMITS_KEY_ALIAS = "e2eUpdateLimitsKey";
|
||||
|
|
|
|||
|
|
@ -24,18 +24,18 @@ INSERT INTO "LiteLLM_OrganizationTable" (
|
|||
'e2e-proxy-admin', 'e2e-proxy-admin'
|
||||
);
|
||||
|
||||
-- 4. Users (password hash is scrypt of "test")
|
||||
-- 4. Users (password hash is scrypt of E2E_SEEDED_USER_PASSWORD from constants.ts)
|
||||
INSERT INTO "LiteLLM_UserTable" ("user_id", "user_email", "user_role", "teams", "password")
|
||||
VALUES
|
||||
('e2e-proxy-admin', 'admin@test.local', 'proxy_admin', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-admin-viewer', 'adminviewer@test.local', 'proxy_admin_viewer', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-internal-user', 'internal@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-org","e2e-team-keygen"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-internal-viewer', 'viewer@test.local', 'internal_user_viewer', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-team-admin', 'teamadmin@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-delete"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-invitable-user', 'invitable@test.local', 'internal_user', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-internal-noteam', 'noteam@test.local', 'internal_user', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-invitable-by-team-admin', 'invitable-team@test.local', 'internal_user', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
|
||||
('e2e-removable-member', 'removable@test.local', 'internal_user', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr');
|
||||
('e2e-proxy-admin', 'admin@test.local', 'proxy_admin', '{"e2e-team-crud"}', 'scrypt:KdnTJwPb3gswdqSznPE5CC6apeFIMycd6BG7yRWndZa3QZPcVs37y7jvrQCaPUNq'),
|
||||
('e2e-admin-viewer', 'adminviewer@test.local', 'proxy_admin_viewer', '{}', 'scrypt:KdnTJwPb3gswdqSznPE5CC6apeFIMycd6BG7yRWndZa3QZPcVs37y7jvrQCaPUNq'),
|
||||
('e2e-internal-user', 'internal@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-org","e2e-team-keygen"}', 'scrypt:KdnTJwPb3gswdqSznPE5CC6apeFIMycd6BG7yRWndZa3QZPcVs37y7jvrQCaPUNq'),
|
||||
('e2e-internal-viewer', 'viewer@test.local', 'internal_user_viewer', '{"e2e-team-crud"}', 'scrypt:KdnTJwPb3gswdqSznPE5CC6apeFIMycd6BG7yRWndZa3QZPcVs37y7jvrQCaPUNq'),
|
||||
('e2e-team-admin', 'teamadmin@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-delete"}', 'scrypt:KdnTJwPb3gswdqSznPE5CC6apeFIMycd6BG7yRWndZa3QZPcVs37y7jvrQCaPUNq'),
|
||||
('e2e-invitable-user', 'invitable@test.local', 'internal_user', '{}', 'scrypt:KdnTJwPb3gswdqSznPE5CC6apeFIMycd6BG7yRWndZa3QZPcVs37y7jvrQCaPUNq'),
|
||||
('e2e-internal-noteam', 'noteam@test.local', 'internal_user', '{}', 'scrypt:KdnTJwPb3gswdqSznPE5CC6apeFIMycd6BG7yRWndZa3QZPcVs37y7jvrQCaPUNq'),
|
||||
('e2e-invitable-by-team-admin', 'invitable-team@test.local', 'internal_user', '{}', 'scrypt:KdnTJwPb3gswdqSznPE5CC6apeFIMycd6BG7yRWndZa3QZPcVs37y7jvrQCaPUNq'),
|
||||
('e2e-removable-member', 'removable@test.local', 'internal_user', '{"e2e-team-crud"}', 'scrypt:KdnTJwPb3gswdqSznPE5CC6apeFIMycd6BG7yRWndZa3QZPcVs37y7jvrQCaPUNq');
|
||||
|
||||
-- 5. Teams (members_with_roles is required JSON)
|
||||
INSERT INTO "LiteLLM_TeamTable" (
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import {
|
||||
ADMIN_STORAGE_PATH,
|
||||
ADMIN_VIEWER_STORAGE_PATH,
|
||||
E2E_SEEDED_USER_PASSWORD,
|
||||
INTERNAL_USER_STORAGE_PATH,
|
||||
INTERNAL_VIEWER_STORAGE_PATH,
|
||||
TEAM_ADMIN_STORAGE_PATH,
|
||||
|
|
@ -23,22 +24,22 @@ export const users: Record<Role, { email: string; password: string; seedApiRole?
|
|||
},
|
||||
[Role.ProxyAdminViewer]: {
|
||||
email: "adminviewer@test.local",
|
||||
password: "test",
|
||||
password: E2E_SEEDED_USER_PASSWORD,
|
||||
seedApiRole: "proxy_admin_viewer",
|
||||
},
|
||||
[Role.InternalUser]: {
|
||||
email: "internal@test.local",
|
||||
password: "test",
|
||||
password: E2E_SEEDED_USER_PASSWORD,
|
||||
seedApiRole: "internal_user",
|
||||
},
|
||||
[Role.InternalUserViewer]: {
|
||||
email: "viewer@test.local",
|
||||
password: "test",
|
||||
password: E2E_SEEDED_USER_PASSWORD,
|
||||
seedApiRole: "internal_user_viewer",
|
||||
},
|
||||
[Role.TeamAdmin]: {
|
||||
email: "teamadmin@test.local",
|
||||
password: "test",
|
||||
password: E2E_SEEDED_USER_PASSWORD,
|
||||
seedApiRole: "internal_user",
|
||||
},
|
||||
};
|
||||
|
|
|
|||
|
|
@ -21,6 +21,8 @@ interface ChatOptions {
|
|||
apiKey?: string;
|
||||
/** Sent as `user`, which lands in the spend log's end_user column. */
|
||||
endUser?: string;
|
||||
/** Sent as `litellm_trace_id`, which lands in the spend log's session_id column. */
|
||||
traceId?: string;
|
||||
}
|
||||
|
||||
/** POST /v1/chat/completions and return the completion id (the Logs Request ID). */
|
||||
|
|
@ -34,6 +36,7 @@ export async function sendChatCompletion(request: APIRequestContext, opts: ChatO
|
|||
model: opts.model,
|
||||
messages: [{ role: "user", content: opts.prompt }],
|
||||
...(opts.endUser ? { user: opts.endUser } : {}),
|
||||
...(opts.traceId ? { litellm_trace_id: opts.traceId } : {}),
|
||||
},
|
||||
});
|
||||
expect(res.ok(), `chat completion for ${opts.model} failed (${res.status()}): ${await res.text()}`).toBe(true);
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import { test, expect } from "@playwright/test";
|
||||
import { navigateToPage } from "../../helpers/navigation";
|
||||
import { Page } from "../../fixtures/pages";
|
||||
import { E2E_SEEDED_USER_PASSWORD } from "../../constants";
|
||||
|
||||
/**
|
||||
* Logs in fresh inside the test rather than reusing a stored session because
|
||||
|
|
@ -15,7 +16,7 @@ test.describe("Internal User with no team memberships", () => {
|
|||
// Log in via the form as the no-team seeded user.
|
||||
await page.goto("/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill("noteam@test.local");
|
||||
await page.getByPlaceholder("Enter your password").fill("test");
|
||||
await page.getByPlaceholder("Enter your password").fill(E2E_SEEDED_USER_PASSWORD);
|
||||
await page.getByRole("button", { name: "Login", exact: true }).click();
|
||||
await expect(page.getByRole("complementary").getByText("Virtual Keys")).toBeVisible({ timeout: 30_000 });
|
||||
expect(new URL(page.url()).pathname).not.toMatch(/\/connect$/);
|
||||
|
|
|
|||
150
tests/e2e/ui/tests/logs/logsPagination.spec.ts
Normal file
150
tests/e2e/ui/tests/logs/logsPagination.spec.ts
Normal file
|
|
@ -0,0 +1,150 @@
|
|||
import { test, expect, type APIRequestContext, type Locator, type Page as PlaywrightPage } from "@playwright/test";
|
||||
import { ADMIN_STORAGE_PATH } from "../../constants";
|
||||
import { navigateToPage, dismissFeedbackPopup } from "../../helpers/navigation";
|
||||
import { Page } from "../../fixtures/pages";
|
||||
import { CHAT_MODEL_A, createVirtualKey, sendChatCompletion, waitForSpendLog } from "../../helpers/traffic";
|
||||
|
||||
/**
|
||||
* Session-grouped pagination (#38060): a page of N rows must render exactly N session rows, a
|
||||
* session must never straddle pages, and two callers reusing one session id stay separate rows.
|
||||
* All traffic is generated per run behind a unique key alias or session id, so concurrent specs
|
||||
* cannot decide the outcome.
|
||||
*/
|
||||
|
||||
const uniqueSuffix = (): string => `${Date.now()}-${Math.random().toString(36).slice(2, 8)}`;
|
||||
|
||||
/** Every tab stays mounted, so the DOM holds four tables at once; scope to the visible one. */
|
||||
const requestLogsRows = (page: PlaywrightPage): Locator =>
|
||||
page.locator("table").filter({ visible: true }).first().locator("tbody tr");
|
||||
|
||||
const visibleTestId = (page: PlaywrightPage, id: string): Locator => page.getByTestId(id).filter({ visible: true });
|
||||
|
||||
async function openLogs(page: PlaywrightPage): Promise<void> {
|
||||
await navigateToPage(page, Page.Logs);
|
||||
await dismissFeedbackPopup(page);
|
||||
await expect(visibleTestId(page, "datatable-search")).toBeVisible({ timeout: 20_000 });
|
||||
}
|
||||
|
||||
async function openFilterDrawer(page: PlaywrightPage): Promise<Locator> {
|
||||
await visibleTestId(page, "datatable-filters-trigger").click();
|
||||
const drawer = page.getByRole("dialog", { name: "Filters" });
|
||||
await expect(drawer).toBeVisible({ timeout: 10_000 });
|
||||
return drawer;
|
||||
}
|
||||
|
||||
async function applyKeyAliasFilter(page: PlaywrightPage, drawer: Locator, alias: string): Promise<void> {
|
||||
await drawer.getByRole("combobox", { name: "Search a key alias" }).click();
|
||||
await page.keyboard.type(alias);
|
||||
await page.getByRole("option", { name: alias, exact: true }).first().click();
|
||||
await drawer.getByRole("button", { name: "Apply Filters" }).click();
|
||||
await expect(drawer).not.toBeVisible({ timeout: 10_000 });
|
||||
}
|
||||
|
||||
async function setRowsPerPage(page: PlaywrightPage, size: "25" | "50" | "100"): Promise<void> {
|
||||
await visibleTestId(page, "pagination-page-size").click();
|
||||
await page.getByRole("option", { name: size, exact: true }).click();
|
||||
}
|
||||
|
||||
test.describe("Logs page session-grouped pagination", () => {
|
||||
test.use({ storageState: ADMIN_STORAGE_PATH });
|
||||
|
||||
test("a 25-row page renders exactly 25 session rows and no session straddles pages", async ({ page, request }) => {
|
||||
const suffix = uniqueSuffix();
|
||||
const alias = `e2e-logs-pgn-${suffix}`;
|
||||
const mine = await createVirtualKey(request, { key_alias: alias });
|
||||
|
||||
const soloIds: string[] = [];
|
||||
for (let i = 0; i < 26; i++) {
|
||||
soloIds.push(
|
||||
await sendChatCompletion(request, {
|
||||
model: CHAT_MODEL_A,
|
||||
prompt: `logs-pgn-solo-${i}-${suffix}`,
|
||||
apiKey: mine.key,
|
||||
}),
|
||||
);
|
||||
}
|
||||
const sessionA = `sess-pgn-a-${suffix}`;
|
||||
const sessionB = `sess-pgn-b-${suffix}`;
|
||||
let lastSessionCallId = "";
|
||||
for (let i = 0; i < 7; i++) {
|
||||
lastSessionCallId = await sendChatCompletion(request, {
|
||||
model: CHAT_MODEL_A,
|
||||
prompt: `logs-pgn-a-${i}-${suffix}`,
|
||||
apiKey: mine.key,
|
||||
traceId: sessionA,
|
||||
});
|
||||
}
|
||||
for (let i = 0; i < 3; i++) {
|
||||
lastSessionCallId = await sendChatCompletion(request, {
|
||||
model: CHAT_MODEL_A,
|
||||
prompt: `logs-pgn-b-${i}-${suffix}`,
|
||||
apiKey: mine.key,
|
||||
traceId: sessionB,
|
||||
});
|
||||
}
|
||||
await waitForSpendLog(request, lastSessionCallId);
|
||||
await waitForSpendLog(request, soloIds[soloIds.length - 1]);
|
||||
|
||||
// 36 calls in 28 session groups: 26 solos plus sessions of 7 and 3.
|
||||
await openLogs(page);
|
||||
const drawer = await openFilterDrawer(page);
|
||||
await applyKeyAliasFilter(page, drawer, alias);
|
||||
await setRowsPerPage(page, "25");
|
||||
|
||||
await expect(visibleTestId(page, "pagination-range")).toHaveText("Showing 1-25 of 28", { timeout: 30_000 });
|
||||
await expect(requestLogsRows(page)).toHaveCount(25);
|
||||
// The sessions are the newest groups, so their single representative rows sit on page 1.
|
||||
await expect(requestLogsRows(page).filter({ hasText: sessionA })).toHaveCount(1);
|
||||
await expect(requestLogsRows(page).filter({ hasText: sessionA })).toContainText("7");
|
||||
await expect(requestLogsRows(page).filter({ hasText: sessionB })).toHaveCount(1);
|
||||
|
||||
await visibleTestId(page, "pagination-next").click();
|
||||
|
||||
await expect(visibleTestId(page, "pagination-range")).toHaveText("Showing 26-28 of 28", { timeout: 30_000 });
|
||||
await expect(requestLogsRows(page)).toHaveCount(3);
|
||||
await expect(requestLogsRows(page).filter({ hasText: sessionA })).toHaveCount(0);
|
||||
await expect(requestLogsRows(page).filter({ hasText: sessionB })).toHaveCount(0);
|
||||
});
|
||||
|
||||
test("two keys reusing one session id stay separate rows", async ({ page, request }) => {
|
||||
const suffix = uniqueSuffix();
|
||||
const mine = await createVirtualKey(request, { key_alias: `e2e-logs-pgn-mine-${suffix}` });
|
||||
const theirs = await createVirtualKey(request, { key_alias: `e2e-logs-pgn-theirs-${suffix}` });
|
||||
const sharedSession = `sess-pgn-shared-${suffix}`;
|
||||
|
||||
let lastId = "";
|
||||
for (let i = 0; i < 2; i++) {
|
||||
lastId = await sendChatCompletion(request, {
|
||||
model: CHAT_MODEL_A,
|
||||
prompt: `logs-pgn-shared-mine-${i}-${suffix}`,
|
||||
apiKey: mine.key,
|
||||
traceId: sharedSession,
|
||||
});
|
||||
}
|
||||
lastId = await sendChatCompletion(request, {
|
||||
model: CHAT_MODEL_A,
|
||||
prompt: `logs-pgn-shared-theirs-${suffix}`,
|
||||
apiKey: theirs.key,
|
||||
traceId: sharedSession,
|
||||
});
|
||||
await waitForSpendLog(request, lastId);
|
||||
|
||||
await openLogs(page);
|
||||
const drawer = await openFilterDrawer(page);
|
||||
await drawer.getByPlaceholder("Enter session ID…").fill(sharedSession);
|
||||
await drawer.getByRole("button", { name: "Apply Filters" }).click();
|
||||
await expect(drawer).not.toBeVisible({ timeout: 10_000 });
|
||||
|
||||
// One row per caller: reusing a session id must not merge two keys' activity into one row.
|
||||
await expect(requestLogsRows(page).filter({ hasText: sharedSession })).toHaveCount(2, { timeout: 30_000 });
|
||||
|
||||
// And each row carries ITS key's totals: two calls badge the first key's row,
|
||||
// while the other key's single call renders as a plain LLM row.
|
||||
const mineRow = requestLogsRows(page).filter({ hasText: sharedSession }).filter({ hasText: mine.token });
|
||||
const theirsRow = requestLogsRows(page).filter({ hasText: sharedSession }).filter({ hasText: theirs.token });
|
||||
await expect(mineRow).toHaveCount(1);
|
||||
await expect(theirsRow).toHaveCount(1);
|
||||
await expect(mineRow.getByText("2", { exact: true })).toBeVisible();
|
||||
await expect(theirsRow.getByText("LLM", { exact: true })).toBeVisible();
|
||||
});
|
||||
});
|
||||
|
|
@ -10,7 +10,7 @@ test.describe("Second proxy admin", () => {
|
|||
test("an invited admin can log in, mint a key, and call a model with it", async ({ page, browser, request }) => {
|
||||
const suffix = Date.now();
|
||||
const email = `second-admin-${suffix}@test.local`;
|
||||
const password = "e2e-second-admin-password";
|
||||
const password = "E2e-Second-Admin-Pass-1!";
|
||||
const auth = { Authorization: `Bearer ${masterKey()}` };
|
||||
|
||||
const inviteAdminUser = async (): Promise<string> => {
|
||||
|
|
|
|||
|
|
@ -71,6 +71,48 @@ def setup_vector_store_registry():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_hook_routes_search_through_proxy_router(
|
||||
setup_vector_store_registry,
|
||||
):
|
||||
proxy_router = Mock()
|
||||
proxy_router.avector_store_search = AsyncMock(
|
||||
return_value=VectorStoreSearchResponse(
|
||||
object="vector_store.search_results.page",
|
||||
search_query="what is litellm?",
|
||||
data=[
|
||||
VectorStoreSearchResult(
|
||||
score=1.0,
|
||||
content=[VectorStoreResultContent(text="routed context", type="text")],
|
||||
)
|
||||
],
|
||||
)
|
||||
)
|
||||
logging_obj = Mock()
|
||||
logging_obj.model_call_details = {
|
||||
"litellm_params": {"metadata": {"user_api_key_team_id": "team-a"}}
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router", proxy_router):
|
||||
_, messages, _ = await VectorStorePreCallHook().async_get_chat_completion_prompt(
|
||||
model="chat-model",
|
||||
messages=[{"role": "user", "content": "what is litellm?"}],
|
||||
non_default_params={"vector_store_ids": ["T37J8R4WTM"]},
|
||||
prompt_id=None,
|
||||
prompt_variables=None,
|
||||
dynamic_callback_params={},
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
proxy_router.avector_store_search.assert_awaited_once_with(
|
||||
vector_store_id="T37J8R4WTM",
|
||||
query="what is litellm?",
|
||||
custom_llm_provider="bedrock",
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
assert messages[0]["content"] == "Context:\n\nrouted context\n\n"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_e2e_bedrock_knowledgebase_retrieval_with_completion(
|
||||
setup_vector_store_registry,
|
||||
|
|
|
|||
|
|
@ -5,17 +5,218 @@ These tests simulate real-world scenarios where headers and configuration
|
|||
need to be properly propagated through the router to the LLM API.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock, patch, AsyncMock
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
LiteLLMVectorStoreEmbeddingExecutor,
|
||||
RouterVectorStoreEmbeddingExecutor,
|
||||
)
|
||||
|
||||
QUERY_VECTOR = [0.5, -0.25, 0.125]
|
||||
OPENAI_EMBEDDINGS_URL = "https://api.openai.com/v1/embeddings"
|
||||
STORE_EMBEDDINGS_URL = "https://embedding.example/v1/embeddings"
|
||||
|
||||
|
||||
def _mock_embedding_route(respx_mock: respx.MockRouter, url: str) -> respx.Route:
|
||||
return respx_mock.post(url).mock(
|
||||
return_value=httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": QUERY_VECTOR}],
|
||||
"model": "text-embedding-3-small",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _sent(route: respx.Route, index: int) -> tuple[str, str, list[str]]:
|
||||
request = route.calls[index].request
|
||||
body = json.loads(request.read())
|
||||
return request.headers["authorization"], body["model"], body["input"]
|
||||
|
||||
|
||||
def _alias_router() -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "team-alias",
|
||||
"litellm_params": {
|
||||
"model": "openai/text-embedding-3-small",
|
||||
"api_key": "deployment-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
class TestRouterEmbeddingIntegration:
|
||||
"""Integration tests for embedding with router configuration."""
|
||||
|
||||
def test_vector_store_request_metadata_prefers_litellm_metadata(self):
|
||||
assert Router._vector_store_request_metadata(
|
||||
{
|
||||
"litellm_metadata": {"user_api_key_team_id": "team-a"},
|
||||
"metadata": {"user_api_key_team_id": "team-b"},
|
||||
}
|
||||
) == {"user_api_key_team_id": "team-a"}
|
||||
|
||||
assert Router._vector_store_request_metadata({"metadata": {"user_api_key_team_id": "team-b"}}) == {
|
||||
"user_api_key_team_id": "team-b"
|
||||
}
|
||||
assert Router._vector_store_request_metadata({}) == {}
|
||||
|
||||
def test_sync_vector_store_wrapper_injects_router_embedding_executor(self):
|
||||
router = Router(model_list=[])
|
||||
original = MagicMock(return_value="searched")
|
||||
wrapped = router.factory_function(original, call_type="vector_store_search")
|
||||
|
||||
assert (
|
||||
wrapped(
|
||||
vector_store_id="store",
|
||||
query="query",
|
||||
custom_llm_provider="valkey",
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
== "searched"
|
||||
)
|
||||
|
||||
call_kwargs = original.call_args.kwargs
|
||||
assert call_kwargs["custom_llm_provider"] == "valkey"
|
||||
executor = call_kwargs["_direct_vector_store_embedding_executor"]
|
||||
assert isinstance(executor, RouterVectorStoreEmbeddingExecutor)
|
||||
assert executor.metadata == {"user_api_key_team_id": "team-a"}
|
||||
|
||||
def test_sync_vector_store_wrapper_preserves_model_routing(self):
|
||||
router = Router(model_list=[])
|
||||
original = MagicMock()
|
||||
wrapped = router.factory_function(original, call_type="vector_store_search")
|
||||
|
||||
with patch.object(router, "_generic_api_call_with_fallbacks", return_value="routed") as fallback:
|
||||
assert wrapped(model="vector-alias", vector_store_id="store", query="query") == "routed"
|
||||
|
||||
assert fallback.call_args.kwargs["model"] == "vector-alias"
|
||||
assert fallback.call_args.kwargs["original_function"] is original
|
||||
assert isinstance(
|
||||
fallback.call_args.kwargs["_direct_vector_store_embedding_executor"],
|
||||
RouterVectorStoreEmbeddingExecutor,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_embedding_executors_cover_sdk_and_router_paths(
|
||||
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
openai_route = _mock_embedding_route(respx_mock, OPENAI_EMBEDDINGS_URL)
|
||||
store_route = _mock_embedding_route(respx_mock, STORE_EMBEDDINGS_URL)
|
||||
sdk_executor = LiteLLMVectorStoreEmbeddingExecutor()
|
||||
|
||||
sync_response = sdk_executor.embed("openai/text-embedding-3-small", "sync", {"api_key": "explicit"})
|
||||
async_response = await sdk_executor.aembed("openai/text-embedding-3-small", "async", {"api_key": "explicit"})
|
||||
|
||||
assert sync_response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert async_response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert _sent(openai_route, 0) == ("Bearer explicit", "text-embedding-3-small", ["sync"])
|
||||
assert _sent(openai_route, 1) == ("Bearer explicit", "text-embedding-3-small", ["async"])
|
||||
|
||||
explicit_config = {
|
||||
"api_base": "https://embedding.example/v1",
|
||||
"api_key": "store-key",
|
||||
"metadata": {
|
||||
"configured": True,
|
||||
"user_api_key_team_id": "untrusted-team",
|
||||
},
|
||||
"model": "untrusted-model",
|
||||
}
|
||||
mock_router = MagicMock()
|
||||
mock_router.embedding.return_value = sync_response
|
||||
router_executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=mock_router,
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
assert router_executor.embed("team-alias", "query", explicit_config) is sync_response
|
||||
mock_router.embedding.assert_called_once_with(
|
||||
model="team-alias",
|
||||
input=["query"],
|
||||
api_base="https://embedding.example/v1",
|
||||
api_key="store-key",
|
||||
metadata={"configured": True, "user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
alias_executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=_alias_router(),
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
sync_alias = alias_executor.embed("team-alias", "sync query", explicit_config)
|
||||
async_alias = await alias_executor.aembed("team-alias", "async query", explicit_config)
|
||||
|
||||
assert sync_alias.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert async_alias.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert openai_route.call_count == 2
|
||||
assert _sent(store_route, 0) == ("Bearer store-key", "text-embedding-3-small", ["sync query"])
|
||||
assert _sent(store_route, 1) == ("Bearer store-key", "text-embedding-3-small", ["async query"])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_executor_falls_back_to_sdk_for_models_the_router_does_not_serve(
|
||||
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
store_route = _mock_embedding_route(respx_mock, STORE_EMBEDDINGS_URL)
|
||||
executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=_alias_router(),
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
inline_config = {"api_base": "https://embedding.example/v1", "api_key": "store-key"}
|
||||
|
||||
sync_response = executor.embed("openai/text-embedding-3-large", "sync query", inline_config)
|
||||
async_response = await executor.aembed("openai/text-embedding-3-large", "async query", inline_config)
|
||||
|
||||
assert sync_response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert async_response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert _sent(store_route, 0) == ("Bearer store-key", "text-embedding-3-large", ["sync query"])
|
||||
assert _sent(store_route, 1) == ("Bearer store-key", "text-embedding-3-large", ["async query"])
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_executor_rejects_unserved_models_without_explicit_config(
|
||||
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "env-key")
|
||||
openai_route = _mock_embedding_route(respx_mock, OPENAI_EMBEDDINGS_URL)
|
||||
executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=_alias_router(),
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError):
|
||||
executor.embed("openai/text-embedding-3-large", "sync query", {})
|
||||
with pytest.raises(litellm.BadRequestError):
|
||||
await executor.aembed("openai/text-embedding-3-large", "async query", {})
|
||||
|
||||
assert openai_route.call_count == 0
|
||||
|
||||
def test_router_executor_routes_deployment_model_names_through_the_router(
|
||||
self, respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
openai_route = _mock_embedding_route(respx_mock, OPENAI_EMBEDDINGS_URL)
|
||||
executor = RouterVectorStoreEmbeddingExecutor(router=_alias_router(), metadata={})
|
||||
|
||||
response = executor.embed("openai/text-embedding-3-small", "query", {})
|
||||
|
||||
assert response.data[0]["embedding"] == QUERY_VECTOR
|
||||
assert _sent(openai_route, 0) == ("Bearer deployment-key", "text-embedding-3-small", ["query"])
|
||||
|
||||
def test_embedding_with_deployment_specific_headers(self):
|
||||
"""
|
||||
Test that deployment-specific headers are propagated.
|
||||
|
|
@ -122,9 +323,7 @@ class TestRouterEmbeddingIntegration:
|
|||
|
||||
router = Router(
|
||||
model_list=model_list,
|
||||
default_litellm_params={
|
||||
"metadata": {"environment": "test", "service": "embedding-service"}
|
||||
},
|
||||
default_litellm_params={"metadata": {"environment": "test", "service": "embedding-service"}},
|
||||
)
|
||||
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
|
|
@ -240,9 +439,7 @@ class TestRouterEmbeddingIntegration:
|
|||
# Make multiple calls and verify headers are always present
|
||||
for i in range(5):
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MagicMock(
|
||||
data=[{"embedding": [0.1, 0.2]}]
|
||||
)
|
||||
mock_embedding.return_value = MagicMock(data=[{"embedding": [0.1, 0.2]}])
|
||||
|
||||
router.embedding(model="shared-embedding-model", input=[f"test {i}"])
|
||||
|
||||
|
|
@ -327,9 +524,7 @@ class TestRouterEmbeddingIntegration:
|
|||
|
||||
router = Router(
|
||||
model_list=model_list,
|
||||
default_litellm_params={
|
||||
"headers": {"X-Custom-Azure-Header": "azure-value"}
|
||||
},
|
||||
default_litellm_params={"headers": {"X-Custom-Azure-Header": "azure-value"}},
|
||||
)
|
||||
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
|
|
|
|||
44
tests/rust-python-harness/AGENTS.md
Normal file
44
tests/rust-python-harness/AGENTS.md
Normal file
|
|
@ -0,0 +1,44 @@
|
|||
# Expected Structure
|
||||
|
||||
```text
|
||||
tests/rust-python-harness/
|
||||
├── __main__.py
|
||||
│
|
||||
├── strategies/
|
||||
│ ├── e2e_parity/
|
||||
│ │ ├── runner.py
|
||||
│ │ ├── sdk/
|
||||
│ │ │ ├── ocr/
|
||||
│ │ │ ├── messages/
|
||||
│ │ │ ├── chat_completions/
|
||||
│ │ │ └── responses/
|
||||
│ │ └── gateway/
|
||||
│ │
|
||||
│ ├── trace_parity/
|
||||
│ │ ├── runner.py
|
||||
│ │ ├── sdk/
|
||||
│ │ └── gateway/
|
||||
│ │
|
||||
│ └── unit_tests/
|
||||
│ ├── runner.py
|
||||
│ ├── mapping_validator.py
|
||||
│ ├── python_runner.py
|
||||
│ └── rust_runner.py
|
||||
│
|
||||
└── shared/
|
||||
├── parity/
|
||||
├── tracing/
|
||||
└── reporting/
|
||||
```
|
||||
|
||||
- Run locally only; no CI integration
|
||||
- `__main__.py` selects strategies and combines their reports; each strategy also runs independently
|
||||
- `e2e_parity/` compares SDK objects, exceptions, callbacks, and streams, or gateway HTTP responses
|
||||
- `trace_parity/` compares mapped operations, call counts, and required execution ordering
|
||||
- E2E and trace runners share orchestration across `sdk/` and `gateway/`; surface-specific execution lives in those folders
|
||||
- `unit_tests/runner.py` combines mapping validation, Python test runs, and native Rust test runs
|
||||
- `mapping_validator.py` matches Python/Rust tests by agreed names or annotations and reports missing or ambiguous counterparts
|
||||
- `python_runner.py` runs existing Python tests with Rust disabled and enabled in separate processes, verifies backend selection, and compares results
|
||||
- `rust_runner.py` runs Cargo tests; native Rust unit tests stay beside their implementation
|
||||
- `shared/` contains reusable parity, tracing, and reporting machinery
|
||||
- Keep fixtures with their owning API and existing Python tests in their current locations
|
||||
|
|
@ -1013,6 +1013,48 @@ def test_an_unmapped_exception_with_no_model_or_provider_is_a_connection_error(q
|
|||
assert "boom" in raised.value.message
|
||||
|
||||
|
||||
def _raise_and_map(
|
||||
model: str | None, original_exception: Exception, custom_llm_provider: str | None
|
||||
) -> None:
|
||||
"""Calls exception_type() from inside the except block, as litellm/main.py does,
|
||||
so traceback.format_exc() has a real stack."""
|
||||
try:
|
||||
raise original_exception
|
||||
except type(original_exception) as caught:
|
||||
exception_type(
|
||||
model=model,
|
||||
original_exception=caught,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
|
||||
def test_an_unmapped_exception_message_keeps_traceback_for_sdk_callers(quiet_exception_mapping):
|
||||
"""Direct SDK callers debug unmapped provider exceptions with this traceback;
|
||||
only the proxy's response boundary strips it."""
|
||||
with pytest.raises(litellm.APIConnectionError) as raised:
|
||||
_raise_and_map(
|
||||
model="MiniMax-M2.5",
|
||||
original_exception=RuntimeError("socket hung up"),
|
||||
custom_llm_provider="minimax",
|
||||
)
|
||||
|
||||
assert "Traceback (most recent call last)" in raised.value.message
|
||||
assert "test_exception_mapping_utils.py" in raised.value.message
|
||||
|
||||
|
||||
def test_an_unmapped_exception_with_no_model_or_provider_message_keeps_traceback(
|
||||
quiet_exception_mapping,
|
||||
):
|
||||
with pytest.raises(litellm.APIConnectionError) as raised:
|
||||
_raise_and_map(
|
||||
model=None,
|
||||
original_exception=ValueError("boom"),
|
||||
custom_llm_provider=None,
|
||||
)
|
||||
|
||||
assert "Traceback (most recent call last)" in raised.value.message
|
||||
|
||||
|
||||
CONTEXT_WINDOW_MESSAGE = "This model's maximum context length is 4096 tokens."
|
||||
CONTENT_POLICY_MESSAGE = (
|
||||
'{"error": {"type": "invalid_request_error", "code": "content_policy_violation"}}'
|
||||
|
|
|
|||
|
|
@ -1314,3 +1314,236 @@ async def test_finalizer_on_live_loop_disposes_foreign_loop_session_without_sche
|
|||
|
||||
assert AsyncHTTPHandler._finalizer_close_tasks == baseline_tasks
|
||||
assert session.closed
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def forward_proxy_server():
|
||||
"""Plain HTTP forward proxy that records the absolute URIs it is asked to fetch."""
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from socketserver import ThreadingMixIn
|
||||
|
||||
seen_uris: list[str] = []
|
||||
|
||||
class RecordingProxyHandler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def do_GET(self):
|
||||
seen_uris.append(self.path)
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Length", "9")
|
||||
self.end_headers()
|
||||
self.wfile.write(b"via-proxy")
|
||||
|
||||
def log_message(self, format, *args):
|
||||
pass
|
||||
|
||||
class ThreadedServer(ThreadingMixIn, HTTPServer):
|
||||
daemon_threads = True
|
||||
|
||||
server = ThreadedServer(("127.0.0.1", 0), RecordingProxyHandler)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
yield f"http://127.0.0.1:{server.server_port}", seen_uris
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
thread.join(timeout=5)
|
||||
|
||||
|
||||
# `.invalid` never resolves (RFC 6761), so the only way this request can succeed is through the proxy
|
||||
_PROXY_ONLY_UPSTREAM_URL = "http://upstream.invalid/v1/models"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("disable_aiohttp_transport", [True, False])
|
||||
@pytest.mark.parametrize("force_ipv4", [True, False])
|
||||
async def test_async_handler_honours_proxy_env_for_every_transport(
|
||||
forward_proxy_server, monkeypatch: pytest.MonkeyPatch, disable_aiohttp_transport: bool, force_ipv4: bool
|
||||
):
|
||||
proxy_url, seen_uris = forward_proxy_server
|
||||
monkeypatch.setenv("HTTP_PROXY", proxy_url)
|
||||
monkeypatch.delenv("NO_PROXY", raising=False)
|
||||
monkeypatch.delenv("no_proxy", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", disable_aiohttp_transport)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", force_ipv4)
|
||||
|
||||
handler = AsyncHTTPHandler()
|
||||
try:
|
||||
response = await handler.get(_PROXY_ONLY_UPSTREAM_URL)
|
||||
finally:
|
||||
await handler.close()
|
||||
|
||||
assert response.text == "via-proxy"
|
||||
assert seen_uris == [_PROXY_ONLY_UPSTREAM_URL]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("force_ipv4", [True, False])
|
||||
def test_sync_handler_honours_proxy_env(forward_proxy_server, monkeypatch: pytest.MonkeyPatch, force_ipv4: bool):
|
||||
proxy_url, seen_uris = forward_proxy_server
|
||||
monkeypatch.setenv("HTTP_PROXY", proxy_url)
|
||||
monkeypatch.delenv("NO_PROXY", raising=False)
|
||||
monkeypatch.delenv("no_proxy", raising=False)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", force_ipv4)
|
||||
|
||||
handler = HTTPHandler()
|
||||
try:
|
||||
response = handler.get(_PROXY_ONLY_UPSTREAM_URL)
|
||||
finally:
|
||||
handler.close()
|
||||
|
||||
assert response.text == "via-proxy"
|
||||
assert seen_uris == [_PROXY_ONLY_UPSTREAM_URL]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_force_ipv4_httpx_transport_honours_no_proxy(keepalive_server, monkeypatch: pytest.MonkeyPatch):
|
||||
"""NO_PROXY hosts must still go direct when the proxy mounts are supplied by litellm instead of httpx."""
|
||||
monkeypatch.setenv("HTTP_PROXY", "http://proxy.invalid:3128")
|
||||
monkeypatch.setenv("NO_PROXY", "127.0.0.1")
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", True)
|
||||
|
||||
handler = AsyncHTTPHandler()
|
||||
try:
|
||||
response = await handler.get(keepalive_server)
|
||||
finally:
|
||||
await handler.close()
|
||||
|
||||
assert response.text == "ok"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def private_ca_tls_upstream(tmp_path: pathlib.Path):
|
||||
"""HTTPS server behind a CONNECT proxy, both on localhost; the server's cert is signed by a test-only CA."""
|
||||
import datetime
|
||||
import select
|
||||
import socket
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from socketserver import ThreadingMixIn
|
||||
|
||||
from cryptography import x509
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from cryptography.x509.oid import NameOID
|
||||
|
||||
key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "upstream.invalid")])
|
||||
now = datetime.datetime.now(datetime.timezone.utc)
|
||||
cert = (
|
||||
x509.CertificateBuilder()
|
||||
.subject_name(name)
|
||||
.issuer_name(name)
|
||||
.public_key(key.public_key())
|
||||
.serial_number(x509.random_serial_number())
|
||||
.not_valid_before(now - datetime.timedelta(minutes=1))
|
||||
.not_valid_after(now + datetime.timedelta(hours=1))
|
||||
.add_extension(x509.SubjectAlternativeName([x509.DNSName("upstream.invalid")]), critical=False)
|
||||
.add_extension(x509.BasicConstraints(ca=True, path_length=None), critical=True)
|
||||
.sign(key, hashes.SHA256())
|
||||
)
|
||||
ca_pem = tmp_path / "ca.pem"
|
||||
ca_pem.write_bytes(cert.public_bytes(serialization.Encoding.PEM))
|
||||
key_pem = tmp_path / "key.pem"
|
||||
key_pem.write_bytes(
|
||||
key.private_bytes(
|
||||
serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()
|
||||
)
|
||||
)
|
||||
|
||||
class OkTlsHandler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def do_GET(self):
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Length", "6")
|
||||
self.end_headers()
|
||||
self.wfile.write(b"ok-tls")
|
||||
|
||||
def log_message(self, format, *args):
|
||||
pass
|
||||
|
||||
class ThreadedServer(ThreadingMixIn, HTTPServer):
|
||||
daemon_threads = True
|
||||
|
||||
tls_server = ThreadedServer(("127.0.0.1", 0), OkTlsHandler)
|
||||
server_ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
|
||||
server_ctx.load_cert_chain(str(ca_pem), str(key_pem))
|
||||
tls_server.socket = server_ctx.wrap_socket(tls_server.socket, server_side=True)
|
||||
tls_port = tls_server.server_port
|
||||
|
||||
class ConnectProxyHandler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
def do_CONNECT(self):
|
||||
upstream = socket.create_connection(("127.0.0.1", tls_port))
|
||||
self.send_response(200, "Connection established")
|
||||
self.end_headers()
|
||||
sockets = [self.connection, upstream]
|
||||
while True:
|
||||
readable, _, _ = select.select(sockets, [], [], 5)
|
||||
if not readable:
|
||||
break
|
||||
for src in readable:
|
||||
data = src.recv(65536)
|
||||
if not data:
|
||||
upstream.close()
|
||||
return
|
||||
(upstream if src is self.connection else self.connection).sendall(data)
|
||||
|
||||
def log_message(self, format, *args):
|
||||
pass
|
||||
|
||||
proxy_server = ThreadedServer(("127.0.0.1", 0), ConnectProxyHandler)
|
||||
threads = [
|
||||
threading.Thread(target=tls_server.serve_forever, daemon=True),
|
||||
threading.Thread(target=proxy_server.serve_forever, daemon=True),
|
||||
]
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
try:
|
||||
yield f"http://127.0.0.1:{proxy_server.server_port}", str(ca_pem)
|
||||
finally:
|
||||
for server in (proxy_server, tls_server):
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
for thread in threads:
|
||||
thread.join(timeout=5)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_force_ipv4_https_proxy_mount_uses_handler_ca_bundle(
|
||||
private_ca_tls_upstream, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
proxy_url, ca_pem = private_ca_tls_upstream
|
||||
monkeypatch.setenv("HTTPS_PROXY", proxy_url)
|
||||
monkeypatch.delenv("NO_PROXY", raising=False)
|
||||
monkeypatch.delenv("no_proxy", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", True)
|
||||
|
||||
handler = AsyncHTTPHandler(ssl_verify=ca_pem)
|
||||
try:
|
||||
response = await handler.get("https://upstream.invalid/v1/models")
|
||||
finally:
|
||||
await handler.close()
|
||||
|
||||
assert response.text == "ok-tls"
|
||||
|
||||
|
||||
def test_sync_force_ipv4_https_proxy_mount_uses_handler_ca_bundle(
|
||||
private_ca_tls_upstream, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
proxy_url, ca_pem = private_ca_tls_upstream
|
||||
monkeypatch.setenv("HTTPS_PROXY", proxy_url)
|
||||
monkeypatch.delenv("NO_PROXY", raising=False)
|
||||
monkeypatch.delenv("no_proxy", raising=False)
|
||||
monkeypatch.setattr(litellm, "force_ipv4", True)
|
||||
|
||||
handler = HTTPHandler(ssl_verify=ca_pem)
|
||||
try:
|
||||
response = handler.get("https://upstream.invalid/v1/models")
|
||||
finally:
|
||||
handler.close()
|
||||
|
||||
assert response.text == "ok-tls"
|
||||
|
|
|
|||
|
|
@ -67,20 +67,52 @@ class FakeAsyncEmbeddingFn(FakeEmbeddingFn):
|
|||
return SimpleNamespace(data=[{"embedding": self.embedding}])
|
||||
|
||||
|
||||
class FakeEmbeddingExecutor:
|
||||
def __init__(self, embedding):
|
||||
self.embedding = embedding
|
||||
self.captured = None
|
||||
|
||||
def embed(self, model, query, configuration):
|
||||
self.captured = (model, query, configuration)
|
||||
return SimpleNamespace(data=[{"embedding": self.embedding}])
|
||||
|
||||
async def aembed(self, model, query, configuration):
|
||||
self.captured = (model, query, configuration)
|
||||
return SimpleNamespace(data=[{"embedding": self.embedding}])
|
||||
|
||||
|
||||
def _doc(doc_id, distance, **fields):
|
||||
return SimpleNamespace(id=doc_id, vector_distance=str(distance), **fields)
|
||||
|
||||
|
||||
def _search(config, client=None, query="what is litellm", optional_params=None, litellm_params=None):
|
||||
def _search(config, client=None, query="what is litellm", optional_params=None, litellm_params=None, executor=None):
|
||||
return config.execute_search_vector_store_request(
|
||||
vector_store_id="my_index",
|
||||
query=query,
|
||||
vector_store_search_optional_params=optional_params or {},
|
||||
litellm_logging_obj=MagicMock(),
|
||||
litellm_params={"litellm_embedding_model": "openai/text-embedding-3-small", **(litellm_params or {})},
|
||||
embedding_executor=executor,
|
||||
)
|
||||
|
||||
|
||||
def test_sync_search_uses_request_embedding_executor_without_overwriting_explicit_config():
|
||||
executor = FakeEmbeddingExecutor([0.1, 0.2])
|
||||
config = ValkeyVectorStoreConfig(sync_client=FakeRedis())
|
||||
embedding_config = {"api_key": "store-specific-key", "aws_region_name": "us-west-2"}
|
||||
|
||||
_search(
|
||||
config,
|
||||
litellm_params={
|
||||
"litellm_embedding_model": "team-embedding-alias",
|
||||
"litellm_embedding_config": embedding_config,
|
||||
},
|
||||
executor=executor,
|
||||
)
|
||||
|
||||
assert executor.captured == ("team-embedding-alias", "what is litellm", embedding_config)
|
||||
|
||||
|
||||
def test_sync_search_builds_knn_query_with_packed_vector():
|
||||
embedding_fn = FakeEmbeddingFn([0.1, 0.2, 0.3])
|
||||
client = FakeRedis()
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ cache's hit/single-flight behavior. Each assertion fails under a real mutation o
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import gc
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -359,6 +360,110 @@ async def test_cache_invalidate_only_evicts_the_named_key():
|
|||
assert calls == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_invalidate_mid_compute_is_not_overwritten_by_that_compute():
|
||||
"""A bearer minted before an invalidation must never be served after it.
|
||||
|
||||
The compute is suspended at the token endpoint when the invalidation lands, so its write is
|
||||
the one that would resurrect the evicted bearer for the rest of its TTL. The caller it was
|
||||
minted for still gets it; the *cache* is what the invalidation is about.
|
||||
"""
|
||||
cache = ExchangedTokenCache()
|
||||
mint_started, release_mint = asyncio.Event(), asyncio.Event()
|
||||
|
||||
async def slow_mint():
|
||||
mint_started.set()
|
||||
await release_mint.wait()
|
||||
return _ok_token("bearer-minted-before-invalidation")
|
||||
|
||||
async def re_mint():
|
||||
return _ok_token("bearer-minted-after-invalidation")
|
||||
|
||||
in_flight = asyncio.create_task(cache.get_or_compute("slot", slow_mint, fingerprint="fp"))
|
||||
await mint_started.wait()
|
||||
|
||||
assert not in_flight.done()
|
||||
cache.invalidate("slot")
|
||||
release_mint.set()
|
||||
|
||||
raced = await in_flight
|
||||
assert isinstance(raced, Ok) and raced.ok == "bearer-minted-before-invalidation"
|
||||
|
||||
after = await cache.get_or_compute("slot", re_mint, fingerprint="fp")
|
||||
assert isinstance(after, Ok) and after.ok == "bearer-minted-after-invalidation"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_invalidate_mid_compute_survives_garbage_collection():
|
||||
"""The record of an invalidation must outlive a collection cycle taken mid-compute.
|
||||
|
||||
Per-key state is held weakly so idle keys do not accumulate. If the state a compute checks
|
||||
before writing were collectible while that compute is suspended, the check would read as
|
||||
"nothing was invalidated" and the stale write would land; the running compute has to pin it.
|
||||
"""
|
||||
cache = ExchangedTokenCache()
|
||||
mint_started, release_mint = asyncio.Event(), asyncio.Event()
|
||||
|
||||
async def slow_mint():
|
||||
mint_started.set()
|
||||
await release_mint.wait()
|
||||
return _ok_token("bearer-minted-before-invalidation")
|
||||
|
||||
async def re_mint():
|
||||
return _ok_token("bearer-minted-after-invalidation")
|
||||
|
||||
in_flight = asyncio.create_task(cache.get_or_compute("slot", slow_mint, fingerprint="fp"))
|
||||
await mint_started.wait()
|
||||
|
||||
assert not in_flight.done()
|
||||
cache.invalidate("slot")
|
||||
gc.collect()
|
||||
release_mint.set()
|
||||
await in_flight
|
||||
|
||||
after = await cache.get_or_compute("slot", re_mint, fingerprint="fp")
|
||||
assert isinstance(after, Ok) and after.ok == "bearer-minted-after-invalidation"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_stores_a_compute_that_started_after_the_invalidation():
|
||||
"""Only the mint that predates the invalidation loses its write.
|
||||
|
||||
A caller queued behind the single-flight lock computes after the eviction, so its token is
|
||||
fresh and must be cached; otherwise the fix would trade one stale bearer for re-minting on
|
||||
every subsequent resolution.
|
||||
"""
|
||||
cache = ExchangedTokenCache()
|
||||
mint_started, release_mint = asyncio.Event(), asyncio.Event()
|
||||
|
||||
async def slow_mint():
|
||||
mint_started.set()
|
||||
await release_mint.wait()
|
||||
return _ok_token("bearer-minted-before-invalidation")
|
||||
|
||||
async def re_mint():
|
||||
return _ok_token("bearer-minted-after-invalidation")
|
||||
|
||||
async def must_not_run():
|
||||
pytest.fail("the mint that followed the invalidation should have been cached")
|
||||
|
||||
in_flight = asyncio.create_task(cache.get_or_compute("slot", slow_mint, fingerprint="fp"))
|
||||
await mint_started.wait()
|
||||
queued = asyncio.create_task(cache.get_or_compute("slot", re_mint, fingerprint="fp"))
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert not queued.done()
|
||||
cache.invalidate("slot")
|
||||
release_mint.set()
|
||||
|
||||
raced, fresh = await asyncio.gather(in_flight, queued)
|
||||
assert isinstance(raced, Ok) and raced.ok == "bearer-minted-before-invalidation"
|
||||
assert isinstance(fresh, Ok) and fresh.ok == "bearer-minted-after-invalidation"
|
||||
|
||||
served = await cache.get_or_compute("slot", must_not_run, fingerprint="fp")
|
||||
assert isinstance(served, Ok) and served.ok == "bearer-minted-after-invalidation"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cache_does_not_store_a_failed_compute():
|
||||
cache = ExchangedTokenCache()
|
||||
|
|
|
|||
|
|
@ -8517,6 +8517,79 @@ class TestPreemptive401ModeAware:
|
|||
await self._run(delegate, self.LITELLM_KEY_HEADERS, has_stored_token=False)
|
||||
|
||||
|
||||
def _make_obo_server(alias: str) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id=f"id-{alias}",
|
||||
name=alias,
|
||||
alias=alias,
|
||||
server_name=alias,
|
||||
url=f"https://{alias}.test/mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2_token_exchange,
|
||||
token_exchange_endpoint="https://idp.test/token",
|
||||
client_id="cid",
|
||||
client_secret="csecret",
|
||||
mcp_info={"server_name": alias},
|
||||
)
|
||||
|
||||
|
||||
class TestOboPreflightScopedToAllowedServers:
|
||||
"""The connect-time OBO exchange is an outbound IdP call whose result is cached, so it must
|
||||
only run for a server the caller's key resolves to through the allowed set, not for any
|
||||
server the requested path happens to name."""
|
||||
|
||||
SUBJECT_HEADERS = {"Authorization": "Bearer upstream-subject-token"}
|
||||
|
||||
async def _run(self, requested: MCPServer, allowed: list[MCPServer], user_api_key_auth: UserAPIKeyAuth | None):
|
||||
from litellm.proxy._experimental.mcp_server import server as server_module
|
||||
|
||||
allowed_lookup = AsyncMock(return_value=allowed)
|
||||
preflight = AsyncMock()
|
||||
with (
|
||||
patch.object( # test-quality-ok: route handler reads the module-level manager, no injection seam
|
||||
server_module.global_mcp_server_manager, "get_mcp_server_by_name", return_value=requested
|
||||
),
|
||||
patch.object( # test-quality-ok: the exchanger is the observable; a real one would call an IdP
|
||||
server_module.global_mcp_server_manager, "preflight_token_exchange", preflight
|
||||
),
|
||||
patch.object( # test-quality-ok: allowed-set resolution needs the DB; the test controls its answer
|
||||
server_module, "_get_allowed_mcp_servers", allowed_lookup
|
||||
),
|
||||
):
|
||||
await server_module._raise_preemptive_401_for_unauthenticated_servers(
|
||||
scope={"type": "http", "method": "POST", "path": f"/mcp/{requested.alias}", "headers": []},
|
||||
mcp_servers=[requested.alias],
|
||||
oauth2_headers=self.SUBJECT_HEADERS,
|
||||
mcp_server_auth_headers=None,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
client_ip="10.0.0.7",
|
||||
)
|
||||
return allowed_lookup, preflight
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unentitled_key_never_reaches_the_exchanger(self):
|
||||
requested = _make_obo_server("obo_tools")
|
||||
key = UserAPIKeyAuth(api_key="sk-plain-only")
|
||||
|
||||
allowed_lookup, preflight = await self._run(
|
||||
requested, allowed=[_make_obo_server("plain_tools")], user_api_key_auth=key
|
||||
)
|
||||
|
||||
preflight.assert_not_awaited()
|
||||
allowed_lookup.assert_awaited_once_with(
|
||||
user_api_key_auth=key, mcp_servers=[requested.alias], client_ip="10.0.0.7"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_entitled_key_still_exchanges_at_connect(self):
|
||||
requested = _make_obo_server("obo_tools")
|
||||
key = UserAPIKeyAuth(api_key="sk-obo")
|
||||
|
||||
_, preflight = await self._run(requested, allowed=[requested], user_api_key_auth=key)
|
||||
|
||||
preflight.assert_awaited_once_with(server=requested, oauth2_headers=self.SUBJECT_HEADERS, user_api_key_auth=key)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_mcp_call_guardrails_return_the_rewritten_result():
|
||||
"""The result a post_mcp_call guardrail rewrote must be what the caller sends back."""
|
||||
|
|
|
|||
|
|
@ -59,6 +59,7 @@ def _mock_cli_sso_start_response(
|
|||
login_id: str = "cli-session-uuid-456",
|
||||
poll_secret: str = "poll-secret",
|
||||
user_code: str = "ABCD-EFGH",
|
||||
**extra_fields: object,
|
||||
) -> Mock:
|
||||
mock_response = Mock()
|
||||
mock_response.status_code = 200
|
||||
|
|
@ -66,6 +67,7 @@ def _mock_cli_sso_start_response(
|
|||
"login_id": login_id,
|
||||
"poll_secret": poll_secret,
|
||||
"user_code": user_code,
|
||||
**extra_fields,
|
||||
}
|
||||
mock_response.raise_for_status = Mock()
|
||||
return mock_response
|
||||
|
|
@ -333,7 +335,9 @@ class TestLoginCommand:
|
|||
call_args = mock_browser.call_args[0][0]
|
||||
assert "https://test.example.com/sso/key/generate" in call_args
|
||||
assert "cli-test-uuid-123" in call_args
|
||||
assert "user_code" not in call_args
|
||||
assert "Verification code: ABCD-EFGH" in result.output
|
||||
assert "pre-filled in the browser" not in result.output
|
||||
mock_post.assert_called_once()
|
||||
mock_get.assert_called()
|
||||
assert mock_get.call_args.kwargs["headers"] == {"x-litellm-cli-poll-secret": "poll-secret"}
|
||||
|
|
@ -347,6 +351,72 @@ class TestLoginCommand:
|
|||
# Verify commands were shown
|
||||
mock_show_commands.assert_called_once()
|
||||
|
||||
def test_login_prefills_the_code_in_the_browser_when_the_proxy_advertises_it(
|
||||
self, isolated_home, secret_vault_factory
|
||||
) -> None:
|
||||
vault = secret_vault_factory()
|
||||
poll_response = Mock()
|
||||
poll_response.status_code = 200
|
||||
poll_response.json.return_value = {
|
||||
"status": "ready",
|
||||
"key": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt",
|
||||
"user_id": "test-user-123",
|
||||
"team_id": "team-1",
|
||||
"teams": ["team-1"],
|
||||
}
|
||||
start_response = _mock_cli_sso_start_response(
|
||||
login_id="cli-test-uuid-123",
|
||||
verification_uri_complete=(
|
||||
"https://internal-hostname.example.com/sso/key/generate"
|
||||
"?source=litellm-cli&key=cli-test-uuid-123&user_code=ABCD-EFGH"
|
||||
),
|
||||
)
|
||||
|
||||
with (
|
||||
patch("webbrowser.open") as mock_browser,
|
||||
patch("requests.post", return_value=start_response),
|
||||
patch("requests.get", return_value=poll_response),
|
||||
):
|
||||
result = self.runner.invoke(login, obj={"base_url": "https://test.example.com", "secret_vault": vault})
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert json.loads(vault.blob)["key"] == "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt"
|
||||
assert json.loads((isolated_home / ".litellm" / "token.json").read_text())["user_id"] == "test-user-123"
|
||||
opened_url = mock_browser.call_args[0][0]
|
||||
assert opened_url.startswith("https://test.example.com/sso/key/generate?")
|
||||
assert "internal-hostname" not in opened_url
|
||||
assert "key=cli-test-uuid-123" in opened_url
|
||||
assert "user_code=ABCD-EFGH" in opened_url
|
||||
assert "Verification code: ABCD-EFGH (pre-filled in the browser, check it matches)" in result.output
|
||||
|
||||
def test_login_keeps_the_code_out_of_the_url_when_the_proxy_sends_a_non_url_verification_uri(
|
||||
self, secret_vault_factory
|
||||
) -> None:
|
||||
poll_response = Mock()
|
||||
poll_response.status_code = 200
|
||||
poll_response.json.return_value = {
|
||||
"status": "ready",
|
||||
"key": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt",
|
||||
"user_id": "test-user-123",
|
||||
"team_id": "team-1",
|
||||
"teams": ["team-1"],
|
||||
}
|
||||
for advertised in (None, True):
|
||||
start_response = _mock_cli_sso_start_response(verification_uri_complete=advertised)
|
||||
|
||||
with (
|
||||
patch("webbrowser.open") as mock_browser,
|
||||
patch("requests.post", return_value=start_response),
|
||||
patch("requests.get", return_value=poll_response),
|
||||
):
|
||||
result = self.runner.invoke(
|
||||
login, obj={"base_url": "https://test.example.com", "secret_vault": secret_vault_factory()}
|
||||
)
|
||||
|
||||
assert result.exit_code == 0, result.output
|
||||
assert "user_code" not in mock_browser.call_args[0][0]
|
||||
assert "pre-filled in the browser" not in result.output
|
||||
|
||||
def test_login_timeout(self):
|
||||
"""Test login timeout scenario"""
|
||||
mock_context = Mock()
|
||||
|
|
|
|||
|
|
@ -3,7 +3,9 @@ from unittest.mock import patch
|
|||
import httpx
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.exceptions import Timeout
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket
|
||||
from litellm.proxy.guardrails.guardrail_hooks.crowdstrike_aidr import initialize_guardrail
|
||||
|
|
@ -12,8 +14,8 @@ from litellm.proxy.guardrails.guardrail_hooks.crowdstrike_aidr.crowdstrike_aidr
|
|||
CrowdStrikeAIDRHandler,
|
||||
)
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse
|
||||
from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams
|
||||
from litellm.types.utils import Delta, GenericGuardrailAPIInputs, ModelResponse, ModelResponseStream
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
|
|
@ -1578,3 +1580,139 @@ async def test_unparseable_transformed_response_fails_closed_under_fail_open() -
|
|||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "failing closed" in exc_info.value.detail["error"]
|
||||
|
||||
|
||||
def _initialize_from_config(**litellm_params_kwargs: object) -> CrowdStrikeAIDRHandler:
|
||||
litellm_params = LitellmParams(
|
||||
guardrail="crowdstrike_aidr",
|
||||
api_key="pts_crowdstrike_tokenid",
|
||||
api_base="https://api.crowdstrike.com/aidr/aiguard",
|
||||
default_on=True,
|
||||
**litellm_params_kwargs,
|
||||
)
|
||||
guardrail = Guardrail(guardrail_name="crowdstrike-aidr-guard", litellm_params=litellm_params)
|
||||
return initialize_guardrail(litellm_params=litellm_params, guardrail=guardrail)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("mode", "runs_pre_call", "runs_post_call"),
|
||||
[("post_call", False, True), ("pre_call", True, False), (["pre_call", "post_call"], True, True)],
|
||||
)
|
||||
def test_initialize_guardrail_honors_configured_mode(
|
||||
mode: str | list[str], runs_pre_call: bool, runs_post_call: bool
|
||||
) -> None:
|
||||
handler = _initialize_from_config(mode=mode)
|
||||
|
||||
assert handler.should_run_guardrail({}, GuardrailEventHooks.pre_call) is runs_pre_call
|
||||
assert handler.should_run_guardrail({}, GuardrailEventHooks.post_call) is runs_post_call
|
||||
|
||||
|
||||
def test_initialize_guardrail_rejects_unsupported_mode_instead_of_running_other_hooks() -> None:
|
||||
with pytest.raises(ValueError, match="during_call is not in the supported event hooks"):
|
||||
_initialize_from_config(mode="during_call")
|
||||
|
||||
|
||||
def test_initialize_guardrail_defaults_streaming_params() -> None:
|
||||
handler = _initialize_from_config(mode="post_call")
|
||||
|
||||
assert handler.streaming_end_of_stream_only is False
|
||||
assert handler.streaming_sampling_rate == 5
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"configured",
|
||||
[
|
||||
{"streaming_end_of_stream_only": True, "streaming_sampling_rate": 50},
|
||||
{"optional_params": {"streaming_end_of_stream_only": True, "streaming_sampling_rate": 50}},
|
||||
],
|
||||
)
|
||||
def test_initialize_guardrail_forwards_streaming_params(configured: dict[str, object]) -> None:
|
||||
handler = _initialize_from_config(mode="post_call", **configured)
|
||||
|
||||
assert handler.streaming_end_of_stream_only is True
|
||||
assert handler.streaming_sampling_rate == 50
|
||||
|
||||
|
||||
def test_initialize_guardrail_rejects_non_positive_sampling_rate() -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
_initialize_from_config(mode="post_call", streaming_sampling_rate=0)
|
||||
|
||||
|
||||
def test_update_in_memory_litellm_params_reapplies_streaming_params() -> None:
|
||||
handler = _initialize_from_config(mode="post_call")
|
||||
|
||||
handler.update_in_memory_litellm_params(
|
||||
LitellmParams(
|
||||
guardrail="crowdstrike_aidr",
|
||||
mode="post_call",
|
||||
streaming_end_of_stream_only=True,
|
||||
streaming_sampling_rate=7,
|
||||
)
|
||||
)
|
||||
|
||||
assert handler.streaming_end_of_stream_only is True
|
||||
assert handler.streaming_sampling_rate == 7
|
||||
|
||||
|
||||
def _stream_chunk(content: str, finish_reason: str | None) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
model="gpt-4",
|
||||
choices=[
|
||||
litellm.StreamingChoices(
|
||||
index=0, delta=Delta(role="assistant", content=content), finish_reason=finish_reason
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
async def _guard_calls_for_stream(handler: CrowdStrikeAIDRHandler, chunk_texts: list[str]) -> int:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import UnifiedLLMGuardrails
|
||||
|
||||
async def stream():
|
||||
for i, content in enumerate(chunk_texts):
|
||||
yield _stream_chunk(content, "stop" if i == len(chunk_texts) - 1 else None)
|
||||
|
||||
calls = 0
|
||||
|
||||
def _allow(request: httpx.Request) -> httpx.Response:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
return httpx.Response(
|
||||
status_code=200, json={"result": {"blocked": False, "transformed": False}}, request=request
|
||||
)
|
||||
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"guardrail_to_apply": handler,
|
||||
"metadata": {"guardrails": ["crowdstrike-aidr-guard"]},
|
||||
}
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(_allow)) as client:
|
||||
await handler.async_handler.close()
|
||||
handler.async_handler.client = client
|
||||
async for _ in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test", request_route="/chat/completions"),
|
||||
response=stream(),
|
||||
request_data=request_data,
|
||||
):
|
||||
pass
|
||||
return calls
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("configured", "expected_calls"),
|
||||
[
|
||||
({}, 3),
|
||||
({"streaming_sampling_rate": 2}, 6),
|
||||
({"streaming_end_of_stream_only": True}, 1),
|
||||
({"streaming_end_of_stream_only": True, "streaming_sampling_rate": 2}, 1),
|
||||
],
|
||||
)
|
||||
async def test_streaming_params_from_config_control_output_scan_cadence(
|
||||
configured: dict[str, object], expected_calls: int
|
||||
) -> None:
|
||||
"""10 chunks: default samples at 5 and 10 plus the final pass, rate 2 samples 5 times plus final, end-of-stream scans once."""
|
||||
handler = _initialize_from_config(mode="post_call", **configured)
|
||||
|
||||
assert await _guard_calls_for_stream(handler, list("ABCDEFGHIJ")) == expected_calls
|
||||
|
|
|
|||
|
|
@ -104,7 +104,7 @@ def mock_in_memory_handler(mocker):
|
|||
mock_handler.get_guardrail_by_id.return_value = MOCK_CONFIG_GUARDRAIL
|
||||
mock_handler.get_source.return_value = "config"
|
||||
mock_handler.initialize_guardrail = mocker.Mock()
|
||||
mock_handler.update_in_memory_guardrail = mocker.Mock()
|
||||
mock_handler.sync_guardrail_from_db = mocker.Mock()
|
||||
mock_handler.delete_in_memory_guardrail = mocker.Mock()
|
||||
mock_handler.reconcile_db_guardrails = mocker.Mock(return_value=[])
|
||||
return mock_handler
|
||||
|
|
@ -1047,13 +1047,15 @@ async def test_create_guardrail_endpoint(
|
|||
"scenario,expected_result,expected_exception",
|
||||
[
|
||||
("success_with_sync", "test-db-guardrail", None),
|
||||
("success_sync_fails", "test-db-guardrail", None),
|
||||
("success_sync_fails_unexpected_error", "test-db-guardrail", None),
|
||||
("sync_fails_invalid_config", None, HTTPException),
|
||||
("database_failure", None, HTTPException),
|
||||
("no_prisma_client", None, HTTPException),
|
||||
],
|
||||
ids=[
|
||||
"success_with_immediate_sync",
|
||||
"success_but_sync_fails",
|
||||
"success_but_sync_fails_with_unexpected_error",
|
||||
"sync_rejects_invalid_config",
|
||||
"database_error",
|
||||
"missing_prisma_client",
|
||||
],
|
||||
|
|
@ -1073,6 +1075,7 @@ async def test_update_guardrail_endpoint(
|
|||
mock_logger = None
|
||||
if scenario == "success_with_sync":
|
||||
mock_prisma_client = mocker.Mock()
|
||||
mock_in_memory_handler.sync_guardrail_from_db = mocker.Mock()
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY",
|
||||
|
|
@ -1083,10 +1086,13 @@ async def test_update_guardrail_endpoint(
|
|||
mock_in_memory_handler,
|
||||
)
|
||||
|
||||
elif scenario == "success_sync_fails":
|
||||
elif scenario == "success_sync_fails_unexpected_error":
|
||||
# A non-ValueError/TypeError failure is not a config-rejection signal,
|
||||
# so it keeps the pre-existing swallow-and-warn behavior rather than
|
||||
# rolling back the DB write.
|
||||
mock_prisma_client = mocker.Mock()
|
||||
mock_in_memory_handler.update_in_memory_guardrail.side_effect = Exception(
|
||||
"Sync failed"
|
||||
mock_in_memory_handler.sync_guardrail_from_db = mocker.Mock(
|
||||
side_effect=Exception("Sync failed")
|
||||
)
|
||||
mock_logger = mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.verbose_proxy_logger"
|
||||
|
|
@ -1102,6 +1108,25 @@ async def test_update_guardrail_endpoint(
|
|||
mock_in_memory_handler,
|
||||
)
|
||||
|
||||
elif scenario == "sync_fails_invalid_config":
|
||||
# Regression for the PUT half of the fix: a TypeError from the sync (the
|
||||
# in-place update_in_memory_guardrail raised exactly this on every PUT)
|
||||
# must roll back the DB write and surface a 422, not persist the
|
||||
# rejected config with a 200.
|
||||
mock_prisma_client = mocker.Mock()
|
||||
mock_in_memory_handler.sync_guardrail_from_db = mocker.Mock(
|
||||
side_effect=TypeError("vars() argument must have __dict__ attribute")
|
||||
)
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # test-quality-ok: reused pattern
|
||||
mocker.patch( # test-quality-ok: reused pattern
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY",
|
||||
mock_guardrail_registry,
|
||||
)
|
||||
mocker.patch( # test-quality-ok: reused pattern
|
||||
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
|
||||
mock_in_memory_handler,
|
||||
)
|
||||
|
||||
elif scenario == "database_failure":
|
||||
mock_prisma_client = mocker.Mock()
|
||||
mock_guardrail_registry.update_guardrail_in_db.side_effect = Exception(
|
||||
|
|
@ -1130,6 +1155,16 @@ async def test_update_guardrail_endpoint(
|
|||
assert "Database error" in str(exc_info.value.detail)
|
||||
elif scenario == "no_prisma_client":
|
||||
assert "Prisma client not initialized" in str(exc_info.value.detail)
|
||||
elif scenario == "sync_fails_invalid_config":
|
||||
assert exc_info.value.status_code == 422
|
||||
assert "update rejected" in str(exc_info.value.detail)
|
||||
# Rolled back: update_guardrail_in_db is called once for the
|
||||
# rejected write and once more to restore the previous config.
|
||||
assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2
|
||||
assert (
|
||||
mock_guardrail_registry.update_guardrail_in_db.call_args.kwargs["guardrail"]
|
||||
== MOCK_DB_GUARDRAIL
|
||||
)
|
||||
|
||||
else:
|
||||
result = await update_guardrail(
|
||||
|
|
@ -1145,11 +1180,11 @@ async def test_update_guardrail_endpoint(
|
|||
prisma_client=mocker.ANY,
|
||||
)
|
||||
|
||||
mock_in_memory_handler.update_in_memory_guardrail.assert_called_once_with(
|
||||
guardrail_id="test-guardrail-id", guardrail=mocker.ANY
|
||||
mock_in_memory_handler.sync_guardrail_from_db.assert_called_once_with(
|
||||
guardrail=mocker.ANY
|
||||
)
|
||||
|
||||
if scenario == "success_sync_fails":
|
||||
if scenario == "success_sync_fails_unexpected_error":
|
||||
assert mock_logger is not None
|
||||
mock_logger.warning.assert_called_once()
|
||||
assert "Failed to update" in str(mock_logger.warning.call_args)
|
||||
|
|
|
|||
|
|
@ -913,3 +913,96 @@ def test_reinitialize_guardrail_restores_previous_on_failure():
|
|||
assert restored.guardrail_name == "restore-me"
|
||||
finally:
|
||||
registry_module.guardrail_initializer_registry.pop("restore_test", None)
|
||||
|
||||
|
||||
def test_reinitialize_guardrail_raises_value_error_for_non_value_error_init_failures():
|
||||
"""Regression for the LIT-6479 fix's 422 path: a constructor failure that is not
|
||||
already a ValueError/TypeError (re.error from an invalid regex has neither in its
|
||||
MRO) must still surface as ValueError, so the PUT/PATCH endpoints' rollback+422
|
||||
catch is exhaustive instead of warn-and-200 persisting a broken config."""
|
||||
import re
|
||||
|
||||
from litellm.proxy.guardrails import guardrail_registry as registry_module
|
||||
|
||||
def _initializer(litellm_params, guardrail):
|
||||
if litellm_params.api_key == "bad-regex":
|
||||
re.compile("([")
|
||||
return CustomGuardrail(
|
||||
guardrail_name=guardrail["guardrail_name"],
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
registry_module.guardrail_initializer_registry["regex_test"] = _initializer
|
||||
try:
|
||||
handler = InMemoryGuardrailHandler()
|
||||
created = handler.initialize_guardrail(
|
||||
guardrail={
|
||||
"guardrail_name": "regex-me",
|
||||
"litellm_params": {"guardrail": "regex_test", "mode": "pre_call", "api_key": "ok"},
|
||||
},
|
||||
)
|
||||
guardrail_id = created["guardrail_id"]
|
||||
|
||||
with pytest.raises(ValueError, match="Guardrail initialization failed") as excinfo:
|
||||
handler.reinitialize_guardrail(
|
||||
guardrail={
|
||||
"guardrail_id": guardrail_id,
|
||||
"guardrail_name": "regex-me",
|
||||
"litellm_params": {"guardrail": "regex_test", "mode": "pre_call", "api_key": "bad-regex"},
|
||||
},
|
||||
)
|
||||
|
||||
assert isinstance(excinfo.value.__cause__, re.error)
|
||||
assert guardrail_id in handler.IN_MEMORY_GUARDRAILS
|
||||
restored = handler.guardrail_id_to_custom_guardrail[guardrail_id]
|
||||
assert restored is not None and restored.guardrail_name == "regex-me"
|
||||
finally:
|
||||
registry_module.guardrail_initializer_registry.pop("regex_test", None)
|
||||
|
||||
|
||||
def test_sync_guardrail_from_db_applies_db_dict_params_to_live_instance():
|
||||
"""
|
||||
Regression for PUT /guardrails/{id}: the DB row arrives with litellm_params as
|
||||
a plain jsonb dict, and the in-place update_in_memory_guardrail cast it to
|
||||
LitellmParams without constructing one, so vars() raised and the running proxy
|
||||
kept enforcing the stale config forever. The PUT endpoint now routes through
|
||||
sync_guardrail_from_db, which must rebuild the live instance from the dict:
|
||||
new blocked words compiled in, old ones gone, and the event hook re-derived
|
||||
from mode (the base-class setattr path wrote self.mode while dispatch reads
|
||||
self.event_hook, so only a full re-init applies a mode change).
|
||||
"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
|
||||
ContentFilterGuardrail,
|
||||
)
|
||||
|
||||
handler = InMemoryGuardrailHandler()
|
||||
gid = "66666666-6666-6666-6666-666666666666"
|
||||
|
||||
def db_guardrail(word: str, mode: str) -> Guardrail:
|
||||
return Guardrail(
|
||||
guardrail_id=gid,
|
||||
guardrail_name="cf-put-sync",
|
||||
litellm_params={
|
||||
"guardrail": "litellm_content_filter",
|
||||
"mode": mode,
|
||||
"default_on": True,
|
||||
"blocked_words": [{"keyword": word, "action": "BLOCK"}],
|
||||
},
|
||||
)
|
||||
|
||||
lists = _all_callback_lists()
|
||||
snapshots = [list(cb_list) for cb_list in lists]
|
||||
try:
|
||||
handler.sync_guardrail_from_db(db_guardrail("foobarblock", "pre_call"))
|
||||
handler.sync_guardrail_from_db(db_guardrail("quxnewblock", "during_call"))
|
||||
|
||||
instance = handler.guardrail_id_to_custom_guardrail[gid]
|
||||
assert isinstance(instance, ContentFilterGuardrail)
|
||||
assert instance._check_blocked_words("hello QUXNEWBLOCK") is not None
|
||||
assert instance._check_blocked_words("hello FOOBARBLOCK") is None
|
||||
assert instance.event_hook == GuardrailEventHooks.during_call
|
||||
assert instance.should_run_guardrail(data={}, event_type=GuardrailEventHooks.during_call) is True
|
||||
finally:
|
||||
for cb_list, snapshot in zip(lists, snapshots):
|
||||
cb_list[:] = snapshot
|
||||
|
|
|
|||
|
|
@ -1077,3 +1077,243 @@ def test_public_mcp_hub_does_not_expose_upstream_url():
|
|||
assert all("url" not in item for item in data)
|
||||
assert secret_url not in response.text
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def reset_autorouter_presets_cache():
|
||||
from litellm.proxy.public_endpoints.public_endpoints import _AutoRouterPresetsCache
|
||||
|
||||
_AutoRouterPresetsCache.presets = None
|
||||
_AutoRouterPresetsCache.lock = None
|
||||
yield
|
||||
_AutoRouterPresetsCache.presets = None
|
||||
_AutoRouterPresetsCache.lock = None
|
||||
|
||||
|
||||
def test_get_autorouter_presets_local_mode_serves_bundled_catalog(
|
||||
monkeypatch, reset_autorouter_presets_cache
|
||||
):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_AUTOROUTER_PRESETS", "True")
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
client = TestClient(app)
|
||||
|
||||
response = client.get("/public/autorouter_presets")
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = response.json()
|
||||
assert "anthropic_family" in payload
|
||||
for preset in payload.values():
|
||||
assert isinstance(preset["label"], str)
|
||||
assert isinstance(preset["description"], str)
|
||||
assert "tiers" in preset["complexity_router_config"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_autorouter_presets_fetches_once_per_process(
|
||||
monkeypatch, reset_autorouter_presets_cache
|
||||
):
|
||||
from litellm.proxy.public_endpoints.public_endpoints import (
|
||||
_AUTOROUTER_PRESETS_ADAPTER,
|
||||
get_autorouter_presets,
|
||||
)
|
||||
|
||||
monkeypatch.delenv("LITELLM_LOCAL_AUTOROUTER_PRESETS", raising=False)
|
||||
remote = _AUTOROUTER_PRESETS_ADAPTER.validate_python(
|
||||
{
|
||||
"remote_only": {
|
||||
"label": "Remote Only",
|
||||
"description": "from the remote catalog",
|
||||
"complexity_router_config": {"tiers": {"SIMPLE": ["m1"], "MEDIUM": ["m2"], "COMPLEX": ["m3"], "REASONING": ["m4"]}},
|
||||
}
|
||||
}
|
||||
)
|
||||
calls = []
|
||||
|
||||
async def fake_fetch(url):
|
||||
calls.append(url)
|
||||
return remote
|
||||
|
||||
first = await get_autorouter_presets(url="https://example.test/presets.json", fetch=fake_fetch)
|
||||
second = await get_autorouter_presets(url="https://example.test/presets.json", fetch=fake_fetch)
|
||||
|
||||
assert first == remote
|
||||
assert second == remote
|
||||
assert calls == ["https://example.test/presets.json"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_autorouter_presets_single_flight_on_concurrent_cold_start(
|
||||
monkeypatch, reset_autorouter_presets_cache
|
||||
):
|
||||
import asyncio
|
||||
|
||||
from litellm.proxy.public_endpoints.public_endpoints import (
|
||||
_AUTOROUTER_PRESETS_ADAPTER,
|
||||
get_autorouter_presets,
|
||||
)
|
||||
|
||||
monkeypatch.delenv("LITELLM_LOCAL_AUTOROUTER_PRESETS", raising=False)
|
||||
remote = _AUTOROUTER_PRESETS_ADAPTER.validate_python(
|
||||
{
|
||||
"remote_only": {
|
||||
"label": "Remote Only",
|
||||
"description": "from the remote catalog",
|
||||
"complexity_router_config": {"tiers": {"SIMPLE": ["m1"], "MEDIUM": ["m2"], "COMPLEX": ["m3"], "REASONING": ["m4"]}},
|
||||
}
|
||||
}
|
||||
)
|
||||
calls = []
|
||||
|
||||
async def slow_fetch(url):
|
||||
calls.append(url)
|
||||
await asyncio.sleep(0.05)
|
||||
return remote
|
||||
|
||||
results = await asyncio.gather(
|
||||
get_autorouter_presets(url="https://example.test/presets.json", fetch=slow_fetch),
|
||||
get_autorouter_presets(url="https://example.test/presets.json", fetch=slow_fetch),
|
||||
get_autorouter_presets(url="https://example.test/presets.json", fetch=slow_fetch),
|
||||
)
|
||||
|
||||
assert all(result == remote for result in results)
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_autorouter_presets_caches_bundled_fallback_on_remote_failure(
|
||||
monkeypatch, reset_autorouter_presets_cache
|
||||
):
|
||||
from litellm.proxy.public_endpoints.public_endpoints import get_autorouter_presets
|
||||
|
||||
monkeypatch.delenv("LITELLM_LOCAL_AUTOROUTER_PRESETS", raising=False)
|
||||
calls = []
|
||||
|
||||
async def broken_fetch(url):
|
||||
calls.append(url)
|
||||
raise ValueError("remote catalog unavailable")
|
||||
|
||||
first = await get_autorouter_presets(url="https://example.test/presets.json", fetch=broken_fetch)
|
||||
second = await get_autorouter_presets(url="https://example.test/presets.json", fetch=broken_fetch)
|
||||
|
||||
assert "anthropic_family" in first
|
||||
assert second == first
|
||||
assert len(calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_autorouter_presets_adapter_rejects_wrong_shapes():
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.proxy.public_endpoints.public_endpoints import _AUTOROUTER_PRESETS_ADAPTER
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
_AUTOROUTER_PRESETS_ADAPTER.validate_python({"bad": {"label": "no description or config"}})
|
||||
with pytest.raises(ValidationError):
|
||||
_AUTOROUTER_PRESETS_ADAPTER.validate_python(["not", "a", "mapping"])
|
||||
with pytest.raises(ValidationError):
|
||||
_AUTOROUTER_PRESETS_ADAPTER.validate_python(
|
||||
{"no_tiers": {"label": "L", "description": "D", "complexity_router_config": {}}}
|
||||
)
|
||||
with pytest.raises(ValidationError):
|
||||
_AUTOROUTER_PRESETS_ADAPTER.validate_python(
|
||||
{
|
||||
"missing_builtin_tier": {
|
||||
"label": "L",
|
||||
"description": "D",
|
||||
"complexity_router_config": {"tiers": {"SIMPLE": ["m1"], "MEDIUM": ["m2"], "COMPLEX": ["m3"]}},
|
||||
}
|
||||
}
|
||||
)
|
||||
with pytest.raises(ValidationError):
|
||||
_AUTOROUTER_PRESETS_ADAPTER.validate_python(
|
||||
{
|
||||
"unknown_tier_name": {
|
||||
"label": "L",
|
||||
"description": "D",
|
||||
"complexity_router_config": {
|
||||
"tiers": {
|
||||
"SIMPLE": ["m1"],
|
||||
"MEDIUM": ["m2"],
|
||||
"COMPLEX": ["m3"],
|
||||
"REASONING": ["m4"],
|
||||
"ULTRA": ["m5"],
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
)
|
||||
with pytest.raises(ValidationError):
|
||||
_AUTOROUTER_PRESETS_ADAPTER.validate_python(
|
||||
{
|
||||
"bad_tiers": {
|
||||
"label": "L",
|
||||
"description": "D",
|
||||
"complexity_router_config": {"tiers": "not-a-mapping"},
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def test_get_autorouter_presets_passes_unknown_catalog_fields_through(
|
||||
monkeypatch, reset_autorouter_presets_cache
|
||||
):
|
||||
from litellm.proxy.public_endpoints.public_endpoints import (
|
||||
_AUTOROUTER_PRESETS_ADAPTER,
|
||||
_AutoRouterPresetsCache,
|
||||
)
|
||||
|
||||
monkeypatch.delenv("LITELLM_LOCAL_AUTOROUTER_PRESETS", raising=False)
|
||||
_AutoRouterPresetsCache.presets = _AUTOROUTER_PRESETS_ADAPTER.validate_python(
|
||||
{
|
||||
"future_preset": {
|
||||
"label": "Future",
|
||||
"description": "carries fields this proxy version does not know",
|
||||
"complexity_router_config": {
|
||||
"tiers": {"SIMPLE": ["m1"], "MEDIUM": ["m2"], "COMPLEX": ["m3"], "REASONING": ["m4"]},
|
||||
"future_config_knob": 3,
|
||||
},
|
||||
"icon": "sparkles",
|
||||
}
|
||||
}
|
||||
)
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
client = TestClient(app)
|
||||
|
||||
response = client.get("/public/autorouter_presets")
|
||||
|
||||
assert response.status_code == 200
|
||||
served = response.json()["future_preset"]
|
||||
assert served["icon"] == "sparkles"
|
||||
assert served["complexity_router_config"]["future_config_knob"] == 3
|
||||
assert served["complexity_router_config"]["tiers"]["SIMPLE"] == ["m1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_remote_autorouter_presets_parses_and_rejects_empty(monkeypatch):
|
||||
import litellm.llms.custom_httpx.http_handler as http_handler_module
|
||||
from litellm.proxy.public_endpoints.public_endpoints import _fetch_remote_autorouter_presets
|
||||
|
||||
catalog = {
|
||||
"remote_only": {
|
||||
"label": "Remote Only",
|
||||
"description": "from the remote catalog",
|
||||
"complexity_router_config": {"tiers": {"SIMPLE": ["m1"], "MEDIUM": ["m2"], "COMPLEX": ["m3"], "REASONING": ["m4"]}},
|
||||
}
|
||||
}
|
||||
response = MagicMock()
|
||||
response.raise_for_status = MagicMock()
|
||||
response.json = MagicMock(return_value=catalog)
|
||||
client = MagicMock()
|
||||
client.get = AsyncMock(return_value=response)
|
||||
monkeypatch.setattr(http_handler_module, "get_async_httpx_client", lambda llm_provider: client)
|
||||
|
||||
presets = await _fetch_remote_autorouter_presets("https://example.test/presets.json")
|
||||
assert presets["remote_only"].label == "Remote Only"
|
||||
response.raise_for_status.assert_called_once()
|
||||
|
||||
response.json = MagicMock(return_value={})
|
||||
with pytest.raises(ValueError, match="empty"):
|
||||
await _fetch_remote_autorouter_presets("https://example.test/presets.json")
|
||||
|
|
|
|||
|
|
@ -3959,18 +3959,18 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
|
|||
]
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_spendlogs.group_by = AsyncMock(
|
||||
return_value=[
|
||||
{"session_id": session_id, "_count": {"session_id": 2}},
|
||||
]
|
||||
)
|
||||
mock_prisma.db.litellm_spendlogs.group_by = AsyncMock()
|
||||
mock_prisma.db.query_raw = AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"session_id": session_id,
|
||||
"api_key": api_key,
|
||||
"session_total_count": 2,
|
||||
"session_total_spend": 15.0,
|
||||
"mcp_tool_call_count": 1,
|
||||
"mcp_tool_call_spend": 10.0,
|
||||
"session_llm_count": 1,
|
||||
"session_agent_count": 0,
|
||||
}
|
||||
]
|
||||
)
|
||||
|
|
@ -3995,6 +3995,8 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
|
|||
assert rows[0]["mcp_tool_call_spend"] == 10.0
|
||||
assert rows[1]["mcp_tool_call_count"] == 1
|
||||
assert rows[1]["mcp_tool_call_spend"] == 10.0
|
||||
assert rows[0]["session_llm_count"] == 1
|
||||
assert rows[0]["session_agent_count"] == 0
|
||||
|
||||
# Every row in the session carries the full session spend, not just its own
|
||||
assert rows[0]["session_total_spend"] == 15.0
|
||||
|
|
@ -4003,13 +4005,126 @@ async def test_build_ui_spend_logs_response_dict_rows_session_counts():
|
|||
# Row without a session_id defaults to 1
|
||||
assert rows[2]["session_total_count"] == 1
|
||||
|
||||
# group_by should have been called with the session_id
|
||||
mock_prisma.db.litellm_spendlogs.group_by.assert_called_once_with(
|
||||
by=["session_id"],
|
||||
where={"session_id": {"in": [session_id]}},
|
||||
count={"session_id": True},
|
||||
# The count is folded into the single aggregate query; no separate group_by call.
|
||||
mock_prisma.db.litellm_spendlogs.group_by.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_ui_spend_logs_response_key_split_session_gets_per_key_aggregates():
|
||||
"""
|
||||
Two keys reusing one session id are separate rows under grouped pagination,
|
||||
and each row must carry ITS key's totals, never the combined session's:
|
||||
the aggregate query and its lookup are keyed by (session_id, api_key).
|
||||
"""
|
||||
from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
||||
_build_ui_spend_logs_response,
|
||||
)
|
||||
|
||||
session_id = "sess-shared"
|
||||
dict_rows = [
|
||||
{"request_id": "req-a", "session_id": session_id, "call_type": "completion", "api_key": "key-a"},
|
||||
{"request_id": "req-b", "session_id": session_id, "call_type": "completion", "api_key": "key-b"},
|
||||
]
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.query_raw = AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"session_id": session_id,
|
||||
"api_key": "key-a",
|
||||
"session_total_count": 2,
|
||||
"session_total_spend": 0.2,
|
||||
"mcp_tool_call_count": 0,
|
||||
"mcp_tool_call_spend": 0.0,
|
||||
"session_cache_hit_count": 1,
|
||||
"session_llm_count": 2,
|
||||
"session_agent_count": 0,
|
||||
},
|
||||
{
|
||||
"session_id": session_id,
|
||||
"api_key": "key-b",
|
||||
"session_total_count": 1,
|
||||
"session_total_spend": 0.7,
|
||||
"mcp_tool_call_count": 0,
|
||||
"mcp_tool_call_spend": 0.0,
|
||||
"session_cache_hit_count": 0,
|
||||
"session_llm_count": 1,
|
||||
"session_agent_count": 0,
|
||||
},
|
||||
]
|
||||
)
|
||||
|
||||
result = await _build_ui_spend_logs_response(
|
||||
prisma_client=mock_prisma,
|
||||
data=dict_rows,
|
||||
total_records=2,
|
||||
page=1,
|
||||
page_size=50,
|
||||
total_pages=1,
|
||||
enrich_session_counts=True,
|
||||
)
|
||||
|
||||
rows = result["data"]
|
||||
assert [(r["session_total_count"], r["session_total_spend"]) for r in rows] == [(2, 0.2), (1, 0.7)]
|
||||
assert [r["session_cache_hit_count"] for r in rows] == [1, 0]
|
||||
assert [r["session_llm_count"] for r in rows] == [2, 1]
|
||||
|
||||
aggregate_sql = mock_prisma.db.query_raw.mock_calls[0][1][0]
|
||||
assert "GROUP BY session_id, api_key" in aggregate_sql
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_ui_spend_logs_response_empty_api_key_keeps_session_aggregates():
|
||||
"""
|
||||
The spend-log schema defaults api_key to an empty string, which is a real
|
||||
group value and not a missing one: a multi-call session logged under an
|
||||
empty key must keep its count and spend instead of degrading to a plain
|
||||
single-call row.
|
||||
"""
|
||||
from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
||||
_build_ui_spend_logs_response,
|
||||
)
|
||||
|
||||
session_id = "sess-keyless"
|
||||
dict_rows = [
|
||||
{"request_id": "req-1", "session_id": session_id, "call_type": "completion", "api_key": ""},
|
||||
]
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.query_raw = AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"session_id": session_id,
|
||||
"api_key": "",
|
||||
"session_total_count": 3,
|
||||
"session_total_spend": 0.09,
|
||||
"mcp_tool_call_count": 0,
|
||||
"mcp_tool_call_spend": 0.0,
|
||||
"session_cache_hit_count": 0,
|
||||
"session_llm_count": 3,
|
||||
"session_agent_count": 0,
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
result = await _build_ui_spend_logs_response(
|
||||
prisma_client=mock_prisma,
|
||||
data=dict_rows,
|
||||
total_records=1,
|
||||
page=1,
|
||||
page_size=50,
|
||||
total_pages=1,
|
||||
enrich_session_counts=True,
|
||||
)
|
||||
|
||||
row = result["data"][0]
|
||||
assert row["session_total_count"] == 3
|
||||
assert row["session_total_spend"] == 0.09
|
||||
|
||||
# The empty key must reach the aggregate's authorized-keys filter too.
|
||||
_, call_args, _ = mock_prisma.db.query_raw.mock_calls[0]
|
||||
assert call_args[2] == [""]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_ui_spend_logs_response_sums_multi_round_session_spend():
|
||||
|
|
@ -4033,14 +4148,13 @@ async def test_build_ui_spend_logs_response_sums_multi_round_session_spend():
|
|||
]
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_spendlogs.group_by = AsyncMock(
|
||||
return_value=[{"session_id": session_id, "_count": {"session_id": 3}}]
|
||||
)
|
||||
# The raw aggregate query returns the full session spend (0.01 + 0.02 + 0.03).
|
||||
mock_prisma.db.query_raw = AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"session_id": session_id,
|
||||
"api_key": api_key,
|
||||
"session_total_count": 3,
|
||||
"session_total_spend": 0.06,
|
||||
"mcp_tool_call_count": 0,
|
||||
"mcp_tool_call_spend": 0.0,
|
||||
|
|
@ -4089,13 +4203,12 @@ async def test_build_ui_spend_logs_response_session_cache_hit_count():
|
|||
]
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_spendlogs.group_by = AsyncMock(
|
||||
return_value=[{"session_id": session_id, "_count": {"session_id": 2}}]
|
||||
)
|
||||
mock_prisma.db.query_raw = AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"session_id": session_id,
|
||||
"api_key": api_key,
|
||||
"session_total_count": 2,
|
||||
"session_total_spend": 0.05,
|
||||
"mcp_tool_call_count": 0,
|
||||
"mcp_tool_call_spend": 0.0,
|
||||
|
|
|
|||
|
|
@ -274,6 +274,9 @@ async def test_spend_logs_ui_uses_bounded_count_not_full_scan(monkeypatch):
|
|||
"the page query must not carry a window count that forces a full-window "
|
||||
f"scan. SQL was:\n{page_sql}"
|
||||
)
|
||||
assert "GROUP BY" not in count_sql and "DISTINCT ON" not in page_sql, (
|
||||
"without group_by_session the endpoint must keep raw per-call pagination"
|
||||
)
|
||||
|
||||
assert response["total"] == 137
|
||||
assert response["total_is_capped"] is False
|
||||
|
|
@ -499,3 +502,106 @@ async def test_global_spend_report_team_group_forwards_team_id(monkeypatch):
|
|||
params = mock_prisma.db.query_raw.call_args[0][1:]
|
||||
assert "team_x" in params, "team_id must be forwarded into the DB query params"
|
||||
assert "sl.team_id = $3" in sql, f"team query must filter on team_id. SQL was:\n{sql}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch):
|
||||
"""
|
||||
With group_by_session=true, /spend/logs/ui must page and count SESSIONS,
|
||||
not raw calls: the page query returns one representative row per session
|
||||
(DISTINCT ON the session group key, preferring non-MCP calls, newest
|
||||
first) and the bounded count counts groups. Otherwise the UI collapses a
|
||||
server page of N calls into fewer visible rows while the footer still
|
||||
claims N (issue #38060).
|
||||
"""
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.spend_tracking.spend_management_endpoints import (
|
||||
SPEND_LOGS_PAGINATION_COUNT_CAP,
|
||||
ui_view_spend_logs,
|
||||
)
|
||||
|
||||
page_rows = [
|
||||
{"request_id": "req-1", "metadata": "{}", "session_id": None},
|
||||
{"request_id": "req-2", "metadata": "{}", "session_id": None},
|
||||
]
|
||||
mock_prisma = _make_ui_spend_logs_mock(count_total=12, page_rows=page_rows)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin")
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/spend/logs/ui"
|
||||
|
||||
response = await ui_view_spend_logs(
|
||||
request=mock_request,
|
||||
api_key=None,
|
||||
user_id=None,
|
||||
request_id=None,
|
||||
start_date="2026-02-16 00:00:00",
|
||||
end_date="2026-02-16 23:59:59",
|
||||
page=1,
|
||||
page_size=50,
|
||||
sort_by="startTime",
|
||||
sort_order="desc",
|
||||
user_api_key_dict=auth,
|
||||
group_by_session=True,
|
||||
)
|
||||
|
||||
group_key = "COALESCE(NULLIF(session_id, ''), request_id), api_key"
|
||||
|
||||
count_call = mock_prisma.db.query_raw.call_args_list[0]
|
||||
count_sql = count_call[0][0]
|
||||
assert f"GROUP BY {group_key}" in count_sql, f"grouped total must count sessions. SQL was:\n{count_sql}"
|
||||
assert "COUNT(*) OVER ()" not in count_sql
|
||||
assert "LIMIT" in count_sql and "FROM (" in count_sql, "the grouped count must stay bounded"
|
||||
assert count_call[0][-1] == SPEND_LOGS_PAGINATION_COUNT_CAP + 1
|
||||
|
||||
page_sql = mock_prisma.db.query_raw.call_args_list[1][0][0]
|
||||
assert f"DISTINCT ON ({group_key})" in page_sql, f"page must return one row per session. SQL was:\n{page_sql}"
|
||||
assert f"ORDER BY {group_key}, call_type IN ('call_mcp_tool', 'list_mcp_tools'), \"startTime\" DESC" in page_sql, (
|
||||
"the session representative must prefer the newest non-MCP call"
|
||||
)
|
||||
assert "COUNT(*) OVER ()" not in page_sql
|
||||
|
||||
assert response["total"] == 12
|
||||
assert response["total_is_capped"] is False
|
||||
assert response["total_pages"] == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_logs_ui_request_id_lookup_with_grouping_returns_exact_row(monkeypatch):
|
||||
"""
|
||||
A request_id lookup with group_by_session=true must still resolve the
|
||||
exact requested row: the filter runs before grouping, so the row is its
|
||||
own group's representative and deep links keep working.
|
||||
"""
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.spend_tracking.spend_management_endpoints import ui_view_spend_logs
|
||||
|
||||
target_row = {"request_id": "req-deep-link", "metadata": "{}", "session_id": None}
|
||||
mock_prisma = _make_ui_spend_logs_mock(count_total=1, page_rows=[target_row])
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
||||
auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin")
|
||||
mock_request = MagicMock()
|
||||
mock_request.url.path = "/spend/logs/ui"
|
||||
|
||||
response = await ui_view_spend_logs(
|
||||
request=mock_request,
|
||||
api_key=None,
|
||||
user_id=None,
|
||||
request_id="req-deep-link",
|
||||
start_date=None,
|
||||
end_date=None,
|
||||
page=1,
|
||||
page_size=1,
|
||||
sort_by="startTime",
|
||||
sort_order="desc",
|
||||
user_api_key_dict=auth,
|
||||
group_by_session=True,
|
||||
)
|
||||
|
||||
page_call = mock_prisma.db.query_raw.call_args_list[1]
|
||||
assert "request_id = $" in page_call[0][0], "the request_id equality filter must survive grouping"
|
||||
assert "req-deep-link" in page_call[0]
|
||||
assert [row["request_id"] for row in response["data"]] == ["req-deep-link"]
|
||||
assert response["total"] == 1
|
||||
|
|
|
|||
|
|
@ -3013,7 +3013,7 @@ class TestHandleLLMApiExceptionDictDetail:
|
|||
assert "NotFoundError" in proxy_exc.message
|
||||
|
||||
async def test_exception_with_status_code_propagates(self):
|
||||
"""Exception with a statically-set status_code should propagate it."""
|
||||
"""Exception with a statically-set status_code should propagate it and its message."""
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAIError
|
||||
|
||||
exc = VertexAIError(
|
||||
|
|
@ -3022,12 +3022,30 @@ class TestHandleLLMApiExceptionDictDetail:
|
|||
)
|
||||
proxy_exc = await self._invoke(exc)
|
||||
assert proxy_exc.code == "429"
|
||||
assert proxy_exc.message == "Rate limit exceeded"
|
||||
|
||||
async def test_exception_without_status_code_defaults_to_500(self):
|
||||
"""Exception with no status_code attribute defaults to 500."""
|
||||
"""Exception with no status_code attribute defaults to 500; a message with nothing
|
||||
to redact still reaches the client, since routes raise plain exceptions as validation text."""
|
||||
exc = ValueError("Something broke")
|
||||
proxy_exc = await self._invoke(exc)
|
||||
assert proxy_exc.code == "500"
|
||||
assert proxy_exc.message == "Something broke"
|
||||
|
||||
async def test_unclassified_exception_redacts_internal_details_from_client_message(self):
|
||||
"""Regression for LIT-6747: an unclassified exception's credential, path, and host
|
||||
must not reach the client."""
|
||||
exc = RuntimeError(
|
||||
"Failed to connect to postgresql://litellm_internal:S3cr3tPGPass@10.20.30.40:5432/litellm_prod "
|
||||
"(config file /etc/litellm/secrets/db.yaml)"
|
||||
)
|
||||
proxy_exc = await self._invoke(exc)
|
||||
assert proxy_exc.code == "500"
|
||||
assert "S3cr3tPGPass" not in proxy_exc.message
|
||||
assert "litellm_internal" not in proxy_exc.message
|
||||
assert "10.20.30.40" not in proxy_exc.message
|
||||
assert "/etc/litellm/secrets/db.yaml" not in proxy_exc.message
|
||||
assert "REDACTED" in proxy_exc.message
|
||||
|
||||
async def test_already_normalized_proxy_exception_is_honored(self):
|
||||
"""A ProxyException raised mid-request (e.g. a guardrail block) is already
|
||||
|
|
@ -3244,6 +3262,42 @@ class TestStreamCloseOnDisconnect:
|
|||
|
||||
assert upstream.aclosed
|
||||
|
||||
async def test_async_streaming_data_generator_redacts_internal_details_on_error(
|
||||
self,
|
||||
):
|
||||
"""Regression for LIT-6747: a mid-stream exception must not hand its raw text or a
|
||||
traceback to serialize_error."""
|
||||
|
||||
class FailingUpstream:
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
raise RuntimeError(
|
||||
"Failed to connect to postgresql://litellm_internal:S3cr3tPGPass@10.20.30.40:5432/litellm_prod "
|
||||
"(config file /etc/litellm/secrets/db.yaml)"
|
||||
)
|
||||
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
captured: list = []
|
||||
gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
|
||||
response=FailingUpstream(),
|
||||
user_api_key_dict=ProxyUserAPIKeyAuth(api_key="sk-test"),
|
||||
request_data={"model": "mock-model"},
|
||||
proxy_logging_obj=ProxyLogging(user_api_key_cache=MagicMock()),
|
||||
serialize_chunk=lambda c: "data: x\n\n",
|
||||
serialize_error=lambda e: captured.append(e) or "data: error\n\n",
|
||||
)
|
||||
|
||||
await gen.__anext__()
|
||||
|
||||
assert len(captured) == 1
|
||||
message = captured[0].message
|
||||
assert "S3cr3tPGPass" not in message
|
||||
assert "10.20.30.40" not in message
|
||||
assert "/etc/litellm/secrets/db.yaml" not in message
|
||||
assert "Traceback (most recent call last)" not in message
|
||||
|
||||
@staticmethod
|
||||
def _request_that_disconnects() -> Request:
|
||||
async def receive():
|
||||
|
|
|
|||
|
|
@ -95,6 +95,7 @@ class TestProxyInitializationHelpers:
|
|||
assert args["app"] == "litellm.proxy.proxy_server:app"
|
||||
assert args["host"] == "localhost"
|
||||
assert args["port"] == 8000
|
||||
assert args["server_header"] is False
|
||||
|
||||
# Test with log_config
|
||||
args = ProxyInitializationHelpers._get_default_unvicorn_init_args(
|
||||
|
|
|
|||
|
|
@ -2,29 +2,24 @@ from datetime import datetime, timezone
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import Request
|
||||
|
||||
|
||||
from fastapi import HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
LiteLLM_ManagedVectorStore,
|
||||
)
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
LiteLLMVectorStoreEmbeddingExecutor,
|
||||
RouterVectorStoreEmbeddingExecutor,
|
||||
)
|
||||
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.vector_store_endpoints.endpoints import (
|
||||
_update_request_data_with_litellm_managed_vector_store_registry,
|
||||
index_create,
|
||||
index_list,
|
||||
)
|
||||
from litellm.proxy.vector_store_files_endpoints.endpoints import (
|
||||
_update_request_data_with_model_routing_hint,
|
||||
)
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import (
|
||||
_check_vector_store_access,
|
||||
_resolve_embedding_config,
|
||||
_resolve_embedding_config_from_db,
|
||||
_resolve_embedding_config_from_router,
|
||||
create_vector_store_in_db,
|
||||
new_vector_store,
|
||||
)
|
||||
|
|
@ -33,8 +28,12 @@ from litellm.proxy.vector_store_endpoints.utils import (
|
|||
is_allowed_to_call_vector_store_endpoint,
|
||||
is_allowed_to_call_vector_store_files_endpoint,
|
||||
)
|
||||
from litellm.proxy.vector_store_files_endpoints.endpoints import (
|
||||
_update_request_data_with_model_routing_hint,
|
||||
)
|
||||
from litellm.types.utils import EmbeddingResponse, LlmProviders
|
||||
from litellm.types.vector_stores import IndexCreateRequest, IndexListResponse
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.vector_stores.main import _direct_vector_store_embedding_executor
|
||||
|
||||
|
||||
def _serialize_litellm_params(litellm_params):
|
||||
|
|
@ -51,17 +50,113 @@ def _serialize_litellm_params(litellm_params):
|
|||
return json.dumps(litellm_params or {})
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_embedding_config_cache():
|
||||
"""The use-time embedding-config resolver caches results in process
|
||||
memory across calls. Reset it before every test so the resolver
|
||||
actually exercises the router/DB path under test instead of returning
|
||||
a value cached by an earlier test."""
|
||||
from litellm.proxy.vector_store_endpoints import management_endpoints
|
||||
def test_direct_vector_store_embedding_executor_rejects_invalid_value():
|
||||
with pytest.raises(TypeError, match="Invalid direct vector store embedding executor"):
|
||||
_direct_vector_store_embedding_executor(object(), None, {})
|
||||
|
||||
management_endpoints._embedding_config_cache = None
|
||||
yield
|
||||
management_endpoints._embedding_config_cache = None
|
||||
|
||||
def test_router_vector_store_search_injects_executor_and_request_metadata():
|
||||
router = litellm.Router(model_list=[])
|
||||
original = MagicMock(return_value="searched")
|
||||
wrapped = router.factory_function(original, call_type="vector_store_search")
|
||||
|
||||
assert (
|
||||
wrapped(
|
||||
vector_store_id="store",
|
||||
query="query",
|
||||
custom_llm_provider="valkey",
|
||||
litellm_metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
== "searched"
|
||||
)
|
||||
|
||||
call_kwargs = original.call_args.kwargs
|
||||
assert call_kwargs["custom_llm_provider"] == "valkey"
|
||||
executor = call_kwargs["_direct_vector_store_embedding_executor"]
|
||||
assert isinstance(executor, RouterVectorStoreEmbeddingExecutor)
|
||||
assert executor.metadata == {"user_api_key_team_id": "team-a"}
|
||||
assert litellm.Router._vector_store_request_metadata({"metadata": {"user_api_key_team_id": "team-b"}}) == {
|
||||
"user_api_key_team_id": "team-b"
|
||||
}
|
||||
assert litellm.Router._vector_store_request_metadata({}) == {}
|
||||
|
||||
with patch.object( # test-quality-ok: fallback dispatch is the boundary this wrapper delegates to
|
||||
router, "_generic_api_call_with_fallbacks", return_value="routed"
|
||||
) as fallback:
|
||||
assert wrapped(model="vector-alias", vector_store_id="store", query="query") == "routed"
|
||||
assert fallback.call_args.kwargs["model"] == "vector-alias"
|
||||
assert fallback.call_args.kwargs["original_function"] is original
|
||||
|
||||
create_original = MagicMock(return_value="created")
|
||||
wrapped_create = router.factory_function(create_original, call_type="vector_store_create")
|
||||
assert wrapped_create(name="store") == "created"
|
||||
create_original.assert_called_once_with(name="store")
|
||||
with patch.object( # test-quality-ok: fallback dispatch is the boundary this wrapper delegates to
|
||||
router, "_generic_api_call_with_fallbacks", return_value="created-through-router"
|
||||
) as fallback:
|
||||
assert wrapped_create(model="vector-alias", name="store") == "created-through-router"
|
||||
fallback.assert_called_once_with(original_function=create_original, model="vector-alias", name="store")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_embedding_executors_preserve_explicit_configuration():
|
||||
response = EmbeddingResponse(data=[{"embedding": [0.1], "index": 0, "object": "embedding"}])
|
||||
sdk_executor = LiteLLMVectorStoreEmbeddingExecutor()
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: isolates SDK dispatch from external embedding providers
|
||||
"litellm.embedding", return_value=response
|
||||
) as embedding,
|
||||
patch( # test-quality-ok: isolates async SDK dispatch from external embedding providers
|
||||
"litellm.aembedding", new=AsyncMock(return_value=response)
|
||||
) as aembedding,
|
||||
):
|
||||
assert sdk_executor.embed("openai/model", "sync", {"api_key": "explicit"}) is response
|
||||
assert await sdk_executor.aembed("openai/model", "async", {"api_key": "explicit"}) is response
|
||||
|
||||
embedding.assert_called_once_with(model="openai/model", input=["sync"], api_key="explicit")
|
||||
aembedding.assert_awaited_once_with(model="openai/model", input=["async"], api_key="explicit")
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.embedding.return_value = response
|
||||
mock_router.aembedding = AsyncMock(return_value=response)
|
||||
router_executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=mock_router,
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
assert router_executor.embed("team-alias", "query", {}) is response
|
||||
mock_router.embedding.assert_called_once_with(
|
||||
model="team-alias",
|
||||
input=["query"],
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: verifies explicit store configuration at the SDK boundary
|
||||
"litellm.embedding", return_value=response
|
||||
) as explicit_embedding,
|
||||
patch( # test-quality-ok: verifies async explicit store configuration at the SDK boundary
|
||||
"litellm.aembedding", new=AsyncMock(return_value=response)
|
||||
) as explicit_aembedding,
|
||||
):
|
||||
assert router_executor.embed("openai/model", "query", {"api_key": "store-key"}) is response
|
||||
assert await router_executor.aembed("openai/model", "query", {"api_key": "store-key"}) is response
|
||||
|
||||
explicit_embedding.assert_not_called()
|
||||
explicit_aembedding.assert_not_awaited()
|
||||
assert mock_router.embedding.call_args.kwargs == {
|
||||
"model": "openai/model",
|
||||
"input": ["query"],
|
||||
"api_key": "store-key",
|
||||
"metadata": {"user_api_key_team_id": "team-a"},
|
||||
}
|
||||
mock_router.aembedding.assert_awaited_once_with(
|
||||
model="openai/model",
|
||||
input=["query"],
|
||||
api_key="store-key",
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -82,10 +177,11 @@ async def test_router_avector_store_search_passes_correct_args():
|
|||
}
|
||||
|
||||
# Call router's avector_store_search
|
||||
result = await router.avector_store_search(
|
||||
await router.avector_store_search(
|
||||
vector_store_id="test_store_id",
|
||||
query="test query",
|
||||
custom_llm_provider="bedrock",
|
||||
metadata={"user_api_key_team_id": "team-a"},
|
||||
)
|
||||
|
||||
# Verify the internal method was called with correct args
|
||||
|
|
@ -96,6 +192,38 @@ async def test_router_avector_store_search_passes_correct_args():
|
|||
assert call_args[1]["vector_store_id"] == "test_store_id"
|
||||
assert call_args[1]["query"] == "test query"
|
||||
assert call_args[1]["custom_llm_provider"] == "bedrock"
|
||||
executor = call_args[1]["_direct_vector_store_embedding_executor"]
|
||||
assert isinstance(executor, RouterVectorStoreEmbeddingExecutor)
|
||||
assert executor.metadata["user_api_key_team_id"] == "team-a"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_embedding_executor_uses_team_scoped_router_deployment():
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "shared-embedding",
|
||||
"litellm_params": {"model": "openai/text-embedding-3-small", "api_key": "team-a-key"},
|
||||
"model_info": {"team_id": "team-a", "team_public_model_name": "shared-embedding"},
|
||||
},
|
||||
{
|
||||
"model_name": "shared-embedding",
|
||||
"litellm_params": {"model": "openai/text-embedding-3-small", "api_key": "team-b-key"},
|
||||
"model_info": {"team_id": "team-b", "team_public_model_name": "shared-embedding"},
|
||||
},
|
||||
]
|
||||
)
|
||||
executor = RouterVectorStoreEmbeddingExecutor(
|
||||
router=router,
|
||||
metadata={"user_api_key_team_id": "team-b"},
|
||||
)
|
||||
response = EmbeddingResponse(data=[{"embedding": [0.1], "index": 0, "object": "embedding"}])
|
||||
|
||||
with patch("litellm.aembedding", new=AsyncMock(return_value=response)) as mock_aembedding:
|
||||
result = await executor.aembed("shared-embedding", "query", {})
|
||||
|
||||
assert result is response
|
||||
assert mock_aembedding.await_args.kwargs["api_key"] == "team-b-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -502,91 +630,30 @@ async def test_update_request_data_with_litellm_managed_vector_store_registry():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_request_data_resolves_embedding_config_at_use_time():
|
||||
"""When the persisted vector store row carries only a
|
||||
``litellm_embedding_model`` reference (the new behaviour after
|
||||
moving the auto-resolve out of write time), the request-handling
|
||||
layer must resolve the embedding config so the downstream embed
|
||||
call still has ``api_key`` / ``api_base`` / ``api_version``. The
|
||||
resolved config lives in this per-request data dict only — never
|
||||
persisted."""
|
||||
mock_vector_store: LiteLLM_ManagedVectorStore = {
|
||||
async def test_managed_vector_store_keeps_embedding_reference_and_explicit_config():
|
||||
explicit_config = {"api_key": "store-specific-key", "api_base": "https://embedding.example"}
|
||||
managed_vector_store: LiteLLM_ManagedVectorStore = {
|
||||
"vector_store_id": "test_store",
|
||||
"custom_llm_provider": "azure_ai",
|
||||
"custom_llm_provider": "valkey",
|
||||
"litellm_params": {
|
||||
"litellm_embedding_model": "azure/text-embedding-3-large",
|
||||
# Note: no litellm_embedding_config persisted
|
||||
"litellm_embedding_model": "team-embedding-alias",
|
||||
"litellm_embedding_config": explicit_config,
|
||||
},
|
||||
}
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = (
|
||||
mock_vector_store
|
||||
)
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = managed_vector_store
|
||||
|
||||
resolved = {
|
||||
"api_key": "use-time-resolved-key",
|
||||
"api_base": "https://my-azure.example",
|
||||
"api_version": "2024-09-01",
|
||||
}
|
||||
|
||||
with (
|
||||
patch.object(litellm, "vector_store_registry", mock_registry),
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints._resolve_embedding_config",
|
||||
new=AsyncMock(return_value=resolved),
|
||||
),
|
||||
):
|
||||
with patch.object(litellm, "vector_store_registry", mock_registry):
|
||||
result = await _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
data={}, vector_store_id="test_store"
|
||||
data={},
|
||||
vector_store_id="test_store",
|
||||
)
|
||||
|
||||
assert result["litellm_embedding_model"] == "azure/text-embedding-3-large"
|
||||
assert result["litellm_embedding_config"] == resolved
|
||||
assert result["litellm_embedding_model"] == "team-embedding-alias"
|
||||
assert result["litellm_embedding_config"] == explicit_config
|
||||
assert managed_vector_store["litellm_params"]["litellm_embedding_config"] == explicit_config
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_request_data_passes_through_legacy_embedding_config():
|
||||
"""A vector store row created by an older proxy version may already
|
||||
carry a fully-resolved ``litellm_embedding_config`` in its persisted
|
||||
``litellm_params`` (the very leak this PR closes). Those legacy rows
|
||||
must still work — the use-time resolver skips re-resolution when
|
||||
the config is already present so the embed call keeps succeeding."""
|
||||
legacy_config = {
|
||||
"api_key": "legacy-cleartext-key",
|
||||
"api_base": "https://legacy-azure.example",
|
||||
"api_version": "2024-01-01",
|
||||
}
|
||||
mock_vector_store: LiteLLM_ManagedVectorStore = {
|
||||
"vector_store_id": "legacy_store",
|
||||
"custom_llm_provider": "azure_ai",
|
||||
"litellm_params": {
|
||||
"litellm_embedding_model": "azure/text-embedding-3-large",
|
||||
"litellm_embedding_config": legacy_config,
|
||||
},
|
||||
}
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.get_litellm_managed_vector_store_from_registry.return_value = (
|
||||
mock_vector_store
|
||||
)
|
||||
|
||||
resolve_mock = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(litellm, "vector_store_registry", mock_registry),
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.endpoints._resolve_embedding_config",
|
||||
new=resolve_mock,
|
||||
),
|
||||
):
|
||||
result = await _update_request_data_with_litellm_managed_vector_store_registry(
|
||||
data={}, vector_store_id="legacy_store"
|
||||
)
|
||||
|
||||
assert result["litellm_embedding_config"] == legacy_config
|
||||
resolve_mock.assert_not_awaited()
|
||||
|
||||
|
||||
class TestCheckVectorStorePermission:
|
||||
"""Test suite for check_vector_store_permission function."""
|
||||
|
|
@ -2003,57 +2070,7 @@ async def test_vector_store_update_and_list_synchronization():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_embedding_config_from_db():
|
||||
"""Test that _resolve_embedding_config_from_db correctly resolves embedding config from database."""
|
||||
mock_prisma_client = MagicMock()
|
||||
|
||||
# Mock database model with litellm_params
|
||||
mock_db_model = MagicMock()
|
||||
mock_db_model.litellm_params = {
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://api.openai.com",
|
||||
"api_version": "2024-01-01",
|
||||
}
|
||||
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock(
|
||||
return_value=mock_db_model
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.decrypt_value_helper",
|
||||
side_effect=lambda value, key, return_original_value: value,
|
||||
):
|
||||
result = await _resolve_embedding_config_from_db(
|
||||
embedding_model="text-embedding-ada-002", prisma_client=mock_prisma_client
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "test-api-key"
|
||||
assert result["api_base"] == "https://api.openai.com"
|
||||
assert result["api_version"] == "2024-01-01"
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first.assert_called_once_with(
|
||||
where={"model_name": "text-embedding-ada-002"}
|
||||
)
|
||||
|
||||
# Test with empty embedding_model
|
||||
result_empty = await _resolve_embedding_config_from_db(
|
||||
embedding_model="", prisma_client=mock_prisma_client
|
||||
)
|
||||
assert result_empty is None
|
||||
|
||||
# Test with model not found
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
result_not_found = await _resolve_embedding_config_from_db(
|
||||
embedding_model="non-existent-model", prisma_client=mock_prisma_client
|
||||
)
|
||||
assert result_not_found is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_vector_store_auto_resolves_embedding_config():
|
||||
"""Test that new_vector_store auto-resolves embedding config when embedding_model is provided but config is not."""
|
||||
async def test_new_vector_store_persists_embedding_reference_without_credentials():
|
||||
import json
|
||||
|
||||
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
|
||||
|
|
@ -2070,14 +2087,6 @@ async def test_new_vector_store_auto_resolves_embedding_config():
|
|||
},
|
||||
}
|
||||
|
||||
# Mock database model lookup for embedding config resolution
|
||||
mock_db_model = MagicMock()
|
||||
mock_db_model.litellm_params = {
|
||||
"api_key": "resolved-api-key",
|
||||
"api_base": "https://api.openai.com",
|
||||
"api_version": "2024-01-01",
|
||||
}
|
||||
|
||||
# Mock user API key
|
||||
mock_user_api_key = MagicMock(spec=UserAPIKeyAuth)
|
||||
mock_user_api_key.user_role = None
|
||||
|
|
@ -2088,10 +2097,6 @@ async def test_new_vector_store_auto_resolves_embedding_config():
|
|||
mock_prisma_client.db.litellm_managedvectorstorestable.find_unique = AsyncMock(
|
||||
return_value=None # Vector store doesn't exist yet
|
||||
)
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock(
|
||||
return_value=mock_db_model
|
||||
)
|
||||
|
||||
# Track what was passed to create
|
||||
captured_create_data = {}
|
||||
|
||||
|
|
@ -2112,261 +2117,21 @@ async def test_new_vector_store_auto_resolves_embedding_config():
|
|||
mock_registry = MagicMock()
|
||||
mock_registry.add_vector_store_to_registry = MagicMock()
|
||||
|
||||
# Mock router to return None (so it falls back to DB resolution)
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment_by_model_group_name.return_value = None
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
||||
patch("litellm.proxy.proxy_server.llm_router", mock_router),
|
||||
patch(
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.decrypt_value_helper",
|
||||
side_effect=lambda value, key, return_original_value: value,
|
||||
),
|
||||
patch.object(litellm, "vector_store_registry", mock_registry),
|
||||
):
|
||||
result = await new_vector_store(
|
||||
vector_store=vector_store_data, user_api_key_dict=mock_user_api_key
|
||||
)
|
||||
result = await new_vector_store(vector_store=vector_store_data, user_api_key_dict=mock_user_api_key)
|
||||
|
||||
assert result["status"] == "success"
|
||||
# Auto-resolve no longer happens at create time — the persisted row
|
||||
# carries only the model reference, never the resolved cleartext
|
||||
# credential. Resolution now happens at request-handling time inside
|
||||
# ``_update_request_data_with_litellm_managed_vector_store_registry``,
|
||||
# where the resolved config lives in per-request memory and is never
|
||||
# written to the database.
|
||||
litellm_params_json = captured_create_data.get("litellm_params")
|
||||
assert litellm_params_json is not None
|
||||
litellm_params_dict = json.loads(litellm_params_json)
|
||||
assert "litellm_embedding_config" not in litellm_params_dict
|
||||
assert litellm_params_dict["litellm_embedding_model"] == "text-embedding-ada-002"
|
||||
|
||||
# The response must also not echo a cleartext credential — even on
|
||||
# the create response, where redaction guards against caller-supplied
|
||||
# cleartext or pre-existing rows that were created by an earlier
|
||||
# proxy version.
|
||||
response_vs = result["vector_store"]
|
||||
assert "resolved-api-key" not in _serialize_litellm_params(
|
||||
response_vs.get("litellm_params")
|
||||
)
|
||||
|
||||
|
||||
def test_resolve_embedding_config_from_router():
|
||||
"""Test that _resolve_embedding_config_from_router correctly extracts credentials from config-defined models."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
# Create a mock router with a model
|
||||
mock_router = MagicMock()
|
||||
|
||||
# Create a mock deployment with litellm_params
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "config-api-key"
|
||||
mock_litellm_params.api_base = "https://config-api-base.com"
|
||||
mock_litellm_params.api_version = "2024-02-01"
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
|
||||
mock_router.get_deployment_by_model_group_name.return_value = mock_deployment
|
||||
|
||||
# Test resolution
|
||||
result = _resolve_embedding_config_from_router(
|
||||
embedding_model="text-embedding-ada-002", llm_router=mock_router
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "config-api-key"
|
||||
assert result["api_base"] == "https://config-api-base.com"
|
||||
assert result["api_version"] == "2024-02-01"
|
||||
|
||||
mock_router.get_deployment_by_model_group_name.assert_called_once_with(
|
||||
model_group_name="text-embedding-ada-002"
|
||||
)
|
||||
|
||||
|
||||
def test_resolve_embedding_config_from_router_with_provider_prefix():
|
||||
"""Test that _resolve_embedding_config_from_router handles provider prefixes like 'azure/model-name'."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
# Create a mock router
|
||||
mock_router = MagicMock()
|
||||
|
||||
# Create a mock deployment
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "azure-api-key"
|
||||
mock_litellm_params.api_base = "https://azure-endpoint.openai.azure.com"
|
||||
mock_litellm_params.api_version = "2024-02-15"
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
|
||||
# First call with full name returns None, second call with stripped name returns deployment
|
||||
mock_router.get_deployment_by_model_group_name.side_effect = [None, mock_deployment]
|
||||
|
||||
result = _resolve_embedding_config_from_router(
|
||||
embedding_model="azure/text-embedding-3-large", llm_router=mock_router
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "azure-api-key"
|
||||
assert result["api_base"] == "https://azure-endpoint.openai.azure.com"
|
||||
assert result["api_version"] == "2024-02-15"
|
||||
|
||||
# Should have tried both the full name and stripped name
|
||||
assert mock_router.get_deployment_by_model_group_name.call_count == 2
|
||||
|
||||
|
||||
def test_resolve_embedding_config_from_router_returns_none_when_not_found():
|
||||
"""Test that _resolve_embedding_config_from_router returns None when model is not in router."""
|
||||
mock_router = MagicMock()
|
||||
mock_router.get_deployment_by_model_group_name.return_value = None
|
||||
|
||||
result = _resolve_embedding_config_from_router(
|
||||
embedding_model="nonexistent-model", llm_router=mock_router
|
||||
)
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_resolve_embedding_config_from_router_handles_os_environ():
|
||||
"""Test that _resolve_embedding_config_from_router handles os.environ/ prefixed values."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
mock_router = MagicMock()
|
||||
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "os.environ/OPENAI_API_KEY"
|
||||
mock_litellm_params.api_base = "https://direct-url.com"
|
||||
mock_litellm_params.api_version = None
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
|
||||
mock_router.get_deployment_by_model_group_name.return_value = mock_deployment
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.get_secret",
|
||||
return_value="resolved-from-env",
|
||||
) as mock_get_secret:
|
||||
result = _resolve_embedding_config_from_router(
|
||||
embedding_model="text-embedding-ada-002", llm_router=mock_router
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "resolved-from-env"
|
||||
assert result["api_base"] == "https://direct-url.com"
|
||||
assert "api_version" not in result
|
||||
|
||||
mock_get_secret.assert_called_once_with("os.environ/OPENAI_API_KEY")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_embedding_config_tries_router_then_db():
|
||||
"""Test that _resolve_embedding_config tries router first, then falls back to DB."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_router = MagicMock()
|
||||
|
||||
# Router has the model
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "router-api-key"
|
||||
mock_litellm_params.api_base = "https://router-api-base.com"
|
||||
mock_litellm_params.api_version = None
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
|
||||
mock_router.get_deployment_by_model_group_name.return_value = mock_deployment
|
||||
|
||||
# DB should NOT be called since router has the model
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock()
|
||||
|
||||
result = await _resolve_embedding_config(
|
||||
embedding_model="text-embedding-ada-002",
|
||||
prisma_client=mock_prisma_client,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "router-api-key"
|
||||
|
||||
# DB should NOT have been called since router found the model
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_embedding_config_caches_result():
|
||||
"""The first lookup should hit the router/DB; subsequent lookups for
|
||||
the same model name should return the cached value without touching
|
||||
the router or the database."""
|
||||
from litellm.types.router import Deployment, LiteLLM_Params
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_router = MagicMock()
|
||||
|
||||
mock_litellm_params = MagicMock(spec=LiteLLM_Params)
|
||||
mock_litellm_params.api_key = "router-api-key"
|
||||
mock_litellm_params.api_base = "https://router-api-base.com"
|
||||
mock_litellm_params.api_version = None
|
||||
|
||||
mock_deployment = MagicMock(spec=Deployment)
|
||||
mock_deployment.litellm_params = mock_litellm_params
|
||||
mock_router.get_deployment_by_model_group_name.return_value = mock_deployment
|
||||
|
||||
first = await _resolve_embedding_config(
|
||||
embedding_model="cached-model",
|
||||
prisma_client=mock_prisma_client,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
assert first is not None
|
||||
assert mock_router.get_deployment_by_model_group_name.call_count == 1
|
||||
|
||||
second = await _resolve_embedding_config(
|
||||
embedding_model="cached-model",
|
||||
prisma_client=mock_prisma_client,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
assert second == first
|
||||
# Router (and by extension the DB) was not consulted again.
|
||||
assert mock_router.get_deployment_by_model_group_name.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolve_embedding_config_falls_back_to_db():
|
||||
"""Test that _resolve_embedding_config falls back to DB when router doesn't have the model."""
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_router = MagicMock()
|
||||
|
||||
# Router doesn't have the model
|
||||
mock_router.get_deployment_by_model_group_name.return_value = None
|
||||
|
||||
# DB has the model
|
||||
mock_db_model = MagicMock()
|
||||
mock_db_model.litellm_params = {
|
||||
"api_key": "db-api-key",
|
||||
"api_base": "https://db-api-base.com",
|
||||
}
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first = AsyncMock(
|
||||
return_value=mock_db_model
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.decrypt_value_helper",
|
||||
side_effect=lambda value, key, return_original_value: value,
|
||||
):
|
||||
result = await _resolve_embedding_config(
|
||||
embedding_model="text-embedding-ada-002",
|
||||
prisma_client=mock_prisma_client,
|
||||
llm_router=mock_router,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert result["api_key"] == "db-api-key"
|
||||
|
||||
# DB should have been called since router didn't find the model
|
||||
mock_prisma_client.db.litellm_proxymodeltable.find_first.assert_called()
|
||||
assert "api_key" not in _serialize_litellm_params(response_vs.get("litellm_params"))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -2425,9 +2190,7 @@ async def test_new_vector_store_auto_resolves_from_router():
|
|||
}
|
||||
return mock_created_vector_store
|
||||
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock(
|
||||
side_effect=mock_create
|
||||
)
|
||||
mock_prisma_client.db.litellm_managedvectorstorestable.create = AsyncMock(side_effect=mock_create)
|
||||
|
||||
mock_registry = MagicMock()
|
||||
mock_registry.add_vector_store_to_registry = MagicMock()
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ completion_start_time = end_time."""
|
|||
import json
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from unittest.mock import Mock
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -378,3 +378,162 @@ def test_stamp_responses_usage_cost_survives_calculator_failure():
|
|||
_stamp_responses_usage_cost(response, logging_obj)
|
||||
|
||||
assert getattr(response.usage, "cost", None) is None
|
||||
|
||||
|
||||
def _capture_dispatch(logged: list):
|
||||
"""Record the object handed to the success handlers.
|
||||
|
||||
``Mock(spec=LiteLLMLoggingObj).dispatch_success_handlers`` is an AsyncMock whose side effect
|
||||
only runs when the coroutine is awaited, so capture with a plain function instead.
|
||||
"""
|
||||
|
||||
async def _noop() -> None:
|
||||
return None
|
||||
|
||||
def _dispatch(result, **kwargs):
|
||||
logged.append(result)
|
||||
return _noop()
|
||||
|
||||
return _dispatch
|
||||
|
||||
|
||||
def _headers_config(*, transform_hidden_params: Optional[dict] = None) -> Mock:
|
||||
"""Config whose completed event carries a real ResponsesAPIResponse, so the logging copy
|
||||
performs a genuine model_dump/model_validate round trip."""
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
|
||||
def _transform(model, parsed_chunk, logging_obj):
|
||||
evt_type = parsed_chunk.get("type")
|
||||
if evt_type != "response.completed":
|
||||
stub = Mock()
|
||||
stub.type = evt_type
|
||||
return stub
|
||||
response = ResponsesAPIResponse(
|
||||
id="resp_headers",
|
||||
created_at=1,
|
||||
output=[],
|
||||
parallel_tool_calls=False,
|
||||
tool_choice="auto",
|
||||
tools=[],
|
||||
)
|
||||
if transform_hidden_params is not None:
|
||||
response._hidden_params.update(transform_hidden_params)
|
||||
return ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=response,
|
||||
)
|
||||
|
||||
mock_config.transform_streaming_response.side_effect = _transform
|
||||
return mock_config
|
||||
|
||||
|
||||
def _make_header_iterator(
|
||||
*,
|
||||
headers: dict,
|
||||
config: Mock,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ResponsesAPIStreamingIterator:
|
||||
async def aiter_bytes():
|
||||
yield _sse_event({"type": "response.completed"})
|
||||
|
||||
mock_response = Mock()
|
||||
mock_response.headers = headers
|
||||
mock_response.aiter_bytes = aiter_bytes
|
||||
|
||||
return ResponsesAPIStreamingIterator(
|
||||
response=mock_response,
|
||||
model="gpt-4o-mini",
|
||||
responses_api_provider_config=config,
|
||||
logging_obj=logging_obj,
|
||||
litellm_metadata={},
|
||||
custom_llm_provider="azure",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_logging_response_carries_provider_response_headers():
|
||||
"""LIT-6055: the provider headers the iterator captured must reach the logged response, so
|
||||
custom loggers can read Azure's apim-request-id from the callback payload."""
|
||||
logging_obj = _logging_obj_stub()
|
||||
logged: list[object] = []
|
||||
logging_obj.dispatch_success_handlers = _capture_dispatch(logged)
|
||||
|
||||
logging_obj._on_deferred_stream_complete = None
|
||||
|
||||
iterator = _make_header_iterator(
|
||||
headers={"apim-request-id": "azure-correlation-1", "x-ms-region": "East US 2"},
|
||||
config=_headers_config(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
async for _ in iterator:
|
||||
pass
|
||||
|
||||
assert len(logged) == 1
|
||||
hidden_params = logged[0].response._hidden_params
|
||||
assert hidden_params["additional_headers"]["llm_provider-apim-request-id"] == "azure-correlation-1"
|
||||
assert hidden_params["additional_headers"]["llm_provider-x-ms-region"] == "East US 2"
|
||||
assert hidden_params["headers"]["apim-request-id"] == "azure-correlation-1"
|
||||
# the proxy builds the client's response headers from the iterator's own dict, so the logged
|
||||
# response must hold copies rather than alias it
|
||||
assert hidden_params["additional_headers"] is not iterator._hidden_params["additional_headers"]
|
||||
assert hidden_params["headers"] is not iterator._raw_response_headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_logging_copy_preserves_transform_hidden_params():
|
||||
"""LIT-6055: model_validate(model_dump()) drops pydantic private attributes, so headers a
|
||||
provider transform already set on the response (fake_stream) must be re-applied."""
|
||||
logging_obj = _logging_obj_stub()
|
||||
logged: list[object] = []
|
||||
logging_obj.dispatch_success_handlers = _capture_dispatch(logged)
|
||||
|
||||
logging_obj._on_deferred_stream_complete = None
|
||||
|
||||
iterator = _make_header_iterator(
|
||||
headers={},
|
||||
config=_headers_config(
|
||||
transform_hidden_params={
|
||||
"additional_headers": {"llm_provider-apim-request-id": "from-transform"},
|
||||
"headers": {"apim-request-id": "from-transform"},
|
||||
"response_cost": 0.5,
|
||||
}
|
||||
),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
async for _ in iterator:
|
||||
pass
|
||||
|
||||
assert len(logged) == 1
|
||||
hidden_params = logged[0].response._hidden_params
|
||||
assert hidden_params["additional_headers"]["llm_provider-apim-request-id"] == "from-transform"
|
||||
assert hidden_params["headers"]["apim-request-id"] == "from-transform"
|
||||
assert iterator.completed_response is not logged[0]
|
||||
# only the header keys travel: response_cost would short-circuit the cost calculator
|
||||
assert "response_cost" not in hidden_params
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_logging_copy_fallback_leaves_caller_event_untouched():
|
||||
"""LIT-6055: when the logging copy falls back to the original event, the header restore must
|
||||
not stamp logging-only state onto the object the caller is iterating."""
|
||||
logging_obj = _logging_obj_stub()
|
||||
logged: list[object] = []
|
||||
logging_obj.dispatch_success_handlers = _capture_dispatch(logged)
|
||||
logging_obj._on_deferred_stream_complete = None
|
||||
|
||||
iterator = _make_header_iterator(
|
||||
headers={"apim-request-id": "azure-correlation-1"},
|
||||
config=_headers_config(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
async for _ in iterator:
|
||||
pass
|
||||
|
||||
assert len(logged) == 1
|
||||
iterator._completed_response_logged = False
|
||||
logged.clear()
|
||||
with patch.object(type(iterator.completed_response), "model_dump", side_effect=ValueError("cannot serialize")):
|
||||
iterator._log_completed_response(is_async=True)
|
||||
|
||||
assert logged == [iterator.completed_response]
|
||||
assert iterator.completed_response.response._hidden_params == {}
|
||||
|
|
|
|||
|
|
@ -172,8 +172,6 @@ class TestLLMHTTPHandlerRealtimeRedaction:
|
|||
|
||||
|
||||
class TestProxyStreamingDataGeneratorRedaction:
|
||||
"""Test _redact_string on traceback.format_exc() — the pattern at common_request_processing.py:1733."""
|
||||
|
||||
def test_redact_traceback_format_exc(self):
|
||||
try:
|
||||
raise RuntimeError(
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import logging
|
|||
import logging.config
|
||||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Callable
|
||||
from io import StringIO
|
||||
from typing import Final
|
||||
|
|
@ -13,11 +14,12 @@ from litellm._logging import (
|
|||
JsonFormatter,
|
||||
_redact_string,
|
||||
_secret_filter,
|
||||
redact_internal_details_from_client_message,
|
||||
verbose_logger,
|
||||
verbose_proxy_logger,
|
||||
verbose_router_logger,
|
||||
)
|
||||
from litellm.litellm_core_utils.secret_redaction import redact_string
|
||||
from litellm.litellm_core_utils.secret_redaction import redact_internal_details, redact_string
|
||||
|
||||
SECRET = "sk-proj-abc123def456ghi789jklmnopqrst"
|
||||
|
||||
|
|
@ -657,3 +659,65 @@ def test_json_formatter_redacts_non_string_extra_values(extra):
|
|||
assert output.strip(), "no record captured"
|
||||
assert SECRET not in output, f"non-string extra leaked a secret: {output}"
|
||||
assert "REDACTED" in output
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"text,leaked",
|
||||
(
|
||||
("config file /etc/litellm/secrets/db.yaml", "/etc/litellm/secrets/db.yaml"),
|
||||
("home dir /Users/admin/.litellm/master_key.txt", "/Users/admin/.litellm/master_key.txt"),
|
||||
("cache at /var/cache/litellm/tokens.db", "/var/cache/litellm/tokens.db"),
|
||||
("path C:\\Users\\admin\\secrets.env", "C:\\Users\\admin\\secrets.env"),
|
||||
("connecting to host 10.20.30.40", "10.20.30.40"),
|
||||
("connecting to host 192.168.1.5", "192.168.1.5"),
|
||||
("connecting to host 172.16.0.9", "172.16.0.9"),
|
||||
("connecting to host 127.0.0.1", "127.0.0.1"),
|
||||
("connecting to db-primary.internal", "db-primary.internal"),
|
||||
("connecting to redis.corp", "redis.corp"),
|
||||
),
|
||||
)
|
||||
def test_redact_internal_details_catches_paths_and_hostnames(text, leaked):
|
||||
result = redact_internal_details(text)
|
||||
assert leaked not in result, f"{leaked!r} was not redacted"
|
||||
assert "REDACTED" in result
|
||||
|
||||
|
||||
def test_redact_internal_details_leaves_public_hostnames_and_routes_alone():
|
||||
"""litellm's own error messages rely on routes like /v1/models staying legible."""
|
||||
safe_strings = (
|
||||
"call https://api.openai.com/v1/chat/completions",
|
||||
"/chat/completions: Invalid model name passed in model=gpt-9",
|
||||
"Call `/v1/models` to view available models for your key",
|
||||
"reducto:// file IDs are not accepted through the proxy OCR API",
|
||||
)
|
||||
for text in safe_strings:
|
||||
assert redact_internal_details(text) == text
|
||||
|
||||
|
||||
def test_redact_internal_details_layers_on_top_of_credential_redaction():
|
||||
text = "postgresql://litellm_internal:S3cr3tPGPass@10.20.30.40:5432/litellm_prod"
|
||||
result = redact_internal_details(text)
|
||||
assert "S3cr3tPGPass" not in result
|
||||
assert "10.20.30.40" not in result
|
||||
|
||||
|
||||
def test_redact_internal_details_drops_embedded_traceback():
|
||||
"""Regression for LIT-6747: the traceback exception_type() embeds for SDK callers
|
||||
must never reach an HTTP client."""
|
||||
try:
|
||||
raise RuntimeError("socket hung up")
|
||||
except RuntimeError:
|
||||
raw_tb = traceback.format_exc()
|
||||
message = f"litellm.APIConnectionError: MinimaxException - socket hung up\n{raw_tb}"
|
||||
|
||||
result = redact_internal_details(message)
|
||||
|
||||
assert result == "litellm.APIConnectionError: MinimaxException - socket hung up"
|
||||
assert "Traceback (most recent call last)" not in result
|
||||
assert __file__.split("/")[-1] not in result
|
||||
|
||||
|
||||
def test_redact_internal_details_from_client_message_respects_disable_flag():
|
||||
with patch("litellm._logging._ENABLE_SECRET_REDACTION", False): # test-quality-ok: the opt-out flag is the SUT
|
||||
text = "config file /etc/litellm/secrets/db.yaml"
|
||||
assert redact_internal_details_from_client_message(text) == text
|
||||
|
|
|
|||
|
|
@ -1,10 +1,19 @@
|
|||
import pytest
|
||||
import litellm
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure_ai.vector_stores.transformation import AzureAIVectorStoreConfig
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.vector_stores import (
|
||||
asearch as vector_store_asearch,
|
||||
)
|
||||
from litellm.vector_stores import (
|
||||
search as vector_store_search,
|
||||
asearch as vector_store_asearch,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -30,10 +39,108 @@ async def test_basic_search_vector_store(sync_mode):
|
|||
if sync_mode:
|
||||
response = vector_store_search(query=default_query, **base_request_args)
|
||||
else:
|
||||
response = await vector_store_asearch(
|
||||
query=default_query, **base_request_args
|
||||
)
|
||||
response = await vector_store_asearch(query=default_query, **base_request_args)
|
||||
except litellm.InternalServerError:
|
||||
pytest.skip("Skipping test due to litellm.InternalServerError")
|
||||
|
||||
print("litellm response=", json.dumps(response, indent=4, default=str))
|
||||
|
||||
|
||||
class RecordingEmbeddingExecutor:
|
||||
def __init__(self, response):
|
||||
self.response = response
|
||||
self.calls = []
|
||||
|
||||
def embed(self, model, query, configuration):
|
||||
self.calls.append((model, query, dict(configuration)))
|
||||
return self.response
|
||||
|
||||
async def aembed(self, model, query, configuration):
|
||||
self.calls.append((model, query, dict(configuration)))
|
||||
return self.response
|
||||
|
||||
|
||||
ALIAS_QUERY_VECTOR = [0.5, -0.25, 0.125]
|
||||
ALIAS_EMBEDDING_RESPONSE = EmbeddingResponse(
|
||||
data=[{"embedding": ALIAS_QUERY_VECTOR, "index": 0, "object": "embedding"}]
|
||||
)
|
||||
STORE_EMBEDDINGS_URL = "https://embedding.example/v1/embeddings"
|
||||
|
||||
|
||||
def _transform_kwargs(executor):
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
return {
|
||||
"vector_store_id": "my-vector-index",
|
||||
"query": "what is azure search?",
|
||||
"vector_store_search_optional_params": {"top_k": 2},
|
||||
"api_base": "https://azure-kb-search.search.windows.net",
|
||||
"litellm_logging_obj": logging_obj,
|
||||
"litellm_params": {
|
||||
"litellm_embedding_model": "multilingual-e5-large",
|
||||
"azure_search_vector_field": "embedding",
|
||||
},
|
||||
"embedding_executor": executor,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transform_uses_injected_executor_without_embedding_config(respx_mock: respx.MockRouter):
|
||||
executor = RecordingEmbeddingExecutor(ALIAS_EMBEDDING_RESPONSE)
|
||||
config = AzureAIVectorStoreConfig()
|
||||
transform_kwargs = _transform_kwargs(executor)
|
||||
|
||||
url, sync_body = config.transform_search_vector_store_request(**transform_kwargs)
|
||||
_, async_body = await config.atransform_search_vector_store_request(**transform_kwargs)
|
||||
|
||||
assert respx_mock.calls.call_count == 0
|
||||
assert executor.calls == [("multilingual-e5-large", "what is azure search?", {})] * 2
|
||||
assert (
|
||||
url == "https://azure-kb-search.search.windows.net/indexes/my-vector-index/docs/search?api-version=2024-07-01"
|
||||
)
|
||||
assert sync_body == async_body
|
||||
assert sync_body["vectorQueries"] == [
|
||||
{"vector": ALIAS_QUERY_VECTOR, "fields": "embedding", "kind": "vector", "k": 2}
|
||||
]
|
||||
assert sync_body["top"] == 2
|
||||
logging_details = transform_kwargs["litellm_logging_obj"].model_call_details
|
||||
assert logging_details["embedding_model"] == "multilingual-e5-large"
|
||||
assert logging_details["top_k"] == 2
|
||||
|
||||
|
||||
def test_transform_falls_back_to_sdk_embedding_without_executor(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = respx_mock.post(STORE_EMBEDDINGS_URL).mock(
|
||||
return_value=httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": ALIAS_QUERY_VECTOR}],
|
||||
"model": "text-embedding-3-small",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
)
|
||||
transform_kwargs = _transform_kwargs(None)
|
||||
transform_kwargs["litellm_params"] = {
|
||||
"litellm_embedding_model": "openai/text-embedding-3-small",
|
||||
"litellm_embedding_config": {"api_base": "https://embedding.example/v1", "api_key": "store-key"},
|
||||
}
|
||||
|
||||
_, body = AzureAIVectorStoreConfig().transform_search_vector_store_request(**transform_kwargs)
|
||||
|
||||
embedding_request = embedding_route.calls.last.request
|
||||
assert embedding_request.headers["authorization"] == "Bearer store-key"
|
||||
assert json.loads(embedding_request.read())["input"] == ["what is azure search?"]
|
||||
assert body["vectorQueries"][0]["vector"] == ALIAS_QUERY_VECTOR
|
||||
assert body["vectorQueries"][0]["fields"] == "contentVector"
|
||||
|
||||
|
||||
def test_transform_requires_embedding_model():
|
||||
transform_kwargs = _transform_kwargs(RecordingEmbeddingExecutor(ALIAS_EMBEDDING_RESPONSE))
|
||||
transform_kwargs["litellm_params"] = {"litellm_embedding_config": {"api_key": "store-key"}}
|
||||
|
||||
with pytest.raises(ValueError, match="litellm_embedding_model is required"):
|
||||
AzureAIVectorStoreConfig().transform_search_vector_store_request(**transform_kwargs)
|
||||
|
|
|
|||
|
|
@ -3,16 +3,19 @@ Tests for Milvus Vector Store
|
|||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.llms.milvus.vector_stores.transformation import MilvusVectorStoreConfig
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
from litellm.vector_stores import asearch as vector_store_asearch
|
||||
from litellm.vector_stores import search as vector_store_search
|
||||
|
||||
|
||||
# Mock response from actual Milvus API
|
||||
MOCK_MILVUS_SEARCH_RESPONSE = {
|
||||
"code": 0,
|
||||
|
|
@ -98,7 +101,7 @@ class TestMilvusVectorStore:
|
|||
mock_response.json.return_value = MOCK_MILVUS_SEARCH_RESPONSE
|
||||
mock_response.text = json.dumps(MOCK_MILVUS_SEARCH_RESPONSE)
|
||||
|
||||
with patch("litellm.embedding") as mock_embedding:
|
||||
with patch("litellm.aembedding", new_callable=AsyncMock) as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
|
|
@ -147,16 +150,10 @@ class TestMilvusVectorStore:
|
|||
else:
|
||||
# Fallback: check for json kwarg or in args
|
||||
request_data = call_args.kwargs.get("json")
|
||||
if (
|
||||
request_data is None
|
||||
and len(call_args.args) > 0
|
||||
and isinstance(call_args.args[0], dict)
|
||||
):
|
||||
if request_data is None and len(call_args.args) > 0 and isinstance(call_args.args[0], dict):
|
||||
request_data = call_args.args[0]
|
||||
|
||||
assert (
|
||||
request_data is not None
|
||||
), f"Could not extract request data. Call args: {call_args}"
|
||||
assert request_data is not None, f"Could not extract request data. Call args: {call_args}"
|
||||
print("Request data:", json.dumps(request_data, indent=2, default=str))
|
||||
|
||||
# Validate request structure
|
||||
|
|
@ -213,9 +210,7 @@ class TestMilvusVectorStore:
|
|||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
|
||||
) as mock_post:
|
||||
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
# Make the search request
|
||||
|
|
@ -252,16 +247,10 @@ class TestMilvusVectorStore:
|
|||
else:
|
||||
# Fallback: check for json kwarg or in args
|
||||
request_data = call_args.kwargs.get("json")
|
||||
if (
|
||||
request_data is None
|
||||
and len(call_args.args) > 0
|
||||
and isinstance(call_args.args[0], dict)
|
||||
):
|
||||
if request_data is None and len(call_args.args) > 0 and isinstance(call_args.args[0], dict):
|
||||
request_data = call_args.args[0]
|
||||
|
||||
assert (
|
||||
request_data is not None
|
||||
), f"Could not extract request data. Call args: {call_args}"
|
||||
assert request_data is not None, f"Could not extract request data. Call args: {call_args}"
|
||||
|
||||
# Validate request structure
|
||||
assert "collectionName" in request_data
|
||||
|
|
@ -316,11 +305,7 @@ class TestMilvusVectorStore:
|
|||
if request_data_str:
|
||||
return json.loads(request_data_str)
|
||||
request_data = call_args.kwargs.get("json")
|
||||
if (
|
||||
request_data is None
|
||||
and len(call_args.args) > 0
|
||||
and isinstance(call_args.args[0], dict)
|
||||
):
|
||||
if request_data is None and len(call_args.args) > 0 and isinstance(call_args.args[0], dict):
|
||||
request_data = call_args.args[0]
|
||||
return request_data
|
||||
|
||||
|
|
@ -334,9 +319,7 @@ class TestMilvusVectorStore:
|
|||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
|
||||
) as mock_post:
|
||||
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
vector_store_search(
|
||||
|
|
@ -375,9 +358,7 @@ class TestMilvusVectorStore:
|
|||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
|
||||
) as mock_post:
|
||||
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
vector_store_search(
|
||||
|
|
@ -413,9 +394,7 @@ class TestMilvusVectorStore:
|
|||
with patch("litellm.embedding") as mock_embedding:
|
||||
mock_embedding.return_value = MOCK_EMBEDDING_RESPONSE
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
|
||||
) as mock_post:
|
||||
with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
|
||||
vector_store_search(
|
||||
|
|
@ -492,3 +471,247 @@ if __name__ == "__main__":
|
|||
test.test_basic_search_with_mock_sync()
|
||||
|
||||
print("\n✅ All mock tests passed!")
|
||||
|
||||
|
||||
class RecordingEmbeddingExecutor:
|
||||
def __init__(self, response):
|
||||
self.response = response
|
||||
self.calls = []
|
||||
|
||||
def embed(self, model, query, configuration):
|
||||
self.calls.append((model, query, dict(configuration)))
|
||||
return self.response
|
||||
|
||||
async def aembed(self, model, query, configuration):
|
||||
self.calls.append((model, query, dict(configuration)))
|
||||
return self.response
|
||||
|
||||
|
||||
ALIAS_QUERY_VECTOR = [0.5, -0.25, 0.125]
|
||||
ALIAS_EMBEDDING_RESPONSE = EmbeddingResponse(
|
||||
data=[{"embedding": ALIAS_QUERY_VECTOR, "index": 0, "object": "embedding"}]
|
||||
)
|
||||
OPENAI_EMBEDDINGS_URL = "https://api.openai.com/v1/embeddings"
|
||||
MILVUS_SEARCH_URL = "https://milvus.example/v2/vectordb/entities/search"
|
||||
ALIAS_SEARCH_KWARGS = {
|
||||
"query": "what is machine learning?",
|
||||
"vector_store_id": "book_2",
|
||||
"custom_llm_provider": "milvus",
|
||||
"api_base": "https://milvus.example",
|
||||
"api_key": "mock_milvus_api_key",
|
||||
"litellm_embedding_model": "multilingual-e5-large",
|
||||
"milvus_text_field": "book_intro_text",
|
||||
}
|
||||
|
||||
|
||||
def _alias_router():
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "multilingual-e5-large",
|
||||
"litellm_params": {
|
||||
"model": "openai/text-embedding-3-small",
|
||||
"api_key": "deployment-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def _mock_embedding_route(respx_mock: respx.MockRouter) -> respx.Route:
|
||||
return respx_mock.post(OPENAI_EMBEDDINGS_URL).mock(
|
||||
return_value=httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"object": "list",
|
||||
"data": [{"object": "embedding", "index": 0, "embedding": ALIAS_QUERY_VECTOR}],
|
||||
"model": "text-embedding-3-small",
|
||||
"usage": {"prompt_tokens": 2, "total_tokens": 2},
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _mock_search_route(respx_mock: respx.MockRouter) -> respx.Route:
|
||||
return respx_mock.post(MILVUS_SEARCH_URL).mock(return_value=httpx.Response(200, json=MOCK_MILVUS_SEARCH_RESPONSE))
|
||||
|
||||
|
||||
def _assert_alias_resolved(embedding_route: respx.Route, search_route: respx.Route, response):
|
||||
embedding_request = embedding_route.calls.last.request
|
||||
assert embedding_request.headers["authorization"] == "Bearer deployment-key"
|
||||
embedding_body = json.loads(embedding_request.read())
|
||||
assert embedding_body["model"] == "text-embedding-3-small"
|
||||
assert embedding_body["input"] == ["what is machine learning?"]
|
||||
search_request = search_route.calls.last.request
|
||||
assert search_request.headers["authorization"] == "Bearer mock_milvus_api_key"
|
||||
assert json.loads(search_request.read())["data"] == [ALIAS_QUERY_VECTOR]
|
||||
assert len(response["data"]) == len(MOCK_MILVUS_SEARCH_RESPONSE["data"])
|
||||
assert response["data"][0]["content"][0]["text"] == MOCK_MILVUS_SEARCH_RESPONSE["data"][0]["book_intro_text"]
|
||||
|
||||
|
||||
def test_router_search_resolves_bare_embedding_alias_sync(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = _alias_router().vector_store_search(**ALIAS_SEARCH_KWARGS)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_search_resolves_bare_embedding_alias_async(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = await _alias_router().avector_store_search(**ALIAS_SEARCH_KWARGS)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
def test_sdk_search_with_router_kwarg_resolves_bare_embedding_alias_sync(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = litellm.vector_stores.search(router=_alias_router(), **ALIAS_SEARCH_KWARGS)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sdk_search_with_router_kwarg_resolves_bare_embedding_alias_async(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = await litellm.vector_stores.asearch(router=_alias_router(), **ALIAS_SEARCH_KWARGS)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
def _team_alias_router():
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "team-a-embedder",
|
||||
"litellm_params": {
|
||||
"model": "openai/text-embedding-3-small",
|
||||
"api_key": "deployment-key",
|
||||
},
|
||||
"model_info": {"team_id": "team-a", "team_public_model_name": "multilingual-e5-large"},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sdk_search_with_router_kwarg_resolves_team_alias_from_request_metadata(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
search_route = _mock_search_route(respx_mock)
|
||||
|
||||
response = await litellm.vector_stores.asearch(
|
||||
router=_team_alias_router(), metadata={"user_api_key_team_id": "team-a"}, **ALIAS_SEARCH_KWARGS
|
||||
)
|
||||
|
||||
_assert_alias_resolved(embedding_route, search_route, response)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sdk_search_with_router_kwarg_rejects_team_alias_without_team_metadata(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
_mock_search_route(respx_mock)
|
||||
|
||||
with pytest.raises(litellm.APIConnectionError):
|
||||
await litellm.vector_stores.asearch(router=_team_alias_router(), **ALIAS_SEARCH_KWARGS)
|
||||
|
||||
assert embedding_route.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transform_uses_injected_executor_without_embedding_config(respx_mock: respx.MockRouter):
|
||||
executor = RecordingEmbeddingExecutor(ALIAS_EMBEDDING_RESPONSE)
|
||||
config = MilvusVectorStoreConfig()
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
transform_kwargs = {
|
||||
"vector_store_id": "book_2",
|
||||
"query": ["what is", "milvus?"],
|
||||
"vector_store_search_optional_params": {"limit": 3},
|
||||
"api_base": "https://milvus.example",
|
||||
"litellm_logging_obj": logging_obj,
|
||||
"litellm_params": {"litellm_embedding_model": "multilingual-e5-large", "milvus_db_name": "docs"},
|
||||
"embedding_executor": executor,
|
||||
}
|
||||
|
||||
url, sync_body = config.transform_search_vector_store_request(**transform_kwargs)
|
||||
_, async_body = await config.atransform_search_vector_store_request(**transform_kwargs)
|
||||
|
||||
assert respx_mock.calls.call_count == 0
|
||||
assert executor.calls == [("multilingual-e5-large", "what is milvus?", {})] * 2
|
||||
assert url == MILVUS_SEARCH_URL
|
||||
assert sync_body == async_body
|
||||
assert sync_body == {
|
||||
"collectionName": "book_2",
|
||||
"data": [ALIAS_QUERY_VECTOR],
|
||||
"annsField": "book_intro_vector",
|
||||
"limit": 3,
|
||||
"dbName": "docs",
|
||||
}
|
||||
assert logging_obj.model_call_details["input"] == "what is milvus?"
|
||||
assert logging_obj.model_call_details["embedding_model"] == "multilingual-e5-large"
|
||||
|
||||
|
||||
def test_transform_falls_back_to_sdk_embedding_without_executor_or_config(
|
||||
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "env-key")
|
||||
embedding_route = _mock_embedding_route(respx_mock)
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.model_call_details = {}
|
||||
|
||||
_, body = MilvusVectorStoreConfig().transform_search_vector_store_request(
|
||||
vector_store_id="book_2",
|
||||
query="q",
|
||||
vector_store_search_optional_params={},
|
||||
api_base="https://milvus.example",
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params={"litellm_embedding_model": "openai/text-embedding-3-small"},
|
||||
)
|
||||
|
||||
embedding_request = embedding_route.calls.last.request
|
||||
assert embedding_request.headers["authorization"] == "Bearer env-key"
|
||||
assert json.loads(embedding_request.read())["input"] == ["q"]
|
||||
assert body["data"] == [ALIAS_QUERY_VECTOR]
|
||||
|
||||
|
||||
def test_transform_requires_embedding_model():
|
||||
with pytest.raises(ValueError, match="litellm_embedding_model is required"):
|
||||
MilvusVectorStoreConfig().transform_search_vector_store_request(
|
||||
vector_store_id="book_2",
|
||||
query="q",
|
||||
vector_store_search_optional_params={},
|
||||
api_base="https://milvus.example",
|
||||
litellm_logging_obj=MagicMock(),
|
||||
litellm_params={"litellm_embedding_config": {"api_key": "store-key"}},
|
||||
embedding_executor=RecordingEmbeddingExecutor(ALIAS_EMBEDDING_RESPONSE),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
"limit": 22330
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26762
|
||||
"limit": 26759
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 261
|
||||
|
|
@ -27,12 +27,12 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16474
|
||||
"limit": 16472
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5520
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4495
|
||||
"limit": 4489
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,16 @@
|
|||
import { AutoRouterPreset, hydratePresets } from "@/lib/autorouter_presets";
|
||||
import { getAutoRouterPresets } from "@/components/networking";
|
||||
import { useQuery } from "@tanstack/react-query";
|
||||
import { createQueryKeys } from "../common/queryKeysFactory";
|
||||
|
||||
const presetKeys = createQueryKeys("autoRouterPresets");
|
||||
|
||||
export const useAutoRouterPresets = () => {
|
||||
const options = {
|
||||
queryKey: presetKeys.list({}),
|
||||
queryFn: async () => hydratePresets(await getAutoRouterPresets()),
|
||||
staleTime: 24 * 60 * 60 * 1000,
|
||||
gcTime: 24 * 60 * 60 * 1000,
|
||||
};
|
||||
return useQuery<AutoRouterPreset[]>(options);
|
||||
};
|
||||
|
|
@ -9,11 +9,19 @@ import { getSubmitBlockedReason } from "./add_auto_router_tab";
|
|||
import { buildModelAvailability } from "@/lib/autorouter_presets";
|
||||
import { testAutoRouterRouting } from "../networking";
|
||||
import { ModelGroup } from "@/components/llm_calls/fetch_models";
|
||||
import { getAllPresets, getPresetByKey, getRequiredModelsInPreset } from "@/lib/autorouter_presets";
|
||||
import { AutoRouterPreset, getRequiredModelsInPreset } from "@/lib/autorouter_presets";
|
||||
import { BUNDLED_PRESETS, LOADED_PRESETS_QUERY, useAutoRouterPresets } from "../../../tests/mocks/autoRouterPresets";
|
||||
vi.mock(
|
||||
"@/app/(dashboard)/hooks/autoRouter/useComplexityScorerDefaults",
|
||||
async () => await import("../../../tests/mocks/complexityScorerDefaults"),
|
||||
);
|
||||
vi.mock(
|
||||
"@/app/(dashboard)/hooks/autoRouter/useAutoRouterPresets",
|
||||
async () => await import("../../../tests/mocks/autoRouterPresets"),
|
||||
);
|
||||
|
||||
const getAllPresets = (): AutoRouterPreset[] => BUNDLED_PRESETS;
|
||||
const getPresetByKey = (key: string): AutoRouterPreset | undefined => BUNDLED_PRESETS.find((p) => p.key === key);
|
||||
|
||||
const ANTHROPIC_PRESET = getPresetByKey("anthropic_family")!;
|
||||
const ANTHROPIC_TIERS = ANTHROPIC_PRESET.complexity_router_config.tiers;
|
||||
|
|
@ -1142,3 +1150,52 @@ describe("getSubmitBlockedReason", () => {
|
|||
);
|
||||
});
|
||||
});
|
||||
|
||||
describe("preset catalog fetch states", () => {
|
||||
afterEach(() => vi.mocked(useAutoRouterPresets).mockReturnValue(LOADED_PRESETS_QUERY));
|
||||
|
||||
it("keeps showing cached presets without the error banner when only a refetch fails", () => {
|
||||
vi.mocked(useAutoRouterPresets).mockReturnValue({
|
||||
...LOADED_PRESETS_QUERY,
|
||||
isError: true,
|
||||
} as never);
|
||||
renderWithProviders(<Harness />);
|
||||
|
||||
expect(screen.queryByText(/Could not load templates/)).not.toBeInTheDocument();
|
||||
|
||||
openTemplateDropdown();
|
||||
expect(screen.queryAllByRole("option").length).toBeGreaterThan(1);
|
||||
});
|
||||
|
||||
it("shows a loading hint while the catalog fetch is pending", () => {
|
||||
vi.mocked(useAutoRouterPresets).mockReturnValue({
|
||||
...LOADED_PRESETS_QUERY,
|
||||
data: undefined,
|
||||
isPending: true,
|
||||
} as never);
|
||||
renderWithProviders(<Harness />);
|
||||
|
||||
expect(screen.getByText("Loading templates...")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("degrades to Custom Configuration with a retry hint that refetches the catalog", async () => {
|
||||
const refetch = vi.fn();
|
||||
vi.mocked(useAutoRouterPresets).mockReturnValue({
|
||||
...LOADED_PRESETS_QUERY,
|
||||
data: undefined,
|
||||
isError: true,
|
||||
refetch,
|
||||
} as never);
|
||||
renderWithProviders(<Harness />);
|
||||
|
||||
expect(await screen.findByText(/Could not load templates/)).toBeInTheDocument();
|
||||
|
||||
openTemplateDropdown();
|
||||
const options = screen.queryAllByRole("option");
|
||||
expect(options).toHaveLength(1);
|
||||
expect(options[0]).toHaveTextContent("Custom Configuration");
|
||||
|
||||
fireEvent.click(screen.getByRole("button", { name: "Retry" }));
|
||||
expect(refetch).toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -50,8 +50,6 @@ import AutoRouterConnectionTest from "./auto_router_connection_test";
|
|||
import AutoRouterRoutingTest from "./AutoRouterRoutingTest";
|
||||
import { toast } from "@/lib/toast";
|
||||
import {
|
||||
getAllPresets,
|
||||
getPresetByKey,
|
||||
getMissingModelsInPreset,
|
||||
getReferencedModelsError,
|
||||
buildEmptyPrefill,
|
||||
|
|
@ -62,6 +60,7 @@ import {
|
|||
PresetPrefill,
|
||||
AutoRouterPreset,
|
||||
} from "@/lib/autorouter_presets";
|
||||
import { useAutoRouterPresets } from "@/app/(dashboard)/hooks/autoRouter/useAutoRouterPresets";
|
||||
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
|
||||
interface AddAutoRouterTabProps {
|
||||
|
|
@ -102,9 +101,7 @@ const presetDisabledHint = (availability: PresetAvailability): string | null =>
|
|||
// caller-specific missing-model reason gets the alarming red treatment.
|
||||
const isPresetHintAlarming = (availability: PresetAvailability): boolean => availability.kind === "missing_models";
|
||||
|
||||
// getAllPresets() already returns a stable, module-level array (see autorouter_presets.ts), so
|
||||
// this is resolved once at import time rather than re-called from inside the component every render.
|
||||
const presets = getAllPresets();
|
||||
const NO_PRESETS: AutoRouterPreset[] = [];
|
||||
|
||||
// A one-line summary of what's configured, shown when the detailed section is collapsed so a
|
||||
// caller can see the shape of the config without opening it.
|
||||
|
|
@ -229,6 +226,14 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
|
|||
});
|
||||
const modelsLoading = groupsLoading || deploymentsLoading;
|
||||
const modelInfo = React.useMemo(() => data ?? [], [data]);
|
||||
const {
|
||||
data: presetsData,
|
||||
isPending: presetsPending,
|
||||
isError: presetsError,
|
||||
refetch: refetchPresets,
|
||||
} = useAutoRouterPresets();
|
||||
const presets = presetsData ?? NO_PRESETS;
|
||||
const presetsUnavailable = presetsError && presetsData === undefined;
|
||||
// react-query keeps the last successful list around when a later refetch fails, so isError alone
|
||||
// can't tell "never loaded" apart from "loaded, then a background refetch errored" - only the
|
||||
// former leaves us with nothing trustworthy to verify a preset's models against.
|
||||
|
|
@ -277,7 +282,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
|
|||
presets
|
||||
.map((preset) => ({ preset, availability: presetAvailability(preset) }))
|
||||
.sort((a, b) => Number(b.availability.kind === "available") - Number(a.availability.kind === "available")),
|
||||
[presetAvailability],
|
||||
[presets, presetAvailability],
|
||||
);
|
||||
|
||||
const templateItems = React.useMemo(
|
||||
|
|
@ -307,7 +312,7 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
|
|||
return;
|
||||
}
|
||||
|
||||
const preset = getPresetByKey(presetKey);
|
||||
const preset = presets.find((p) => p.key === presetKey);
|
||||
// Refuse to apply a preset whose models are not verified available. The dropdown disables
|
||||
// these options, so this is a guard against a stale click resolving after the list changed.
|
||||
if (!preset) return;
|
||||
|
|
@ -538,6 +543,15 @@ const AddAutoRouterTab: React.FC<AddAutoRouterTabProps> = ({
|
|||
</button>
|
||||
</div>
|
||||
)}
|
||||
{presetsPending && <div className="text-xs mt-1 text-muted-foreground">Loading templates...</div>}
|
||||
{presetsUnavailable && (
|
||||
<div className="text-xs mt-1 text-destructive">
|
||||
Could not load templates, so only Custom Configuration is shown.{" "}
|
||||
<button type="button" className="underline" onClick={() => void refetchPresets()}>
|
||||
Retry
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{requiresTeamScope && (
|
||||
|
|
|
|||
|
|
@ -90,6 +90,7 @@ import type {
|
|||
} from "@/app/(dashboard)/caching/_components/coordination_redis_settings/types";
|
||||
import { MCP_TOOLS_PREVIEW_FORBIDDEN_MESSAGE } from "./mcp_tools/constants";
|
||||
import type { ComplexityRouterConfigPayload } from "./add_model/build_complexity_router_config";
|
||||
import type { AutoRouterPresetsResponse } from "@/lib/autorouter_presets";
|
||||
import type { VectorStoreIndex } from "@/app/(dashboard)/vector-stores/_components/IndexesTab";
|
||||
import type { RoutingDecision } from "./view_logs/LogDetailsDrawer/RoutingDecisionCard";
|
||||
import {
|
||||
|
|
@ -410,6 +411,15 @@ export const getComplexityScorerDefaults = async (): Promise<ComplexityScorerDef
|
|||
return await apiClient.get(`/public/complexity_router/scorer_defaults`);
|
||||
};
|
||||
|
||||
export const getAutoRouterPresets = async (): Promise<AutoRouterPresetsResponse> => {
|
||||
/**
|
||||
* Fetch the auto-router preset catalog from the proxy's public endpoint. The template picker
|
||||
* renders from this rather than from a copy in the dashboard, so a catalog update propagates
|
||||
* without a dashboard release.
|
||||
*/
|
||||
return await apiClient.get(`/public/autorouter_presets`);
|
||||
};
|
||||
|
||||
export const getAgentCreateMetadata = async (): Promise<AgentCreateInfo[]> => {
|
||||
/**
|
||||
* Fetch agent type metadata from the proxy's public endpoint.
|
||||
|
|
@ -2042,6 +2052,7 @@ interface UiSpendLogsParams {
|
|||
min_spend?: number;
|
||||
max_spend?: number;
|
||||
exclude_internal_health_checks?: boolean;
|
||||
group_by_session?: boolean;
|
||||
}
|
||||
|
||||
interface UiSpendLogsCallOptions {
|
||||
|
|
@ -6876,13 +6887,14 @@ export const buildMcpOAuthAuthorizeUrl = ({
|
|||
const base = getProxyBaseUrl();
|
||||
const normalizedServerId = encodeURIComponent(serverId.trim());
|
||||
const url = `${base}/v1/mcp/server/oauth/${normalizedServerId}/authorize`;
|
||||
const params = new URLSearchParams({
|
||||
const authorizeParams = {
|
||||
redirect_uri: redirectUri,
|
||||
state,
|
||||
response_type: "code",
|
||||
code_challenge: codeChallenge,
|
||||
code_challenge_method: "S256",
|
||||
});
|
||||
};
|
||||
const params = new URLSearchParams(authorizeParams);
|
||||
if (clientId && clientId.trim().length > 0) {
|
||||
params.set("client_id", clientId);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -138,33 +138,39 @@ describe("RequestLogsPanel", () => {
|
|||
respondWith([]);
|
||||
});
|
||||
|
||||
describe("multi-call session collapsing", () => {
|
||||
const sessionRows = [
|
||||
logEntry({ request_id: "req-mcp", call_type: "call_mcp_tool", session_id: "sess-1", session_total_count: 3 }),
|
||||
logEntry({ request_id: "req-llm", call_type: "acompletion", session_id: "sess-1", session_total_count: 3 }),
|
||||
logEntry({ request_id: "req-llm-2", call_type: "acompletion", session_id: "sess-1", session_total_count: 3 }),
|
||||
];
|
||||
|
||||
it("collapses a multi-call session to a single representative row", async () => {
|
||||
respondWith(sessionRows);
|
||||
describe("server-grouped session pagination (#38060)", () => {
|
||||
it("requests session-grouped pages of 10 rows by default", async () => {
|
||||
renderPanel();
|
||||
|
||||
await waitFor(() => expect(row("req-mcp") ?? row("req-llm") ?? row("req-llm-2")).not.toBeNull());
|
||||
|
||||
const rendered = ["req-mcp", "req-llm", "req-llm-2"].filter((id) => row(id) !== null);
|
||||
expect(rendered).toHaveLength(1);
|
||||
await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled());
|
||||
expect(lastCall()?.params?.group_by_session).toBe(true);
|
||||
expect(lastCall()?.page_size).toBe(10);
|
||||
});
|
||||
|
||||
it("prefers an LLM call over an MCP call as the session's representative", async () => {
|
||||
respondWith(sessionRows);
|
||||
it("renders every row the server returns without client-side collapsing", async () => {
|
||||
respondWith([
|
||||
logEntry({ request_id: "req-a", call_type: "acompletion", session_id: "sess-1", session_total_count: 3 }),
|
||||
logEntry({ request_id: "req-b", call_type: "acompletion", session_id: "sess-1", session_total_count: 3 }),
|
||||
logEntry({ request_id: "req-c", call_type: "acompletion", session_id: "sess-1", session_total_count: 3 }),
|
||||
]);
|
||||
renderPanel();
|
||||
|
||||
await waitFor(() => expect(row("req-llm")).not.toBeNull());
|
||||
expect(row("req-mcp")).toBeNull();
|
||||
await waitFor(() => expect(row("req-a")).not.toBeNull());
|
||||
expect(row("req-b")).not.toBeNull();
|
||||
expect(row("req-c")).not.toBeNull();
|
||||
});
|
||||
|
||||
it("shows the session's call count and composition on the representative row", async () => {
|
||||
respondWith(sessionRows);
|
||||
it("shows the session's call count on the server-picked representative row", async () => {
|
||||
respondWith([
|
||||
logEntry({
|
||||
request_id: "req-llm",
|
||||
call_type: "acompletion",
|
||||
session_id: "sess-1",
|
||||
session_total_count: 3,
|
||||
session_llm_count: 2,
|
||||
mcp_tool_call_count: 1,
|
||||
}),
|
||||
]);
|
||||
renderPanel();
|
||||
|
||||
await waitFor(() => expect(row("req-llm")).not.toBeNull());
|
||||
|
|
@ -296,6 +302,7 @@ describe("RequestLogsPanel", () => {
|
|||
if (!byIdCall) throw new Error("expected a by-id uiSpendLogsCall");
|
||||
expect(byIdCall.page).toBe(1);
|
||||
expect(byIdCall.page_size).toBe(1);
|
||||
expect(byIdCall.params?.group_by_session).toBeUndefined();
|
||||
});
|
||||
|
||||
it("closing the drawer removes ?log_id= from the URL and closes the drawer", async () => {
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import type { KeyResponse } from "../key_team_helpers/key_list";
|
|||
import { keyInfoV1Call, uiSpendLogsCall } from "../networking";
|
||||
import KeyInfoView from "../templates/key_info_view";
|
||||
import type { LogEntry } from "./columns";
|
||||
import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "./constants";
|
||||
import { LOGS_PAGE_SIZE_OPTIONS } from "./constants";
|
||||
import {
|
||||
DEFAULT_LOGS_SORTING,
|
||||
formatLogsWindow,
|
||||
|
|
@ -24,7 +24,7 @@ import { LogDetailsDrawer } from "./LogDetailsDrawer";
|
|||
import { LiveTailBanner, LogsTableToolbar } from "./LogsTableToolbar";
|
||||
import { RequestLogsTable } from "./RequestLogsTable";
|
||||
|
||||
const PAGE_SIZE = 50;
|
||||
const PAGE_SIZE = LOGS_PAGE_SIZE_OPTIONS[0];
|
||||
const DEFAULT_INTERVAL = { value: 24, unit: "hours" };
|
||||
|
||||
interface RequestLogsPanelProps {
|
||||
|
|
@ -35,12 +35,6 @@ interface RequestLogsPanelProps {
|
|||
isActive: boolean;
|
||||
}
|
||||
|
||||
interface SessionComposition {
|
||||
llm: number;
|
||||
agent: number;
|
||||
mcp: number;
|
||||
}
|
||||
|
||||
export default function RequestLogsPanel({ accessToken, token, userRole, userID, isActive }: RequestLogsPanelProps) {
|
||||
const [pagination, setPagination] = useState<PaginationState>({ pageIndex: 0, pageSize: PAGE_SIZE });
|
||||
const [sorting, setSorting] = useState<SortingState>(DEFAULT_LOGS_SORTING);
|
||||
|
|
@ -157,49 +151,7 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID,
|
|||
|
||||
const isDrawerOpen = displayLog !== null || displaySessionId !== null;
|
||||
|
||||
const rows = useMemo<LogEntry[]>(() => {
|
||||
const searchedLogs = filteredLogs.data;
|
||||
|
||||
const sessionCompositionById = searchedLogs.reduce<Record<string, SessionComposition>>((acc, log) => {
|
||||
if (!log.session_id) return acc;
|
||||
if (!acc[log.session_id]) {
|
||||
acc[log.session_id] = { llm: 0, agent: 0, mcp: 0 };
|
||||
}
|
||||
if (MCP_CALL_TYPES.includes(log.call_type)) {
|
||||
acc[log.session_id].mcp += 1;
|
||||
} else if (AGENT_CALL_TYPES.includes(log.call_type)) {
|
||||
acc[log.session_id].agent += 1;
|
||||
} else {
|
||||
acc[log.session_id].llm += 1;
|
||||
}
|
||||
return acc;
|
||||
}, {});
|
||||
|
||||
const sessionRepresentativeMap = new Map<string, { requestId: string; isMcp: boolean }>();
|
||||
for (const log of searchedLogs) {
|
||||
if (!log.session_id || (log.session_total_count || 1) <= 1) continue;
|
||||
const isMcp = MCP_CALL_TYPES.includes(log.call_type);
|
||||
const existing = sessionRepresentativeMap.get(log.session_id);
|
||||
if (!existing || (existing.isMcp && !isMcp)) {
|
||||
sessionRepresentativeMap.set(log.session_id, { requestId: log.request_id, isMcp });
|
||||
}
|
||||
}
|
||||
|
||||
return searchedLogs
|
||||
.map((log) => {
|
||||
const sessionComposition = log.session_id ? sessionCompositionById[log.session_id] : undefined;
|
||||
return {
|
||||
...log,
|
||||
session_llm_count: sessionComposition?.llm ?? undefined,
|
||||
session_mcp_count: sessionComposition?.mcp ?? undefined,
|
||||
session_agent_count: sessionComposition?.agent ?? undefined,
|
||||
};
|
||||
})
|
||||
.filter((log) => {
|
||||
if (!log.session_id || (log.session_total_count || 1) <= 1) return true;
|
||||
return sessionRepresentativeMap.get(log.session_id)?.requestId === log.request_id;
|
||||
});
|
||||
}, [filteredLogs.data]);
|
||||
const rows: LogEntry[] = filteredLogs.data;
|
||||
|
||||
const searchTerm = useMemo(() => {
|
||||
const entry = columnFilters.find((filter) => filter.id === LOG_FILTER_IDS.REQUEST_ID);
|
||||
|
|
@ -258,13 +210,12 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID,
|
|||
);
|
||||
|
||||
const handleSessionClick = useCallback(
|
||||
(sessionId: string) => {
|
||||
if (!sessionId) return;
|
||||
const log = rows.find((candidate) => candidate.session_id === sessionId) ?? null;
|
||||
(log: LogEntry) => {
|
||||
if (!log.session_id) return;
|
||||
setSelectedLog(log);
|
||||
openSession(sessionId, log?.request_id ?? null);
|
||||
openSession(log.session_id, log.request_id);
|
||||
},
|
||||
[rows, openSession],
|
||||
[openSession],
|
||||
);
|
||||
|
||||
const handleSelectLog = useCallback(
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import { DataTable, DataTableFilterDrawer, DataTableToolbar } from "@/components
|
|||
|
||||
import type { Team } from "../key_team_helpers/key_list";
|
||||
import type { LogEntry } from "./columns";
|
||||
import { LOGS_PAGE_SIZE_OPTIONS } from "./constants";
|
||||
import { LOG_FILTER_LABELS, type LogsWindow } from "./log_filter_logic";
|
||||
import { RequestLogsFilters } from "./RequestLogsFilters";
|
||||
import { getRequestLogsTableColumns } from "./RequestLogsTableColumns";
|
||||
|
|
@ -28,7 +29,7 @@ interface RequestLogsTableProps {
|
|||
onRefresh: () => void;
|
||||
onRowClick: (log: LogEntry) => void;
|
||||
onKeyHashClick: (keyHash: string) => void;
|
||||
onSessionClick: (sessionId: string) => void;
|
||||
onSessionClick: (log: LogEntry) => void;
|
||||
teams: Team[];
|
||||
logsWindow: LogsWindow;
|
||||
toolbarChildren?: ReactNode;
|
||||
|
|
@ -91,6 +92,7 @@ export function RequestLogsTable({
|
|||
paginationMode="server"
|
||||
pagination={pagination}
|
||||
onPaginationChange={onPaginationChange}
|
||||
pageSizeOptions={LOGS_PAGE_SIZE_OPTIONS}
|
||||
rowCount={rowCount}
|
||||
filterMode="server"
|
||||
columnFilters={columnFilters}
|
||||
|
|
|
|||
|
|
@ -85,13 +85,19 @@ describe("row action cells", () => {
|
|||
expect(deps.onKeyHashClick).toHaveBeenCalledWith("sk-hash-9");
|
||||
});
|
||||
|
||||
it("reports the session id from the session cell", async () => {
|
||||
it("reports the clicked row from the session cell, so two rows sharing a session id stay distinguishable", async () => {
|
||||
const user = userEvent.setup();
|
||||
const deps = { onKeyHashClick: vi.fn(), onSessionClick: vi.fn() };
|
||||
renderRows([logEntry({ request_id: "req-sess", session_id: "sess-42" })], deps);
|
||||
renderRows(
|
||||
[
|
||||
logEntry({ request_id: "req-key-a", session_id: "sess-42", api_key: "key-a" }),
|
||||
logEntry({ request_id: "req-key-b", session_id: "sess-42", api_key: "key-b" }),
|
||||
],
|
||||
deps,
|
||||
);
|
||||
|
||||
await user.click(screen.getByText("sess-42"));
|
||||
expect(deps.onSessionClick).toHaveBeenCalledWith("sess-42");
|
||||
await user.click(screen.getAllByText("sess-42")[1]);
|
||||
expect(deps.onSessionClick).toHaveBeenCalledWith(expect.objectContaining({ request_id: "req-key-b" }));
|
||||
});
|
||||
});
|
||||
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ import { AgentBadge, AgentIcon, LlmBadge, McpBadge, SparkleIcon, WrenchIcon } fr
|
|||
|
||||
export interface RequestLogsTableColumnsDeps {
|
||||
onKeyHashClick: (keyHash: string) => void;
|
||||
onSessionClick: (sessionId: string) => void;
|
||||
onSessionClick: (log: LogEntry) => void;
|
||||
}
|
||||
|
||||
const readMetaString = (metadata: Record<string, unknown> | undefined, key: string): string | undefined => {
|
||||
|
|
@ -61,7 +61,7 @@ export const getRequestLogsTableColumns = ({
|
|||
const isAgent = AGENT_CALL_TYPES.includes(log.call_type);
|
||||
const sessionLlmCount = log.session_llm_count ?? (isMcp || isAgent ? 0 : sessionCount);
|
||||
const sessionAgentCount = log.session_agent_count ?? (isAgent ? sessionCount : 0);
|
||||
const sessionMcpCount = log.session_mcp_count ?? (isMcp ? sessionCount : 0);
|
||||
const sessionMcpCount = log.mcp_tool_call_count ?? (isMcp ? sessionCount : 0);
|
||||
|
||||
if (isMcp) return <McpBadge />;
|
||||
if (isAgent && sessionCount <= 1) return <AgentBadge />;
|
||||
|
|
@ -113,7 +113,7 @@ export const getRequestLogsTableColumns = ({
|
|||
header: "Session ID",
|
||||
size: 120,
|
||||
enableSorting: false,
|
||||
cell: ({ row }) => <IdCell value={row.original.session_id} onClick={onSessionClick} />,
|
||||
cell: ({ row }) => <IdCell value={row.original.session_id} onClick={() => onSessionClick(row.original)} />,
|
||||
},
|
||||
{
|
||||
id: "request_id",
|
||||
|
|
|
|||
|
|
@ -46,6 +46,5 @@ export type LogEntry = {
|
|||
mcp_tool_call_count?: number;
|
||||
mcp_tool_call_spend?: number;
|
||||
session_llm_count?: number;
|
||||
session_mcp_count?: number;
|
||||
session_agent_count?: number;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -12,6 +12,9 @@ export const ERROR_CODE_OPTIONS: { label: string; value: string }[] = [
|
|||
{ label: "529 - Overloaded", value: "529" },
|
||||
];
|
||||
|
||||
/** Page sizes the logs tables offer; the first entry is the default. */
|
||||
export const LOGS_PAGE_SIZE_OPTIONS = [10, 25, 50, 100];
|
||||
|
||||
/** Call types that represent MCP tool invocations (shared across columns, index, drawer). */
|
||||
export const MCP_CALL_TYPES = ["call_mcp_tool", "list_mcp_tools"];
|
||||
|
||||
|
|
|
|||
|
|
@ -181,6 +181,7 @@ export function useLogFilterLogic({
|
|||
sort_by: sortBy,
|
||||
sort_order: sortOrder,
|
||||
exclude_internal_health_checks: excludeInternalHealthChecks,
|
||||
group_by_session: true,
|
||||
},
|
||||
});
|
||||
},
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
import { describe, it, expect } from "vitest";
|
||||
import bundledPresets from "../../../../litellm/proxy/public_endpoints/autorouter_presets.json";
|
||||
import {
|
||||
getAllPresets,
|
||||
getPresetByKey,
|
||||
hydratePresets,
|
||||
AutoRouterPreset,
|
||||
AutoRouterPresetsResponse,
|
||||
getRequiredModelsInPreset,
|
||||
getMissingModelsInPreset,
|
||||
getRequiredModels,
|
||||
|
|
@ -18,8 +20,13 @@ import { DEFAULT_ESCALATION_KEYWORDS } from "@/components/add_model/EscalationKe
|
|||
|
||||
const groupsOnly = (models: Iterable<string>) => buildModelAvailability(models, []);
|
||||
|
||||
// Hydrated from the real bundled catalog so a catalog edit flows into these expectations.
|
||||
const PRESETS = hydratePresets(bundledPresets as AutoRouterPresetsResponse);
|
||||
const getAllPresets = (): AutoRouterPreset[] => PRESETS;
|
||||
const getPresetByKey = (key: string): AutoRouterPreset | undefined => PRESETS.find((p) => p.key === key);
|
||||
|
||||
describe("autorouter_presets", () => {
|
||||
it("loads exactly the bundled presets", () => {
|
||||
it("hydrates exactly the bundled presets", () => {
|
||||
const presets = getAllPresets();
|
||||
expect(presets.map((p) => p.label).sort()).toEqual(["Anthropic Family", "Gemini Family", "Lite", "OpenAI Family"]);
|
||||
// Every preset carries all four fields the UI relies on; a JSON typo dropping one fails here.
|
||||
|
|
|
|||
|
|
@ -20,7 +20,6 @@ import {
|
|||
} from "@/components/add_model/complexity_router_tiers";
|
||||
import { DEFAULT_ESCALATION_KEYWORDS } from "@/components/add_model/EscalationKeywords";
|
||||
import { DEFAULT_MATCH_THRESHOLD } from "@/components/add_model/SemanticKeywordMatching";
|
||||
import presetsRaw from "@/autorouter_presets.json";
|
||||
|
||||
// `key` is the stable JSON object key (e.g. "anthropic_family"); `label` is display text and
|
||||
// never an identity.
|
||||
|
|
@ -31,16 +30,10 @@ export interface AutoRouterPreset {
|
|||
complexity_router_config: ComplexityRouterConfigPayload;
|
||||
}
|
||||
|
||||
// The bundled JSON is a developer-authored, build-time asset, so it is trusted at the import
|
||||
// boundary rather than re-validated at runtime (resolveJsonModule widens its string literals,
|
||||
// hence this one cast). autorouter_presets.test.ts pins the parsed shape, so a JSON typo fails CI.
|
||||
const RAW = presetsRaw as Record<string, Omit<AutoRouterPreset, "key">>;
|
||||
export type AutoRouterPresetsResponse = Record<string, Omit<AutoRouterPreset, "key">>;
|
||||
|
||||
const PRESETS: AutoRouterPreset[] = Object.entries(RAW).map(([key, preset]) => ({ key, ...preset }));
|
||||
|
||||
export const getAllPresets = (): AutoRouterPreset[] => PRESETS;
|
||||
|
||||
export const getPresetByKey = (key: string): AutoRouterPreset | undefined => PRESETS.find((p) => p.key === key);
|
||||
export const hydratePresets = (raw: AutoRouterPresetsResponse): AutoRouterPreset[] =>
|
||||
Object.entries(raw).map(([key, preset]) => ({ key, ...preset }));
|
||||
|
||||
// Generalized over ComplexityRouterConfigPayload so the same accessors check either a preset's own
|
||||
// bundled config or a caller's actually-built config - the two need to agree, since a preset only
|
||||
|
|
|
|||
94
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
94
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -12129,6 +12129,31 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/public/autorouter_presets": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* Get Public Autorouter Presets
|
||||
* @description Return the auto-router preset catalog the dashboard's template picker renders.
|
||||
*
|
||||
* Resolved once per process, like the model cost map: fetched from ``litellm.autorouter_presets_url``
|
||||
* (override with ``LITELLM_AUTOROUTER_PRESETS_URL``) on the first request, falling back to the
|
||||
* catalog bundled with the package on any failure. Set ``LITELLM_LOCAL_AUTOROUTER_PRESETS=True``
|
||||
* to serve the bundled catalog only. A restart picks up a newly published catalog.
|
||||
*/
|
||||
get: operations["get_public_autorouter_presets_public_autorouter_presets_get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/public/complexity_router/scorer_defaults": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -23393,6 +23418,49 @@ export interface components {
|
|||
/** Tier Definitions */
|
||||
tier_definitions: components["schemas"]["TierDefinition"][];
|
||||
};
|
||||
/**
|
||||
* AutoRouterPresetConfig
|
||||
* @description The complexity_router_config a preset prefills.
|
||||
*
|
||||
* Only tiers is validated, because every dashboard consumer dereferences it; everything else
|
||||
* passes through verbatim with unknown fields kept (extra="allow"), so a catalog published after
|
||||
* this proxy shipped still serves its new fields intact.
|
||||
*/
|
||||
AutoRouterPresetConfig: {
|
||||
tiers: components["schemas"]["AutoRouterPresetTiers"];
|
||||
} & {
|
||||
[key: string]: unknown;
|
||||
};
|
||||
/**
|
||||
* AutoRouterPresetRecord
|
||||
* @description One auto-router preset as served to the dashboard's template picker.
|
||||
*/
|
||||
AutoRouterPresetRecord: {
|
||||
complexity_router_config: components["schemas"]["AutoRouterPresetConfig"];
|
||||
/** Description */
|
||||
description: string;
|
||||
/** Label */
|
||||
label: string;
|
||||
} & {
|
||||
[key: string]: unknown;
|
||||
};
|
||||
/**
|
||||
* AutoRouterPresetTiers
|
||||
* @description Exactly the four built-in tiers the dashboard's preset prefill can apply.
|
||||
*
|
||||
* extra="forbid" on purpose: a tier name this dashboard cannot apply would grey out or crash the
|
||||
* picker, so such a catalog is rejected wholesale and the bundled one serves instead.
|
||||
*/
|
||||
AutoRouterPresetTiers: {
|
||||
/** Complex */
|
||||
COMPLEX: string[];
|
||||
/** Medium */
|
||||
MEDIUM: string[];
|
||||
/** Reasoning */
|
||||
REASONING: string[];
|
||||
/** Simple */
|
||||
SIMPLE: string[];
|
||||
};
|
||||
/**
|
||||
* AutoRouterRoutingTestRequest
|
||||
* @description A single request to classify against a complexity-router config that need not be saved yet.
|
||||
|
|
@ -54681,6 +54749,28 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
get_public_autorouter_presets_public_autorouter_presets_get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": {
|
||||
[key: string]: components["schemas"]["AutoRouterPresetRecord"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
get_complexity_scorer_defaults_public_complexity_router_scorer_defaults_get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -56757,6 +56847,8 @@ export interface operations {
|
|||
sort_order?: string | null;
|
||||
/** @description Exclude LiteLLM internal health check requests from results */
|
||||
exclude_internal_health_checks?: boolean;
|
||||
/** @description Paginate over sessions instead of raw logs: one representative row per session, total counts sessions */
|
||||
group_by_session?: boolean;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
|
|
@ -56869,6 +56961,8 @@ export interface operations {
|
|||
sort_order?: string | null;
|
||||
/** @description Exclude LiteLLM internal health check requests from results */
|
||||
exclude_internal_health_checks?: boolean;
|
||||
/** @description Paginate over sessions instead of raw logs: one representative row per session, total counts sessions */
|
||||
group_by_session?: boolean;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
|
|
|
|||
16
ui/litellm-dashboard/tests/mocks/autoRouterPresets.ts
Normal file
16
ui/litellm-dashboard/tests/mocks/autoRouterPresets.ts
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
import { vi } from "vitest";
|
||||
import bundledPresets from "../../../../litellm/proxy/public_endpoints/autorouter_presets.json";
|
||||
import { hydratePresets, type AutoRouterPresetsResponse } from "@/lib/autorouter_presets";
|
||||
|
||||
// Derived from the real bundled catalog so a preset edit there flows into test expectations
|
||||
// instead of redding on a stale copy. Exported as vi.fn so a test can override the query state.
|
||||
export const BUNDLED_PRESETS = hydratePresets(bundledPresets as AutoRouterPresetsResponse);
|
||||
|
||||
export const LOADED_PRESETS_QUERY = {
|
||||
data: BUNDLED_PRESETS,
|
||||
isPending: false,
|
||||
isError: false,
|
||||
refetch: vi.fn(),
|
||||
};
|
||||
|
||||
export const useAutoRouterPresets = vi.fn(() => LOADED_PRESETS_QUERY);
|
||||
Loading…
Add table
Reference in a new issue