mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_decrease_anys_fable6
# Conflicts: # basedpyright-code-budget.json # litellm/integrations/websearch_interception/handler.py # litellm/proxy/response_polling/background_streaming.py # ruff-strict-budget.json # type-discipline-budget.json
This commit is contained in:
commit
fbf7644676
51 changed files with 1091 additions and 580 deletions
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 16984
|
||||
"limit": 16171
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2535
|
||||
"limit": 2226
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"limit": 319
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 5442
|
||||
"limit": 5199
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5655
|
||||
"limit": 5611
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15419
|
||||
"limit": 15350
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -72,7 +72,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportOptionalMemberAccess": {
|
||||
"limit": 1055
|
||||
"limit": 0
|
||||
},
|
||||
"reportOptionalOperand": {
|
||||
"limit": 0
|
||||
|
|
@ -90,7 +90,7 @@
|
|||
"limit": 8
|
||||
},
|
||||
"reportReturnType": {
|
||||
"limit": 213
|
||||
"limit": 181
|
||||
},
|
||||
"reportTypedDictNotRequiredAccess": {
|
||||
"limit": 24
|
||||
|
|
@ -99,25 +99,25 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44505
|
||||
"limit": 44368
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38689
|
||||
"limit": 38468
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19770
|
||||
"limit": 19665
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 30264
|
||||
"limit": 30066
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 117
|
||||
"limit": 111
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"limit": 697
|
||||
"limit": 695
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
"limit": 5
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ GET - /audit/{id} - Get audit log by id
|
|||
GET - /audit - Get all audit logs
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Final, Optional
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
#### AUDIT LOGGING ####
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
|
@ -58,33 +58,33 @@ async def get_audit_logs(
|
|||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(10, ge=1, le=100),
|
||||
# Filter parameters
|
||||
changed_by: Optional[str] = Query(
|
||||
changed_by: str | None = Query(
|
||||
None, description="Filter by user or system that performed the action"
|
||||
),
|
||||
changed_by_api_key: Optional[str] = Query(
|
||||
changed_by_api_key: str | None = Query(
|
||||
None, description="Filter by API key hash that performed the action"
|
||||
),
|
||||
action: Optional[str] = Query(
|
||||
action: str | None = Query(
|
||||
None, description="Filter by action type (create, update, delete)"
|
||||
),
|
||||
table_name: Optional[str] = Query(
|
||||
table_name: str | None = Query(
|
||||
None, description="Filter by table name that was modified"
|
||||
),
|
||||
object_id: Optional[str] = Query(
|
||||
object_id: str | None = Query(
|
||||
None, description="Filter by ID of the object that was modified"
|
||||
),
|
||||
start_date: Optional[str] = Query(None, description="Filter logs after this date"),
|
||||
end_date: Optional[str] = Query(None, description="Filter logs before this date"),
|
||||
object_team_id: Optional[str] = Query(
|
||||
start_date: str | None = Query(None, description="Filter logs after this date"),
|
||||
end_date: str | None = Query(None, description="Filter logs before this date"),
|
||||
object_team_id: str | None = Query(
|
||||
None,
|
||||
description="Filter by team_id present in before_value or updated_values JSON (PostgreSQL only)",
|
||||
),
|
||||
object_key_hash: Optional[str] = Query(
|
||||
object_key_hash: str | None = Query(
|
||||
None,
|
||||
description="Filter by token (key hash) present in before_value or updated_values JSON (PostgreSQL only)",
|
||||
),
|
||||
# Sorting parameters
|
||||
sort_by: Optional[str] = Query(
|
||||
sort_by: str | None = Query(
|
||||
None,
|
||||
description="Column to sort by (e.g. 'updated_at', 'action', 'table_name')",
|
||||
),
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ until they're actually needed.
|
|||
|
||||
import importlib
|
||||
import sys
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping
|
||||
from types import ModuleType
|
||||
from typing import TYPE_CHECKING, Any, Final, cast
|
||||
|
||||
|
|
@ -57,10 +57,11 @@ from ._lazy_imports_registry import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import httpx
|
||||
from tiktoken import Encoding
|
||||
|
||||
|
||||
def get_litellm_globals() -> dict:
|
||||
def get_litellm_globals() -> dict[str, object]:
|
||||
"""
|
||||
Get the globals dictionary of the litellm module.
|
||||
|
||||
|
|
@ -70,7 +71,7 @@ def get_litellm_globals() -> dict:
|
|||
return sys.modules["litellm"].__dict__
|
||||
|
||||
|
||||
def _get_utils_globals() -> dict:
|
||||
def _get_utils_globals() -> dict[str, object]:
|
||||
"""
|
||||
Get the globals dictionary of the utils module.
|
||||
|
||||
|
|
@ -80,6 +81,11 @@ def _get_utils_globals() -> dict:
|
|||
return sys.modules["litellm.utils"].__dict__
|
||||
|
||||
|
||||
def _get_module_level_client_timeout(litellm_globals: Mapping[str, Any]) -> "float | httpx.Timeout | None":
|
||||
"""Read the configured `litellm.request_timeout` used for the module level http clients."""
|
||||
return litellm_globals.get("request_timeout")
|
||||
|
||||
|
||||
# These are special lazy loaders for things that are used internally
|
||||
# They're separate from the main lazy import system because they have specific use cases
|
||||
|
||||
|
|
@ -435,8 +441,8 @@ def _lazy_import_http_handlers(name: str) -> object:
|
|||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
|
||||
# Get timeout from module config (if set)
|
||||
timeout = _globals.get("request_timeout")
|
||||
params: Final = {"timeout": timeout, "client_alias": "module level aclient"}
|
||||
async_timeout: Final = _get_module_level_client_timeout(_globals)
|
||||
params: Final = {"timeout": async_timeout, "client_alias": "module level aclient"}
|
||||
|
||||
# Create the client instance
|
||||
provider_id: Final = cast(Any, "litellm_module_level_client")
|
||||
|
|
@ -453,8 +459,8 @@ def _lazy_import_http_handlers(name: str) -> object:
|
|||
# Create a sync HTTP client
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
|
||||
timeout = _globals.get("request_timeout")
|
||||
sync_client: Final = HTTPHandler(timeout=timeout)
|
||||
sync_timeout: Final = _get_module_level_client_timeout(_globals)
|
||||
sync_client: Final = HTTPHandler(timeout=sync_timeout)
|
||||
|
||||
# Cache it
|
||||
_globals["module_level_client"] = sync_client
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ RedisSemanticCache since those are backend agnostic.
|
|||
import asyncio
|
||||
import hashlib
|
||||
import os
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Final
|
||||
|
||||
|
|
@ -64,7 +65,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
async_client: AsyncRedis | None = None,
|
||||
embedding_max_input_tokens: int | None = None,
|
||||
embedding_timeout: float | None = None,
|
||||
**kwargs: Any,
|
||||
**kwargs: object,
|
||||
):
|
||||
if similarity_threshold is None:
|
||||
raise ValueError("similarity_threshold must be provided, passed None")
|
||||
|
|
@ -87,11 +88,13 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
self.key_prefix = f"{self.index_name}:"
|
||||
self._index_dim: int | None = None
|
||||
|
||||
resolved_url = None
|
||||
if sync_client is None or async_client is None:
|
||||
resolved_url = redis_url or self._build_valkey_url(host, port, password, ssl)
|
||||
self.sync_client = sync_client if sync_client is not None else Redis.from_url(resolved_url)
|
||||
self.async_client = async_client if async_client is not None else AsyncRedis.from_url(resolved_url)
|
||||
if sync_client is not None and async_client is not None:
|
||||
self.sync_client = sync_client
|
||||
self.async_client = async_client
|
||||
else:
|
||||
resolved_url: Final = redis_url or self._build_valkey_url(host, port, password, ssl)
|
||||
self.sync_client = sync_client if sync_client is not None else Redis.from_url(resolved_url)
|
||||
self.async_client = async_client if async_client is not None else AsyncRedis.from_url(resolved_url)
|
||||
|
||||
print_verbose(f"Valkey semantic-cache initializing index - {self.index_name}")
|
||||
|
||||
|
|
@ -118,7 +121,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
return hashlib.sha256(str(key).encode("utf-8")).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _embedding_to_bytes(embedding: list[float]) -> bytes:
|
||||
def _embedding_to_bytes(embedding: Sequence[float]) -> bytes:
|
||||
return pack_vector(embedding)
|
||||
|
||||
def _index_schema(self, dim: int) -> tuple[TagField, VectorField]:
|
||||
|
|
@ -192,7 +195,9 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
def _doc_key(self, key: str) -> str:
|
||||
return f"{self.key_prefix}{self._scope_tag(key)}:{uuid.uuid4()}"
|
||||
|
||||
def _doc_mapping(self, key: str, prompt: str, value_str: str, embedding: list[float]) -> dict:
|
||||
def _doc_mapping(
|
||||
self, key: str, prompt: str, value_str: str, embedding: Sequence[float]
|
||||
) -> Mapping[str | bytes, str | bytes]:
|
||||
return {
|
||||
self.CACHE_KEY_FIELD_NAME: self._scope_tag(key),
|
||||
self.PROMPT_FIELD_NAME: prompt,
|
||||
|
|
@ -208,30 +213,49 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
)
|
||||
return Query(query_string).return_fields(self.RESPONSE_FIELD_NAME, self.DISTANCE_FIELD_NAME).dialect(2)
|
||||
|
||||
async def _async_search(self, key: str, embedding: Sequence[float]) -> object:
|
||||
"""Run the KNN query on the async client, stopping the untyped search surface here."""
|
||||
return await self.async_client.ft(self.index_name).search(
|
||||
self._knn_query(key),
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)}, # pyright: ignore[reportArgumentType] # redis stubs omit bytes; KNN vectors are raw bytes at runtime
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _first_hit(cls, search_result: Any) -> _ValkeyCacheHit | None:
|
||||
docs: Final = getattr(search_result, "docs", [])
|
||||
def _first_hit(cls, search_result: object) -> _ValkeyCacheHit | None:
|
||||
docs: Final[Sequence[object]] = getattr(search_result, "docs", [])
|
||||
if not docs:
|
||||
return None
|
||||
doc: Final = docs[0]
|
||||
response_field: Final[object] = getattr(doc, cls.RESPONSE_FIELD_NAME)
|
||||
distance_field: Final[str | bytes | float] = getattr(doc, cls.DISTANCE_FIELD_NAME)
|
||||
return _ValkeyCacheHit(
|
||||
response=str(getattr(doc, cls.RESPONSE_FIELD_NAME)),
|
||||
distance=float(getattr(doc, cls.DISTANCE_FIELD_NAME)),
|
||||
response=str(response_field),
|
||||
distance=float(distance_field),
|
||||
)
|
||||
|
||||
def _resolve_hit(self, hit: _ValkeyCacheHit | None, key: str, **kwargs: Any) -> Any:
|
||||
@staticmethod
|
||||
def _record_similarity(kwargs: dict[str, Any], similarity: float) -> None:
|
||||
"""Stamp the semantic-similarity score onto the request metadata carried in ``kwargs``."""
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
|
||||
|
||||
@staticmethod
|
||||
def _embedding_metadata(kwargs: dict[str, Any]) -> dict[str, Any] | None:
|
||||
"""The request metadata forwarded to the embedding call."""
|
||||
return kwargs.get("metadata")
|
||||
|
||||
def _resolve_hit(self, hit: _ValkeyCacheHit | None, key: str, **kwargs: object) -> object:
|
||||
if hit is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
self._record_similarity(kwargs, 0.0)
|
||||
return None
|
||||
|
||||
similarity: Final = 1 - hit.distance
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
|
||||
self._record_similarity(kwargs, similarity)
|
||||
|
||||
if similarity < self.similarity_threshold:
|
||||
return None
|
||||
return self._get_cache_logic(cached_response=hit.response)
|
||||
|
||||
def set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
|
||||
def set_cache(self, key: str, value: object, **kwargs: object) -> None:
|
||||
print_verbose(f"Valkey semantic-cache set_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
|
|
@ -250,12 +274,12 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
except Exception as e:
|
||||
print_verbose(f"Error in Valkey semantic-cache set_cache: {e}")
|
||||
|
||||
def get_cache(self, key: str, **kwargs: Any) -> Any:
|
||||
def get_cache(self, key: str, **kwargs: object) -> object:
|
||||
print_verbose(f"Valkey semantic-cache get_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
self._record_similarity(kwargs, 0.0)
|
||||
return None
|
||||
|
||||
embedding: Final = self._get_embedding(prompt)
|
||||
|
|
@ -263,14 +287,14 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
|
||||
search_result: Final = self.sync_client.ft(self.index_name).search(
|
||||
self._knn_query(key),
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)},
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)}, # pyright: ignore[reportArgumentType] # redis stubs omit bytes; KNN vectors are raw bytes at runtime
|
||||
)
|
||||
return self._resolve_hit(self._first_hit(search_result), key, **kwargs)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in Valkey semantic-cache get_cache: {e}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
self._record_similarity(kwargs, 0.0)
|
||||
|
||||
async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
|
||||
async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None:
|
||||
print_verbose(f"Async Valkey semantic-cache set_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
|
|
@ -278,7 +302,7 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
embedding: Final = await self._get_async_embedding(prompt, metadata=self._embedding_metadata(kwargs))
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
doc_key: Final = self._doc_key(key)
|
||||
|
|
@ -289,31 +313,28 @@ class ValkeySemanticCache(RedisSemanticCache):
|
|||
except Exception as e:
|
||||
print_verbose(f"Error in async Valkey semantic-cache set_cache: {e}")
|
||||
|
||||
async def async_get_cache(self, key: str, **kwargs: Any) -> Any:
|
||||
async def async_get_cache(self, key: str, **kwargs: object) -> object:
|
||||
print_verbose(f"Async Valkey semantic-cache get_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt: Final = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
self._record_similarity(kwargs, 0.0)
|
||||
return None
|
||||
|
||||
embedding: Final = await self._get_async_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
embedding: Final = await self._get_async_embedding(prompt, metadata=self._embedding_metadata(kwargs))
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
search_result: Final = await self.async_client.ft(self.index_name).search(
|
||||
self._knn_query(key),
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)},
|
||||
)
|
||||
search_result: Final[object] = await self._async_search(key, embedding)
|
||||
return self._resolve_hit(self._first_hit(search_result), key, **kwargs)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in async Valkey semantic-cache get_cache: {e}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
self._record_similarity(kwargs, 0.0)
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs: Any) -> None:
|
||||
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, object]], **kwargs: object) -> None:
|
||||
try:
|
||||
await asyncio.gather(*[self.async_set_cache(key, value, **kwargs) for key, value in cache_list])
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in Valkey semantic-cache async_set_cache_pipeline: {e}")
|
||||
|
||||
async def _index_info(self) -> dict:
|
||||
async def _index_info(self) -> Mapping[str, object]:
|
||||
return await self.async_client.ft(self.index_name).info()
|
||||
|
|
|
|||
|
|
@ -163,15 +163,11 @@ class BitBucketClient:
|
|||
response.raise_for_status()
|
||||
|
||||
data: Final[BitBucketSrcListing] = response.json()
|
||||
files: Final[list[str]] = []
|
||||
|
||||
for item in data.get("values", []):
|
||||
if item.get("type") == "commit_file":
|
||||
file_path = item.get("path", "")
|
||||
if file_path.endswith(file_extension):
|
||||
files.append(file_path)
|
||||
|
||||
return files
|
||||
return [
|
||||
file_path
|
||||
for item in data.get("values", [])
|
||||
if item.get("type") == "commit_file" and (file_path := item.get("path", "")).endswith(file_extension)
|
||||
]
|
||||
|
||||
except Exception as e:
|
||||
# Check if it's an HTTP error
|
||||
|
|
|
|||
|
|
@ -7,7 +7,10 @@ litellm_content_retrieve tool calls server-side via the typed agentic loop plan.
|
|||
|
||||
import time
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, cast
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, ClassVar, Final, Protocol, cast
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.compression import compress
|
||||
|
|
@ -22,13 +25,23 @@ from litellm.types.integrations.custom_logger import (
|
|||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
LITELLM_CONTENT_RETRIEVE_TOOL_NAME: Final = "litellm_content_retrieve"
|
||||
_CACHE_TTL_SECONDS: Final = 15 * 60
|
||||
|
||||
|
||||
class _AgenticLoopParams(TypedDict, total=False):
|
||||
"""The ``agentic_loop_params`` entry the agentic loop driver records on the logging object."""
|
||||
|
||||
model: ReadOnly[str]
|
||||
|
||||
|
||||
class _AgenticLoopLoggingObj(Protocol):
|
||||
"""Logging object view exposing the untyped call details this handler reads."""
|
||||
|
||||
@property
|
||||
def model_call_details(self) -> Mapping[str, _AgenticLoopParams]: ...
|
||||
|
||||
|
||||
def _compression_savings_from_counts(
|
||||
original_tokens: object, compressed_tokens: object
|
||||
) -> CompressionSavingsMetadata | None:
|
||||
|
|
@ -83,7 +96,7 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
compression_trigger: int = 200_000,
|
||||
compression_target: int | None = None,
|
||||
embedding_model: str | None = None,
|
||||
embedding_model_params: dict[str, Any] | None = None,
|
||||
embedding_model_params: dict[str, object] | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.enabled = enabled
|
||||
|
|
@ -106,7 +119,7 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
@staticmethod
|
||||
def initialize_from_proxy_config(
|
||||
litellm_settings: dict[str, Any],
|
||||
callback_specific_params: dict[str, Any],
|
||||
callback_specific_params: Mapping[str, object],
|
||||
) -> "CompressionInterceptionLogger":
|
||||
compression_params: CompressionInterceptionConfig = {}
|
||||
if "compression_interception_params" in litellm_settings:
|
||||
|
|
@ -120,7 +133,9 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
)
|
||||
return CompressionInterceptionLogger.from_config_yaml(compression_params)
|
||||
|
||||
async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None:
|
||||
async def async_pre_call_deployment_hook(
|
||||
self, kwargs: dict[str, Any], call_type: CallTypes | None
|
||||
) -> dict[str, object] | None:
|
||||
if not self.enabled:
|
||||
return None
|
||||
if call_type is not None and call_type != CallTypes.anthropic_messages:
|
||||
|
|
@ -150,7 +165,7 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
|
||||
cache: Final = cast(dict[str, str], compressed.get("cache", {}))
|
||||
skip_reason: Final = cast(str | None, compressed.get("compression_skipped_reason"))
|
||||
compressed_tools: Final = cast(list[dict[str, Any]], compressed.get("tools", []))
|
||||
compressed_tools: Final = cast(list[dict[str, object]], compressed.get("tools", []))
|
||||
|
||||
# Only mutate kwargs when compression actually produced a result.
|
||||
# If compression was a no-op (below trigger, invalid tool sequence, etc.),
|
||||
|
|
@ -161,7 +176,7 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
kwargs["messages"] = compressed["messages"]
|
||||
if compressed_tools:
|
||||
kwargs["tools"] = self._merge_tools(
|
||||
existing_tools=cast(list[dict[str, Any]] | None, kwargs.get("tools")),
|
||||
existing_tools=cast(list[dict[str, object]] | None, kwargs.get("tools")),
|
||||
compressed_tools=compressed_tools,
|
||||
)
|
||||
call_id = cast(str | None, kwargs.get("litellm_call_id"))
|
||||
|
|
@ -194,14 +209,14 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
|
||||
async def async_should_run_agentic_loop(
|
||||
self,
|
||||
response: Any,
|
||||
response: object,
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
tools: list[dict] | None,
|
||||
messages: Sequence[Mapping[str, object]],
|
||||
tools: Sequence[Mapping[str, object]] | None,
|
||||
stream: bool,
|
||||
custom_llm_provider: str,
|
||||
kwargs: dict,
|
||||
) -> tuple[bool, dict]:
|
||||
kwargs: Mapping[str, object],
|
||||
) -> tuple[bool, dict[str, object]]:
|
||||
if not self.enabled:
|
||||
return False, {}
|
||||
if not self._has_retrieval_tool(tools):
|
||||
|
|
@ -219,19 +234,19 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
|
||||
async def async_build_agentic_loop_plan(
|
||||
self,
|
||||
tools: dict,
|
||||
tools: Mapping[str, object],
|
||||
model: str,
|
||||
messages: list[dict],
|
||||
response: Any,
|
||||
anthropic_messages_provider_config: Any,
|
||||
anthropic_messages_optional_request_params: dict,
|
||||
logging_obj: "LiteLLMLoggingObj | None",
|
||||
messages: list[dict[str, object]],
|
||||
response: object,
|
||||
anthropic_messages_provider_config: object,
|
||||
anthropic_messages_optional_request_params: Mapping[str, object],
|
||||
logging_obj: _AgenticLoopLoggingObj | None,
|
||||
stream: bool,
|
||||
kwargs: dict,
|
||||
kwargs: Mapping[str, object],
|
||||
) -> AgenticLoopPlan:
|
||||
self._prune_expired_cache()
|
||||
tool_calls: Final = cast(list[dict[str, Any]], tools.get("tool_calls", []))
|
||||
thinking_blocks: Final = cast(list[dict[str, Any]], tools.get("thinking_blocks", []))
|
||||
tool_calls: Final = cast(list[dict[str, object]], tools.get("tool_calls", []))
|
||||
thinking_blocks: Final = cast(list[dict[str, object]], tools.get("thinking_blocks", []))
|
||||
|
||||
call_id: Final = self._resolve_call_id(logging_obj=logging_obj, kwargs=kwargs)
|
||||
cache: Final = self._get_cache(call_id=call_id)
|
||||
|
|
@ -274,7 +289,7 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
full_model_name = model
|
||||
if logging_obj is not None:
|
||||
agentic_params: Final = logging_obj.model_call_details.get("agentic_loop_params", {})
|
||||
full_model_name = cast(str, agentic_params.get("model", model))
|
||||
full_model_name = agentic_params.get("model", model)
|
||||
|
||||
request_patch: Final = AgenticLoopRequestPatch(
|
||||
model=full_model_name,
|
||||
|
|
@ -309,15 +324,15 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
return {}
|
||||
return cache_entry[0]
|
||||
|
||||
def _resolve_call_id(self, logging_obj: Any, kwargs: dict[str, Any]) -> str | None:
|
||||
def _resolve_call_id(self, logging_obj: _AgenticLoopLoggingObj | None, kwargs: Mapping[str, object]) -> str | None:
|
||||
if logging_obj is not None:
|
||||
logging_call_id: Final = getattr(logging_obj, "litellm_call_id", None)
|
||||
if isinstance(logging_call_id, str) and logging_call_id:
|
||||
return logging_call_id
|
||||
kwargs_call_id: Final = kwargs.get("litellm_call_id")
|
||||
return cast(str | None, kwargs_call_id if isinstance(kwargs_call_id, str) else None)
|
||||
return kwargs_call_id if isinstance(kwargs_call_id, str) else None
|
||||
|
||||
def _resolve_retrieval_content(self, tool_call: dict[str, Any], cache: dict[str, str]) -> str:
|
||||
def _resolve_retrieval_content(self, tool_call: Mapping[str, object], cache: Mapping[str, str]) -> str:
|
||||
raw_input: Final = tool_call.get("input", {})
|
||||
key = ""
|
||||
if isinstance(raw_input, dict):
|
||||
|
|
@ -328,7 +343,9 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
return cache[key]
|
||||
return f"[compressed content key '{key}' not found]"
|
||||
|
||||
def _extract_retrieval_tool_calls(self, response: Any) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||
def _extract_retrieval_tool_calls(
|
||||
self, response: object
|
||||
) -> tuple[list[dict[str, object]], list[dict[str, object]]]:
|
||||
if isinstance(response, dict):
|
||||
content = response.get("content", [])
|
||||
else:
|
||||
|
|
@ -337,8 +354,8 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
if not isinstance(content, list):
|
||||
return [], []
|
||||
|
||||
tool_calls: Final[list[dict[str, Any]]] = []
|
||||
thinking_blocks: Final[list[dict[str, Any]]] = []
|
||||
tool_calls: Final[list[dict[str, object]]] = []
|
||||
thinking_blocks: Final[list[dict[str, object]]] = []
|
||||
|
||||
for block in content:
|
||||
if isinstance(block, dict):
|
||||
|
|
@ -385,13 +402,13 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
|
||||
return tool_calls, thinking_blocks
|
||||
|
||||
def _prepare_followup_kwargs(self, kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
def _prepare_followup_kwargs(self, kwargs: Mapping[str, object]) -> dict[str, object]:
|
||||
internal_keys: Final = {"litellm_logging_obj"}
|
||||
return {
|
||||
k: v for k, v in kwargs.items() if not k.startswith("_compression_interception") and k not in internal_keys
|
||||
}
|
||||
|
||||
def _has_retrieval_tool(self, tools: Any) -> bool:
|
||||
def _has_retrieval_tool(self, tools: object) -> bool:
|
||||
if not isinstance(tools, list):
|
||||
return False
|
||||
for tool in tools:
|
||||
|
|
@ -407,9 +424,9 @@ class CompressionInterceptionLogger(CustomLogger):
|
|||
|
||||
def _merge_tools(
|
||||
self,
|
||||
existing_tools: list[dict[str, Any]] | None,
|
||||
compressed_tools: list[dict[str, Any]],
|
||||
) -> list[dict[str, Any]]:
|
||||
existing_tools: Sequence[Mapping[str, object]] | None,
|
||||
compressed_tools: Sequence[Mapping[str, object]],
|
||||
) -> list[Mapping[str, object]]:
|
||||
merged: Final = list(existing_tools or [])
|
||||
if self._has_retrieval_tool(merged):
|
||||
return merged
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
# On success, logs events to Promptlayer
|
||||
import re
|
||||
import traceback
|
||||
from collections.abc import AsyncGenerator, Mapping
|
||||
from collections.abc import AsyncGenerator, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -123,11 +123,11 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
return []
|
||||
|
||||
callbacks: Final = AllCallbacks()
|
||||
callback_info: Final = getattr(callbacks, lookup_name, None)
|
||||
callback_info: Final[object] = getattr(callbacks, lookup_name, None)
|
||||
if callback_info is None:
|
||||
return []
|
||||
|
||||
params: Final = getattr(callback_info, "litellm_callback_params", None)
|
||||
params: Final[Sequence[str] | None] = getattr(callback_info, "litellm_callback_params", None)
|
||||
if not params:
|
||||
return []
|
||||
|
||||
|
|
@ -851,7 +851,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
- Converting to string and then truncating the logged content catches this
|
||||
2. We want to avoid modifying the original `messages`, `response`, and `error_str` in the logging payload since these are in kwargs and could be returned to the user
|
||||
"""
|
||||
field_value: Final = standard_logging_object.get(field_name)
|
||||
field_value: Final[object] = standard_logging_object.get(field_name)
|
||||
if field_value:
|
||||
str_value: Final = str(field_value)
|
||||
if len(str_value) > max_length:
|
||||
|
|
@ -1005,8 +1005,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
• Keep untyped or text content.
|
||||
• Recursively redact inline base64 blobs in *any* string field, at any depth.
|
||||
"""
|
||||
raw_messages: Final[Any] = payload.get("messages", [])
|
||||
messages: Final[list[Any]] = raw_messages if isinstance(raw_messages, list) else []
|
||||
raw_messages: Final[object] = payload.get("messages", [])
|
||||
messages: Final[list[object]] = raw_messages if isinstance(raw_messages, list) else []
|
||||
verbose_logger.debug("[CustomLogger] Stripping base64 from %s messages", len(messages))
|
||||
|
||||
if messages:
|
||||
|
|
@ -1037,8 +1037,8 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
• Keep untyped or text content.
|
||||
• Recursively redact inline base64 blobs in *any* string field, at any depth.
|
||||
"""
|
||||
raw_messages: Final[Any] = payload.get("messages", [])
|
||||
messages: Final[list[Any]] = raw_messages if isinstance(raw_messages, list) else []
|
||||
raw_messages: Final[object] = payload.get("messages", [])
|
||||
messages: Final[list[object]] = raw_messages if isinstance(raw_messages, list) else []
|
||||
verbose_logger.debug("[CustomLogger] Stripping base64 from %s messages", len(messages))
|
||||
|
||||
if messages:
|
||||
|
|
@ -1059,7 +1059,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
value: Any,
|
||||
depth: int = 0,
|
||||
max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER,
|
||||
) -> Any:
|
||||
) -> object:
|
||||
"""Recursively redact inline base64 from any nested structure with a max recursion depth limit."""
|
||||
if depth > max_depth:
|
||||
verbose_logger.warning("[CustomLogger] Max recursion depth %s reached while redacting base64", max_depth)
|
||||
|
|
@ -1090,16 +1090,16 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
|
|||
|
||||
def _process_messages(
|
||||
self,
|
||||
messages: list[Any],
|
||||
messages: list[object],
|
||||
max_depth: int = DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER,
|
||||
) -> list[dict[str, Any]]:
|
||||
filtered_messages: Final[list[dict[str, Any]]] = []
|
||||
) -> list[dict[str, object]]:
|
||||
filtered_messages: Final[list[dict[str, object]]] = []
|
||||
for msg in messages:
|
||||
if not isinstance(msg, dict):
|
||||
continue
|
||||
contents: Any = msg.get("content")
|
||||
contents: object = msg.get("content")
|
||||
if isinstance(contents, list):
|
||||
cleaned: list[Any] = []
|
||||
cleaned: list[object] = []
|
||||
for c in contents:
|
||||
if self._should_keep_content(content=c):
|
||||
cleaned.append(self._redact_base64(value=c, max_depth=max_depth))
|
||||
|
|
|
|||
|
|
@ -17,12 +17,12 @@ def build_trace_payload(
|
|||
end_time: datetime,
|
||||
input_data: Any,
|
||||
output_data: Any,
|
||||
metadata: dict[str, Any],
|
||||
metadata: dict[str, object],
|
||||
tags: list[str],
|
||||
thread_id: str | None,
|
||||
) -> types.TracePayload:
|
||||
"""Build a complete trace payload."""
|
||||
trace_name: Final = response_obj.get("object", "unknown type")
|
||||
trace_name: Final[str] = response_obj.get("object", "unknown type")
|
||||
|
||||
return types.TracePayload(
|
||||
project_name=project_name,
|
||||
|
|
@ -47,7 +47,7 @@ def build_span_payload(
|
|||
end_time: datetime,
|
||||
input_data: Any,
|
||||
output_data: Any,
|
||||
metadata: dict[str, Any],
|
||||
metadata: dict[str, object],
|
||||
tags: list[str],
|
||||
usage: dict[str, int],
|
||||
provider: str | None = None,
|
||||
|
|
@ -56,9 +56,9 @@ def build_span_payload(
|
|||
"""Build a complete span payload."""
|
||||
span_id: Final = utils.create_uuid7()
|
||||
|
||||
model: Final = response_obj.get("model", "unknown-model")
|
||||
obj_type: Final = response_obj.get("object", "unknown-object")
|
||||
created: Final = response_obj.get("created", 0)
|
||||
model: Final[str] = response_obj.get("model", "unknown-model")
|
||||
obj_type: Final[str] = response_obj.get("object", "unknown-object")
|
||||
created: Final[int] = response_obj.get("created", 0)
|
||||
span_name: Final = f"{model}_{obj_type}_{created}"
|
||||
|
||||
_logging.verbose_logger.debug("OpikLogger creating span with id %s for trace %s", span_id, trace_id)
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import asyncio
|
|||
import math
|
||||
import uuid
|
||||
from collections.abc import AsyncIterator, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Never, TypedDict, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Never, TypedDict, TypeVar, cast
|
||||
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
|
|
@ -85,6 +85,10 @@ WEBSEARCH_EMIT_NATIVE_BLOCKS_KEY: Final = "_websearch_interception_emit_native_b
|
|||
# ``web_search_tool_result`` blocks to inject into the final response.
|
||||
WEBSEARCH_NATIVE_BLOCKS_METADATA_KEY: Final = "websearch_native_blocks"
|
||||
|
||||
_RESPONSE_CONTENT_FIELD: Final = "content"
|
||||
|
||||
_ResponseT: Final = TypeVar("_ResponseT")
|
||||
|
||||
|
||||
class _PlanMetadataView(TypedDict):
|
||||
websearch_native_blocks: Sequence[Mapping[str, object]] | None
|
||||
|
|
@ -1031,17 +1035,17 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _inject_native_blocks(response: Any, native_blocks: Sequence[Mapping[str, object]]) -> Any:
|
||||
def _inject_native_blocks(response: _ResponseT, native_blocks: Sequence[Mapping[str, object]]) -> _ResponseT:
|
||||
"""Prepend native blocks to response content, dict or object form."""
|
||||
if not native_blocks:
|
||||
return response
|
||||
if isinstance(response, dict):
|
||||
existing = response.get("content") or []
|
||||
response["content"] = list(native_blocks) + list(existing)
|
||||
existing = response.get(_RESPONSE_CONTENT_FIELD) or []
|
||||
response[_RESPONSE_CONTENT_FIELD] = list(native_blocks) + list(existing)
|
||||
return response
|
||||
existing = getattr(response, "content", None) or []
|
||||
existing = getattr(response, _RESPONSE_CONTENT_FIELD, None) or []
|
||||
try:
|
||||
response.content = list(native_blocks) + list(existing)
|
||||
setattr(response, _RESPONSE_CONTENT_FIELD, list(native_blocks) + list(existing))
|
||||
except (AttributeError, TypeError):
|
||||
# Object refused write — fall through and leave the response
|
||||
# untouched rather than crash the request.
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@ import datetime
|
|||
from collections.abc import Mapping
|
||||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.constants import LITELLM_DETAILED_TIMING
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.litellm_core_utils.llm_response_utils.get_api_base import get_api_base
|
||||
|
|
@ -59,11 +61,7 @@ class ResponseMetadata:
|
|||
@property
|
||||
def supports_response_time(self) -> bool:
|
||||
"""Check if response type supports timing metrics"""
|
||||
return (
|
||||
isinstance(self.result, ModelResponse)
|
||||
or isinstance(self.result, EmbeddingResponse)
|
||||
or isinstance(self.result, TranscriptionResponse)
|
||||
)
|
||||
return isinstance(self.result, (ModelResponse, EmbeddingResponse, TranscriptionResponse))
|
||||
|
||||
def set_hidden_params(self, logging_obj: LiteLLMLoggingObject, model: str | None, kwargs: dict) -> None:
|
||||
"""Set hidden parameters on the response"""
|
||||
|
|
@ -79,7 +77,7 @@ class ResponseMetadata:
|
|||
result=self.result, litellm_model_name=model, router_model_id=model_id
|
||||
),
|
||||
"additional_headers": process_response_headers(
|
||||
self._get_value_from_hidden_params("additional_headers") or {},
|
||||
self._get_additional_headers_from_hidden_params() or {},
|
||||
preserve_litellm_internal_headers=True,
|
||||
),
|
||||
"litellm_model_name": model,
|
||||
|
|
@ -98,12 +96,12 @@ class ResponseMetadata:
|
|||
for key, value in new_params.items():
|
||||
setattr(self._hidden_params, key, value)
|
||||
|
||||
def _get_value_from_hidden_params(self, key: str) -> Any | None:
|
||||
"""Get value from hidden params - handles when self._hidden_params is a dict or HiddenParams object"""
|
||||
def _get_additional_headers_from_hidden_params(self) -> httpx.Headers | dict[str, str] | None:
|
||||
"""Get `additional_headers` from hidden params - handles when self._hidden_params is a dict or HiddenParams object"""
|
||||
if isinstance(self._hidden_params, dict):
|
||||
return self._hidden_params.get(key, None)
|
||||
return self._hidden_params.get("additional_headers", None)
|
||||
elif isinstance(self._hidden_params, HiddenParams):
|
||||
return getattr(self._hidden_params, key, None)
|
||||
return getattr(self._hidden_params, "additional_headers", None)
|
||||
|
||||
def set_timing_metrics(
|
||||
self,
|
||||
|
|
@ -129,7 +127,7 @@ class ResponseMetadata:
|
|||
#########################################################
|
||||
# 2. Add callback processing duration
|
||||
#########################################################
|
||||
callback_duration_ms: Final = getattr(logging_obj, "callback_duration_ms", None)
|
||||
callback_duration_ms: Final[float | None] = getattr(logging_obj, "callback_duration_ms", None)
|
||||
if callback_duration_ms is not None:
|
||||
self._update_hidden_params(
|
||||
{
|
||||
|
|
@ -142,17 +140,17 @@ class ResponseMetadata:
|
|||
#########################################################
|
||||
llm_api_duration_ms: Final = logging_obj.model_call_details.get("llm_api_duration_ms")
|
||||
if LITELLM_DETAILED_TIMING and llm_api_duration_ms is not None:
|
||||
detailed: Final[dict] = {
|
||||
detailed: Final[dict[str, float]] = {
|
||||
"timing_llm_api_ms": round(llm_api_duration_ms, 4),
|
||||
}
|
||||
|
||||
# message copy time from Logging.__init__()
|
||||
msg_copy_ms: Final = getattr(logging_obj, "message_copy_duration_ms", None)
|
||||
msg_copy_ms: Final[float | None] = getattr(logging_obj, "message_copy_duration_ms", None)
|
||||
if msg_copy_ms is not None:
|
||||
detailed["timing_message_copy_ms"] = round(msg_copy_ms, 4)
|
||||
|
||||
# pre-processing = time from request start to LLM API call start
|
||||
api_call_start: Final = logging_obj.model_call_details.get("api_call_start_time")
|
||||
api_call_start: Final[datetime.datetime | None] = logging_obj.model_call_details.get("api_call_start_time")
|
||||
if api_call_start is not None and start_time is not None:
|
||||
pre_ms: Final = (api_call_start - start_time).total_seconds() * 1000
|
||||
detailed["timing_pre_processing_ms"] = round(pre_ms, 4)
|
||||
|
|
|
|||
|
|
@ -93,7 +93,7 @@ def print_verbose(print_statement: object):
|
|||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ProviderChunkParsed:
|
||||
response_obj: dict[str, Any]
|
||||
response_obj: dict[str, object]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -1288,7 +1288,7 @@ class CustomStreamWrapper:
|
|||
for key, value in anthropic_response_obj["provider_specific_fields"].items():
|
||||
setattr(model_response, key, value)
|
||||
|
||||
response_obj = cast(dict[str, Any], anthropic_response_obj)
|
||||
response_obj = cast(dict[str, object], anthropic_response_obj)
|
||||
elif self.model == "replicate" or self.custom_llm_provider == "replicate":
|
||||
response_obj = self.handle_replicate_chunk(chunk)
|
||||
completion_obj["content"] = response_obj["text"]
|
||||
|
|
@ -1444,7 +1444,7 @@ class CustomStreamWrapper:
|
|||
if not isinstance(chunk, str):
|
||||
raise ValueError(f"chunk is not a string: {chunk}")
|
||||
response_obj = cast(
|
||||
dict[str, Any],
|
||||
dict[str, object],
|
||||
litellm.CodestralTextCompletionConfig()._chunk_parser(chunk),
|
||||
)
|
||||
completion_obj["content"] = response_obj["text"]
|
||||
|
|
@ -2551,7 +2551,7 @@ def calculate_total_usage(chunks: list[ModelResponse]) -> Usage:
|
|||
|
||||
prompt_tokens: int = 0
|
||||
completion_tokens: int = 0
|
||||
latest_usage_chunk = None
|
||||
latest_usage_chunk: Usage | Mapping[str, int] | None = None
|
||||
prompt_tokens_details: PromptTokensDetailsWrapper | None = None
|
||||
completion_tokens_details: CompletionTokensDetailsWrapper | None = None
|
||||
cache_creation_token_details: CacheCreationTokenDetails | None = None
|
||||
|
|
|
|||
|
|
@ -417,7 +417,7 @@ def _extract_redirect_url(response: httpx.Response, request_url: str) -> str:
|
|||
return str(httpx.URL(request_url).join(location))
|
||||
|
||||
|
||||
def safe_get(client: Any, url: str, **kwargs: Any) -> Any:
|
||||
def safe_get(client: Any, url: str, **kwargs: Any) -> httpx.Response:
|
||||
"""
|
||||
Fetch a user-supplied URL with SSRF protection on every redirect hop.
|
||||
|
||||
|
|
@ -460,7 +460,7 @@ def safe_get(client: Any, url: str, **kwargs: Any) -> Any:
|
|||
raise SSRFError("Too many redirects")
|
||||
|
||||
|
||||
async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any:
|
||||
async def async_safe_get(client: Any, url: str, **kwargs: Any) -> httpx.Response:
|
||||
"""Async version of safe_get."""
|
||||
if not getattr(litellm, "user_url_validation", True):
|
||||
kwargs.setdefault("follow_redirects", True)
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
import json
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, cast
|
||||
|
||||
import httpx
|
||||
from httpx import Headers, Response
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
|
||||
|
|
@ -21,6 +23,29 @@ else:
|
|||
LoggingClass = Any
|
||||
|
||||
|
||||
class AnthropicBatchRequestCounts(TypedDict, total=False):
|
||||
"""The ``request_counts`` object of an Anthropic Message Batch."""
|
||||
|
||||
processing: ReadOnly[int]
|
||||
succeeded: ReadOnly[int]
|
||||
errored: ReadOnly[int]
|
||||
canceled: ReadOnly[int]
|
||||
expired: ReadOnly[int]
|
||||
|
||||
|
||||
class AnthropicMessageBatch(TypedDict, total=False):
|
||||
"""The fields of an Anthropic Message Batch that map onto an OpenAI Batch."""
|
||||
|
||||
id: ReadOnly[str]
|
||||
processing_status: ReadOnly[str]
|
||||
created_at: ReadOnly[str | None]
|
||||
ended_at: ReadOnly[str | None]
|
||||
expires_at: ReadOnly[str | None]
|
||||
cancel_initiated_at: ReadOnly[str | None]
|
||||
archived_at: ReadOnly[str | None]
|
||||
request_counts: ReadOnly[AnthropicBatchRequestCounts]
|
||||
|
||||
|
||||
class AnthropicBatchesConfig(BaseBatchesConfig):
|
||||
def __init__(self):
|
||||
from ..chat.transformation import AnthropicConfig
|
||||
|
|
@ -85,7 +110,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
|
|||
create_batch_data: CreateBatchRequest,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> bytes | str | dict[str, Any]:
|
||||
) -> bytes | str | dict[str, object]:
|
||||
"""
|
||||
Transform the batch creation request to Anthropic format.
|
||||
|
||||
|
|
@ -135,7 +160,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
|
|||
batch_id: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> bytes | str | dict[str, Any]:
|
||||
) -> bytes | str | dict[str, object]:
|
||||
"""
|
||||
Transform batch retrieval request for Anthropic.
|
||||
|
||||
|
|
@ -154,7 +179,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
|
|||
) -> LiteLLMBatch:
|
||||
"""Transform Anthropic MessageBatch retrieval response to LiteLLM format."""
|
||||
try:
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final[AnthropicMessageBatch] = raw_response.json()
|
||||
except Exception as e:
|
||||
raise ValueError(f"Failed to parse Anthropic batch response: {e}")
|
||||
|
||||
|
|
@ -163,18 +188,20 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
|
|||
processing_status: Final = response_data.get("processing_status", "in_progress")
|
||||
|
||||
# Map Anthropic processing_status to OpenAI status
|
||||
status_mapping: dict[
|
||||
str,
|
||||
Literal[
|
||||
"validating",
|
||||
"failed",
|
||||
"in_progress",
|
||||
"finalizing",
|
||||
"completed",
|
||||
"expired",
|
||||
"cancelling",
|
||||
"cancelled",
|
||||
],
|
||||
status_mapping: Final[
|
||||
Mapping[
|
||||
str,
|
||||
Literal[
|
||||
"validating",
|
||||
"failed",
|
||||
"in_progress",
|
||||
"finalizing",
|
||||
"completed",
|
||||
"expired",
|
||||
"cancelling",
|
||||
"cancelled",
|
||||
],
|
||||
]
|
||||
] = {
|
||||
"in_progress": "in_progress",
|
||||
"canceling": "cancelling",
|
||||
|
|
@ -281,7 +308,7 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
|
|||
if not line:
|
||||
continue
|
||||
try:
|
||||
response_json = json.loads(line)
|
||||
response_json: Mapping[str, Mapping[str, dict[str, object]]] = json.loads(line)
|
||||
# Update model_response with the parsed JSON
|
||||
completion_response = response_json["result"]["message"]
|
||||
transformed_response = self.anthropic_chat_config.transform_parsed_response(
|
||||
|
|
|
|||
|
|
@ -160,7 +160,7 @@ class LiteLLMMessagesToResponsesAPIHandler:
|
|||
top_k: int | None = None,
|
||||
top_p: float | None = None,
|
||||
output_format: AnthropicOutputSchema | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
) -> AnthropicMessagesResponse | AsyncIterator[bytes]:
|
||||
responses_kwargs: Final = _build_responses_kwargs(
|
||||
max_tokens=max_tokens,
|
||||
|
|
@ -214,7 +214,7 @@ class LiteLLMMessagesToResponsesAPIHandler:
|
|||
top_p: float | None = None,
|
||||
output_format: AnthropicOutputSchema | None = None,
|
||||
_is_async: bool = False,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
) -> (
|
||||
AnthropicMessagesResponse
|
||||
| AsyncIterator[bytes]
|
||||
|
|
|
|||
|
|
@ -94,7 +94,7 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
aws_sts_endpoint: str | None = None,
|
||||
aws_bedrock_runtime_endpoint: str | None = None,
|
||||
aws_external_id: str | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
):
|
||||
"""
|
||||
Establish bidirectional streaming connection with Bedrock Nova Sonic.
|
||||
|
|
@ -166,13 +166,16 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
)
|
||||
bedrock_client: Final = BedrockRuntimeClient(config=config)
|
||||
|
||||
async def open_bidirectional_stream() -> BedrockBidirectionalStream:
|
||||
return await bedrock_client.invoke_model_with_bidirectional_stream(
|
||||
InvokeModelWithBidirectionalStreamOperationInput(model_id=model)
|
||||
)
|
||||
|
||||
transformation_config: Final = BedrockRealtimeConfig()
|
||||
|
||||
try:
|
||||
# Initialize the bidirectional stream
|
||||
bedrock_stream: Final = await bedrock_client.invoke_model_with_bidirectional_stream(
|
||||
InvokeModelWithBidirectionalStreamOperationInput(model_id=model)
|
||||
)
|
||||
bedrock_stream: Final = await open_bidirectional_stream()
|
||||
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Bidirectional stream established")
|
||||
|
||||
|
|
@ -243,10 +246,11 @@ class BedrockRealtime(BaseAWSLLM):
|
|||
InvokeModelWithBidirectionalStreamInputChunk,
|
||||
)
|
||||
|
||||
def build_input_chunk(payload: bytes) -> object:
|
||||
return InvokeModelWithBidirectionalStreamInputChunk(value=BidirectionalInputPayloadPart(bytes_=payload))
|
||||
|
||||
async def send_to_bedrock(bedrock_message: str) -> None:
|
||||
event: Final = InvokeModelWithBidirectionalStreamInputChunk(
|
||||
value=BidirectionalInputPayloadPart(bytes_=bedrock_message.encode("utf-8"))
|
||||
)
|
||||
event: Final = build_input_chunk(bedrock_message.encode("utf-8"))
|
||||
await bedrock_stream.input_stream.send(event)
|
||||
verbose_proxy_logger.debug("Bedrock Realtime: Sent to Bedrock: %s", bedrock_message[:200])
|
||||
|
||||
|
|
|
|||
|
|
@ -49,10 +49,10 @@ class ChatGPTToolCallNormalizer:
|
|||
def __getattr__(self, name: str) -> object:
|
||||
return getattr(self._stream, name)
|
||||
|
||||
def __iter__(self):
|
||||
def __iter__(self) -> "ChatGPTToolCallNormalizer":
|
||||
return self
|
||||
|
||||
def __aiter__(self):
|
||||
def __aiter__(self) -> "ChatGPTToolCallNormalizer":
|
||||
return self
|
||||
|
||||
def __next__(self) -> ModelResponseStream:
|
||||
|
|
|
|||
|
|
@ -2,13 +2,16 @@
|
|||
CompactifAI chat completion transformation
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
|
@ -23,6 +26,18 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class CompactifAIResponseFields(TypedDict, total=False):
|
||||
"""The chat completion fields of a CompactifAI response body."""
|
||||
|
||||
id: ReadOnly[str]
|
||||
choices: ReadOnly[Sequence[Mapping[str, object]]]
|
||||
created: ReadOnly[int]
|
||||
model: ReadOnly[str]
|
||||
system_fingerprint: ReadOnly[str | None]
|
||||
usage: ReadOnly[Mapping[str, object]]
|
||||
object: ReadOnly[str]
|
||||
|
||||
|
||||
class CompactifAIChatConfig(OpenAIGPTConfig):
|
||||
"""
|
||||
Configuration class for CompactifAI chat completions.
|
||||
|
|
@ -47,10 +62,10 @@ class CompactifAIChatConfig(OpenAIGPTConfig):
|
|||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
request_data: dict,
|
||||
messages: list,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
request_data: Mapping[str, object],
|
||||
messages: Sequence[AllMessageValues],
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
|
|
@ -81,14 +96,18 @@ class CompactifAIChatConfig(OpenAIGPTConfig):
|
|||
message["content"] = tool_calls[0]["function"].get("arguments", "")
|
||||
message["tool_calls"] = None
|
||||
|
||||
returned_response: Final = ModelResponse(**response_json)
|
||||
response_fields: Final[CompactifAIResponseFields] = response_json
|
||||
|
||||
returned_response: Final = ModelResponse(**response_fields)
|
||||
|
||||
# Set model name with provider prefix
|
||||
returned_response.model = f"compactifai/{model}"
|
||||
|
||||
return returned_response
|
||||
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException:
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: dict[str, str] | httpx.Headers
|
||||
) -> BaseLLMException:
|
||||
"""
|
||||
Get the appropriate error class for CompactifAI errors.
|
||||
Since CompactifAI is OpenAI-compatible, we use OpenAI error handling.
|
||||
|
|
|
|||
|
|
@ -6,11 +6,12 @@ endpoint defined in endpoints.json, eliminating the need for individual handler
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Coroutine
|
||||
from collections.abc import Coroutine, Mapping, Sequence
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
|
|
@ -32,26 +33,58 @@ if TYPE_CHECKING:
|
|||
from litellm.llms.base_llm.containers.transformation import BaseContainerConfig
|
||||
|
||||
|
||||
class EndpointConfig(TypedDict):
|
||||
"""One endpoint entry of ``litellm/containers/endpoints.json``."""
|
||||
|
||||
name: ReadOnly[str]
|
||||
async_name: ReadOnly[str]
|
||||
path: ReadOnly[str]
|
||||
method: ReadOnly[str]
|
||||
path_params: ReadOnly[Sequence[str]]
|
||||
query_params: ReadOnly[Sequence[str]]
|
||||
response_type: ReadOnly[str]
|
||||
is_multipart: NotRequired[ReadOnly[bool]]
|
||||
returns_binary: NotRequired[ReadOnly[bool]]
|
||||
|
||||
|
||||
class EndpointsConfig(TypedDict):
|
||||
"""The parsed ``litellm/containers/endpoints.json`` document."""
|
||||
|
||||
endpoints: ReadOnly[Sequence[EndpointConfig]]
|
||||
|
||||
|
||||
class ContainerErrorDetail(TypedDict, total=False):
|
||||
"""The ``error`` object of a container API error body."""
|
||||
|
||||
message: ReadOnly[str]
|
||||
|
||||
|
||||
class ContainerResponseBody(TypedDict, total=False):
|
||||
"""The fields this handler reads off a container API JSON body."""
|
||||
|
||||
error: ReadOnly[ContainerErrorDetail]
|
||||
|
||||
|
||||
_ContainerResponseModel = ContainerFileListResponse | ContainerFileObject | DeleteContainerFileResponse
|
||||
|
||||
# Response type mapping
|
||||
RESPONSE_TYPES: Final[dict[str, type]] = {
|
||||
RESPONSE_TYPES: Final[Mapping[str, type[_ContainerResponseModel]]] = {
|
||||
"ContainerFileListResponse": ContainerFileListResponse,
|
||||
"ContainerFileObject": ContainerFileObject,
|
||||
"DeleteContainerFileResponse": DeleteContainerFileResponse,
|
||||
}
|
||||
|
||||
ContainerEndpointResponse = (
|
||||
ContainerFileListResponse | ContainerFileObject | DeleteContainerFileResponse | bytes | dict[str, object]
|
||||
)
|
||||
ContainerEndpointResponse = _ContainerResponseModel | bytes | ContainerResponseBody
|
||||
|
||||
|
||||
def _load_endpoints_config() -> dict:
|
||||
def _load_endpoints_config() -> EndpointsConfig:
|
||||
"""Load the endpoints configuration from JSON file."""
|
||||
config_path: Final = Path(__file__).parent.parent.parent / "containers" / "endpoints.json"
|
||||
with open(config_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def _get_endpoint_config(endpoint_name: str) -> dict | None:
|
||||
def _get_endpoint_config(endpoint_name: str) -> EndpointConfig | None:
|
||||
"""Get config for a specific endpoint by name."""
|
||||
config: Final = _load_endpoints_config()
|
||||
for endpoint in config["endpoints"]:
|
||||
|
|
@ -60,10 +93,15 @@ def _get_endpoint_config(endpoint_name: str) -> dict | None:
|
|||
return None
|
||||
|
||||
|
||||
def _response_model(response_type_name: str) -> type[_ContainerResponseModel] | None:
|
||||
"""The pydantic model a container endpoint's ``response_type`` names."""
|
||||
return RESPONSE_TYPES.get(response_type_name)
|
||||
|
||||
|
||||
def _build_url(
|
||||
api_base: str,
|
||||
path_template: str,
|
||||
path_params: dict[str, str],
|
||||
path_params: Mapping[str, object],
|
||||
) -> str:
|
||||
"""Build the full URL by substituting path parameters.
|
||||
|
||||
|
|
@ -93,16 +131,12 @@ def _build_url(
|
|||
|
||||
|
||||
def _build_query_params(
|
||||
query_param_names: list,
|
||||
kwargs: dict[str, Any],
|
||||
) -> dict[str, str]:
|
||||
query_param_names: Sequence[str],
|
||||
kwargs: Mapping[str, object],
|
||||
) -> dict[str, object]:
|
||||
"""Build query parameters from kwargs."""
|
||||
params: Final = {}
|
||||
for param_name in query_param_names:
|
||||
value = kwargs.get(param_name)
|
||||
if value is not None:
|
||||
params[param_name] = str(value) if not isinstance(value, str) else value
|
||||
return params
|
||||
supplied: Final = ((param_name, kwargs.get(param_name)) for param_name in query_param_names)
|
||||
return {name: value if isinstance(value, str) else str(value) for name, value in supplied if value is not None}
|
||||
|
||||
|
||||
def _error_message_from_response(response: httpx.Response) -> str:
|
||||
|
|
@ -136,24 +170,24 @@ def _transform_response(
|
|||
if returns_binary:
|
||||
return response.content
|
||||
|
||||
response_json: Final = response.json()
|
||||
response_json: Final[ContainerResponseBody] = response.json()
|
||||
if "error" in response_json:
|
||||
raise BaseLLMException(
|
||||
status_code=response.status_code,
|
||||
message=response_json.get("error", {}).get("message", str(response_json)),
|
||||
message=response_json["error"].get("message", str(response_json)),
|
||||
headers=dict(response.headers),
|
||||
)
|
||||
|
||||
response_type: Final = RESPONSE_TYPES.get(response_type_name)
|
||||
response_type: Final = _response_model(response_type_name)
|
||||
if response_type:
|
||||
return response_type(**response_json)
|
||||
return response_type.model_validate(response_json)
|
||||
return response_json
|
||||
|
||||
|
||||
def _prepare_multipart_file_upload(
|
||||
file: Any,
|
||||
headers: dict[str, Any],
|
||||
) -> tuple:
|
||||
headers: dict[str, object],
|
||||
) -> tuple[dict[str, tuple[str, bytes, str]], dict[str, object]]:
|
||||
"""
|
||||
Prepare file and headers for multipart upload.
|
||||
|
||||
|
|
@ -178,6 +212,52 @@ def _prepare_multipart_file_upload(
|
|||
return files, headers_copy
|
||||
|
||||
|
||||
def _request_headers(
|
||||
container_provider_config: "BaseContainerConfig",
|
||||
extra_headers: dict[str, object] | None,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
) -> dict[str, object]:
|
||||
"""The provider auth headers for a container request."""
|
||||
return container_provider_config.validate_environment(
|
||||
headers=extra_headers or {},
|
||||
api_key=litellm_params.get("api_key", None),
|
||||
)
|
||||
|
||||
|
||||
def _request_api_base(
|
||||
container_provider_config: "BaseContainerConfig",
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
) -> str:
|
||||
"""The provider base URL for a container request."""
|
||||
return container_provider_config.get_complete_url(
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
|
||||
|
||||
def _sync_http_client(
|
||||
client: HTTPHandler | AsyncHTTPHandler | None,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
) -> HTTPHandler:
|
||||
"""The sync HTTP client for a container request, reusing the caller's when usable."""
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
return _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)})
|
||||
return client
|
||||
|
||||
|
||||
def _async_http_client(
|
||||
client: HTTPHandler | AsyncHTTPHandler | None,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
) -> AsyncHTTPHandler:
|
||||
"""The async HTTP client for a container request, reusing the caller's when usable."""
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
return get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders.OPENAI,
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
)
|
||||
return client
|
||||
|
||||
|
||||
class GenericContainerHandler:
|
||||
"""
|
||||
Generic handler for container file API endpoints.
|
||||
|
|
@ -192,13 +272,13 @@ class GenericContainerHandler:
|
|||
container_provider_config: "BaseContainerConfig",
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout = 600,
|
||||
_is_async: bool = False,
|
||||
client: HTTPHandler | AsyncHTTPHandler | None = None,
|
||||
**kwargs,
|
||||
) -> Any | Coroutine[Any, Any, Any]:
|
||||
**kwargs: object,
|
||||
) -> Any | Coroutine[object, object, Any]:
|
||||
"""
|
||||
Generic handler for any container file endpoint.
|
||||
|
||||
|
|
@ -245,11 +325,11 @@ class GenericContainerHandler:
|
|||
container_provider_config: "BaseContainerConfig",
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout = 600,
|
||||
client: HTTPHandler | AsyncHTTPHandler | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
) -> Any:
|
||||
"""Synchronous request handler."""
|
||||
endpoint_config: Final = _get_endpoint_config(endpoint_name)
|
||||
|
|
@ -257,23 +337,14 @@ class GenericContainerHandler:
|
|||
raise ValueError(f"Unknown endpoint: {endpoint_name}")
|
||||
|
||||
# Get HTTP client
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
http_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)})
|
||||
else:
|
||||
http_client = client
|
||||
http_client: Final = _sync_http_client(client, litellm_params)
|
||||
|
||||
# Build request
|
||||
headers = container_provider_config.validate_environment(
|
||||
headers=extra_headers or {},
|
||||
api_key=litellm_params.get("api_key", None),
|
||||
)
|
||||
headers = _request_headers(container_provider_config, extra_headers, litellm_params)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base: Final = container_provider_config.get_complete_url(
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
api_base: Final = _request_api_base(container_provider_config, litellm_params)
|
||||
|
||||
# Build URL with path params
|
||||
path_params: Final = {p: kwargs.get(p, "") for p in endpoint_config.get("path_params", [])}
|
||||
|
|
@ -334,11 +405,11 @@ class GenericContainerHandler:
|
|||
container_provider_config: "BaseContainerConfig",
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
extra_headers: dict[str, Any] | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_headers: dict[str, object] | None = None,
|
||||
extra_query: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout = 600,
|
||||
client: HTTPHandler | AsyncHTTPHandler | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
) -> Any:
|
||||
"""Asynchronous request handler."""
|
||||
endpoint_config: Final = _get_endpoint_config(endpoint_name)
|
||||
|
|
@ -346,26 +417,14 @@ class GenericContainerHandler:
|
|||
raise ValueError(f"Unknown endpoint: {endpoint_name}")
|
||||
|
||||
# Get HTTP client
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
http_client = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders.OPENAI,
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
)
|
||||
else:
|
||||
http_client = client
|
||||
http_client: Final = _async_http_client(client, litellm_params)
|
||||
|
||||
# Build request
|
||||
headers = container_provider_config.validate_environment(
|
||||
headers=extra_headers or {},
|
||||
api_key=litellm_params.get("api_key", None),
|
||||
)
|
||||
headers = _request_headers(container_provider_config, extra_headers, litellm_params)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
api_base: Final = container_provider_config.get_complete_url(
|
||||
api_base=litellm_params.get("api_base", None),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
api_base: Final = _request_api_base(container_provider_config, litellm_params)
|
||||
|
||||
# Build URL with path params
|
||||
path_params: Final = {p: kwargs.get(p, "") for p in endpoint_config.get("path_params", [])}
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ 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, Optional, TypeAlias, TypedDict
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, Optional, TypeAlias, TypedDict
|
||||
|
||||
import certifi
|
||||
import httpx
|
||||
|
|
@ -447,7 +447,7 @@ def _safe_read_response(response: httpx.Response, timeout: float | None = None)
|
|||
return b""
|
||||
|
||||
|
||||
def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> None:
|
||||
def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> NoReturn:
|
||||
"""Raise a MaskedHTTPStatusError for sync HTTP handlers."""
|
||||
if stream:
|
||||
try:
|
||||
|
|
@ -467,7 +467,7 @@ def _raise_masked_sync_error(e: httpx.HTTPStatusError, stream: bool) -> None:
|
|||
raise MaskedHTTPStatusError(e, message=_text, text=_text) from None
|
||||
|
||||
|
||||
async def _raise_masked_async_error(e: httpx.HTTPStatusError, stream: bool) -> None:
|
||||
async def _raise_masked_async_error(e: httpx.HTTPStatusError, stream: bool) -> NoReturn:
|
||||
"""Raise a MaskedHTTPStatusError for async HTTP handlers."""
|
||||
if stream:
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ Talks to e2b's REST API directly over httpx (no e2b SDK dependency):
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Final, cast
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -68,13 +68,10 @@ class E2BSandboxConfig(BaseSandboxConfig):
|
|||
if metadata:
|
||||
body["metadata"] = metadata
|
||||
|
||||
response: Final = cast(
|
||||
httpx.Response,
|
||||
await self._http(client).post(
|
||||
url=f"{base}/sandboxes",
|
||||
headers={"X-API-Key": key, "Content-Type": "application/json"},
|
||||
json=body,
|
||||
),
|
||||
response: Final = await self._http(client).post(
|
||||
url=f"{base}/sandboxes",
|
||||
headers={"X-API-Key": key, "Content-Type": "application/json"},
|
||||
json=body,
|
||||
)
|
||||
data: Final = response.json()
|
||||
|
||||
|
|
@ -117,14 +114,11 @@ class E2BSandboxConfig(BaseSandboxConfig):
|
|||
headers["E2B-Traffic-Access-Token"] = traffic_token
|
||||
|
||||
url: Final = f"https://{JUPYTER_PORT}-{handle.id}.{handle.domain}/execute"
|
||||
response: Final = cast(
|
||||
httpx.Response,
|
||||
await self._http(client).post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
json={"code": code, "context_id": None, "env_vars": env_vars},
|
||||
stream=True,
|
||||
),
|
||||
response: Final = await self._http(client).post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
json={"code": code, "context_id": None, "env_vars": env_vars},
|
||||
stream=True,
|
||||
)
|
||||
lines: Final = await self._read_capped_lines(response)
|
||||
return self._parse_lines(lines)
|
||||
|
|
@ -142,12 +136,9 @@ class E2BSandboxConfig(BaseSandboxConfig):
|
|||
key: Final = api_key or handle._hidden_params.get("api_key") or self.validate_environment()
|
||||
base: Final = api_base or handle._hidden_params.get("api_base") or E2B_API_BASE
|
||||
try:
|
||||
response: Final = cast(
|
||||
httpx.Response,
|
||||
await self._http(client).delete(
|
||||
url=f"{base}/sandboxes/{handle.id}",
|
||||
headers={"X-API-Key": key},
|
||||
),
|
||||
response: Final = await self._http(client).delete(
|
||||
url=f"{base}/sandboxes/{handle.id}",
|
||||
headers={"X-API-Key": key},
|
||||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
if e.response.status_code == 404:
|
||||
|
|
|
|||
|
|
@ -6,11 +6,25 @@ import json
|
|||
import os
|
||||
import re
|
||||
import threading
|
||||
from typing import Any, Final
|
||||
from collections.abc import Callable
|
||||
from typing import Any, Final, Protocol
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import litellm
|
||||
from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
class _GDCHAudienceCredentials(Protocol):
|
||||
"""A GDCH service account credential already bound to an audience, ready to mint a bearer token."""
|
||||
|
||||
@property
|
||||
def valid(self) -> bool: ...
|
||||
|
||||
@property
|
||||
def token(self) -> str: ...
|
||||
|
||||
def refresh(self, request: object) -> None: ...
|
||||
|
||||
|
||||
class GDCGeminiConfig(OpenAILikeChatConfig):
|
||||
|
|
@ -21,7 +35,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
|
|||
def __init__(self, **kwargs: Any) -> None:
|
||||
super().__init__(**kwargs)
|
||||
self._creds_lock = threading.Lock()
|
||||
self._gdch_creds_cache: dict = {}
|
||||
self._gdch_creds_cache: dict[tuple[str, str], _GDCHAudienceCredentials] = {}
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
return [
|
||||
|
|
@ -110,7 +124,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
|
|||
|
||||
return f"{api_base}/v1/projects/{project}/locations/{location}/chat/completions"
|
||||
|
||||
def _read_env_bool(self, val: Any, env_var: str, default: bool = True) -> bool | str:
|
||||
def _read_env_bool(self, val: bool | str | None, env_var: str, default: bool = True) -> bool | str:
|
||||
def _parse(s: str) -> bool | str:
|
||||
cleaned: Final = s.strip().lower()
|
||||
if cleaned in ("false", "0", "no", "off"):
|
||||
|
|
@ -129,7 +143,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
|
|||
return default
|
||||
return _parse(_env_val)
|
||||
|
||||
def _fetch_auth(self, gdch_creds: Any, ssl_verify: bool | str) -> None:
|
||||
def _fetch_auth(self, gdch_creds: _GDCHAudienceCredentials, ssl_verify: bool | str) -> None:
|
||||
import requests
|
||||
from google.auth.transport import requests as auth_requests
|
||||
|
||||
|
|
@ -138,13 +152,24 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
|
|||
auth_request: Final = auth_requests.Request(session=auth_session)
|
||||
gdch_creds.refresh(auth_request)
|
||||
|
||||
def _cached_fetch_token(self, creds: Any, audience: str, ssl_verify: bool | str, api_key: str | None = None) -> str:
|
||||
def _with_gdch_audience(self, creds: object, audience: str) -> _GDCHAudienceCredentials:
|
||||
"""The credential rebound to ``audience``, which GDCH requires before a token refresh."""
|
||||
bind_audience: Final[Callable[[str], _GDCHAudienceCredentials] | None] = getattr(
|
||||
creds, "with_gdch_audience", None
|
||||
)
|
||||
if bind_audience is None:
|
||||
raise AttributeError("GDC credentials must expose with_gdch_audience to be bound to a request audience")
|
||||
return bind_audience(audience)
|
||||
|
||||
def _cached_fetch_token(
|
||||
self, creds: object, audience: str, ssl_verify: bool | str, api_key: str | None = None
|
||||
) -> str:
|
||||
# Key cache by both audience and credential identity to prevent cross-caller contamination
|
||||
cache_key: Final = (audience.rstrip("/"), api_key or str(id(creds)))
|
||||
|
||||
with self._creds_lock:
|
||||
if cache_key not in self._gdch_creds_cache:
|
||||
self._gdch_creds_cache[cache_key] = creds.with_gdch_audience(audience.rstrip("/"))
|
||||
self._gdch_creds_cache[cache_key] = self._with_gdch_audience(creds, audience.rstrip("/"))
|
||||
|
||||
gdch_creds: Final = self._gdch_creds_cache[cache_key]
|
||||
|
||||
|
|
@ -155,7 +180,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
|
|||
|
||||
return token
|
||||
|
||||
def _load_creds_from_key(self, api_key: str) -> tuple[Any, bool]:
|
||||
def _load_creds_from_key(self, api_key: str) -> tuple[object | None, bool]:
|
||||
import google.auth
|
||||
|
||||
try:
|
||||
|
|
@ -175,7 +200,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
|
|||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: list[Any],
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: str | None = None,
|
||||
|
|
@ -230,7 +255,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
|
|||
if self._read_env_bool(litellm_params.get("gdc_token_caching"), "GDC_TOKEN_CACHING", default=False):
|
||||
token = self._cached_fetch_token(creds, audience, ssl_verify, api_key)
|
||||
else:
|
||||
gdch_creds: Final = creds.with_gdch_audience(audience)
|
||||
gdch_creds: Final = self._with_gdch_audience(creds, audience)
|
||||
self._fetch_auth(gdch_creds, ssl_verify)
|
||||
token = gdch_creds.token
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
|
|
@ -252,7 +277,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig):
|
|||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: list[Any],
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -4,10 +4,11 @@ Transformation logic from Cohere's /v1/rerank format to Infinity's `/v1/rerank`
|
|||
Why separate file? Make it easy to see how transformation works
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -26,6 +27,31 @@ from litellm.types.rerank import (
|
|||
from ..common_utils import InfinityError
|
||||
|
||||
|
||||
class _InfinityRerankUsage(TypedDict, extra_items=ReadOnly[int]):
|
||||
"""The token counters Infinity reports in the ``usage`` block of a rerank response."""
|
||||
|
||||
|
||||
class _InfinityRerankResult(TypedDict):
|
||||
"""One scored document in an Infinity ``/v1/rerank`` response."""
|
||||
|
||||
index: ReadOnly[int]
|
||||
relevance_score: ReadOnly[float]
|
||||
document: ReadOnly[str]
|
||||
|
||||
|
||||
class _InfinityRerankResponse(TypedDict):
|
||||
"""The JSON body returned by Infinity's ``/v1/rerank`` endpoint."""
|
||||
|
||||
id: ReadOnly[NotRequired[str]]
|
||||
usage: ReadOnly[NotRequired[_InfinityRerankUsage]]
|
||||
results: ReadOnly[Sequence[_InfinityRerankResult]]
|
||||
|
||||
|
||||
def _parse_rerank_response(raw_response: httpx.Response) -> _InfinityRerankResponse:
|
||||
"""Read the untyped JSON body of an Infinity rerank response."""
|
||||
return raw_response.json()
|
||||
|
||||
|
||||
class InfinityRerankConfig(CohereRerankConfig):
|
||||
def get_complete_url(
|
||||
self,
|
||||
|
|
@ -82,7 +108,7 @@ class InfinityRerankConfig(CohereRerankConfig):
|
|||
No transformation required, Infinity follows Cohere API response format
|
||||
"""
|
||||
try:
|
||||
raw_response_json: Final = raw_response.json()
|
||||
raw_response_json: Final = _parse_rerank_response(raw_response)
|
||||
except Exception:
|
||||
raise InfinityError(message=raw_response.text, status_code=raw_response.status_code)
|
||||
|
||||
|
|
|
|||
|
|
@ -13,12 +13,57 @@ Generated files are returned directly in the response - no separate storage need
|
|||
|
||||
import base64
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from enum import Enum
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, Protocol
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
|
||||
class _ToolCallFunction(Protocol):
|
||||
"""Function payload of an assistant tool call."""
|
||||
|
||||
name: str | None
|
||||
arguments: str
|
||||
|
||||
|
||||
class _ToolCall(Protocol):
|
||||
"""Tool call requested by the assistant on a chat completion choice."""
|
||||
|
||||
id: str
|
||||
function: _ToolCallFunction
|
||||
|
||||
|
||||
class _AssistantMessage(Protocol):
|
||||
"""Assistant message carried by a chat completion choice."""
|
||||
|
||||
content: str | None
|
||||
tool_calls: Sequence[_ToolCall] | None
|
||||
|
||||
|
||||
class _CompletionChoice(Protocol):
|
||||
"""Single choice of a chat completion response."""
|
||||
|
||||
finish_reason: str
|
||||
message: _AssistantMessage
|
||||
|
||||
|
||||
class _SandboxFile(TypedDict):
|
||||
"""File generated inside the sandbox during a code execution run."""
|
||||
|
||||
name: ReadOnly[str]
|
||||
mime_type: ReadOnly[str]
|
||||
content_base64: ReadOnly[str]
|
||||
|
||||
|
||||
class _CodeExecutionArguments(TypedDict):
|
||||
"""Arguments the model passes to the `litellm_code_execution` tool."""
|
||||
|
||||
code: NotRequired[ReadOnly[str]]
|
||||
|
||||
|
||||
class LiteLLMInternalTools(str, Enum):
|
||||
"""
|
||||
Enum for internal LiteLLM tools that are injected into requests.
|
||||
|
|
@ -30,7 +75,7 @@ class LiteLLMInternalTools(str, Enum):
|
|||
CODE_EXECUTION = "litellm_code_execution"
|
||||
|
||||
|
||||
def get_litellm_code_execution_tool() -> dict[str, Any]:
|
||||
def get_litellm_code_execution_tool() -> dict[str, object]:
|
||||
"""
|
||||
Returns the litellm_code_execution tool definition in OpenAI format.
|
||||
|
||||
|
|
@ -51,7 +96,7 @@ def get_litellm_code_execution_tool() -> dict[str, Any]:
|
|||
}
|
||||
|
||||
|
||||
def get_litellm_code_execution_tool_anthropic() -> dict[str, Any]:
|
||||
def get_litellm_code_execution_tool_anthropic() -> dict[str, object]:
|
||||
"""
|
||||
Returns the litellm_code_execution tool definition in Anthropic/messages API format.
|
||||
|
||||
|
|
@ -103,7 +148,7 @@ class CodeExecutionHandler:
|
|||
skill_files: dict[str, bytes],
|
||||
skill_id: str | None = None,
|
||||
**kwargs,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Execute an LLM call with automatic code execution handling.
|
||||
|
||||
|
|
@ -134,8 +179,8 @@ class CodeExecutionHandler:
|
|||
)
|
||||
|
||||
current_messages: Final = list(messages)
|
||||
generated_files: Final[list[dict[str, Any]]] = [] # Files returned directly
|
||||
execution_results: Final[list[dict]] = []
|
||||
generated_files: Final[list[dict[str, object]]] = [] # Files returned directly
|
||||
execution_results: Final[list[dict[str, object]]] = []
|
||||
|
||||
executor: Final = SkillsSandboxExecutor(timeout=self.sandbox_timeout)
|
||||
response: Any = None # Initialize to avoid possibly unbound error
|
||||
|
|
@ -151,11 +196,12 @@ class CodeExecutionHandler:
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
assistant_message = response.choices[0].message
|
||||
stop_reason = response.choices[0].finish_reason
|
||||
choice: _CompletionChoice = response.choices[0]
|
||||
assistant_message = choice.message
|
||||
stop_reason: str = choice.finish_reason
|
||||
|
||||
# Build assistant message for conversation history
|
||||
assistant_msg_dict: dict[str, Any] = {
|
||||
assistant_msg_dict: dict[str, object] = {
|
||||
"role": "assistant",
|
||||
"content": assistant_message.content,
|
||||
}
|
||||
|
|
@ -190,8 +236,8 @@ class CodeExecutionHandler:
|
|||
if tool_name == LiteLLMInternalTools.CODE_EXECUTION.value:
|
||||
# Execute code in sandbox
|
||||
try:
|
||||
args = json.loads(tool_call.function.arguments)
|
||||
code = args.get("code", "")
|
||||
args: _CodeExecutionArguments = json.loads(tool_call.function.arguments)
|
||||
code: str = args.get("code", "")
|
||||
|
||||
verbose_logger.debug("CodeExecutionHandler: Executing code (%s chars)", len(code))
|
||||
|
||||
|
|
@ -202,13 +248,15 @@ class CodeExecutionHandler:
|
|||
|
||||
verbose_logger.debug("CodeExecutionHandler: Execution result: %s", exec_result)
|
||||
|
||||
sandbox_files: Sequence[_SandboxFile] = exec_result["files"]
|
||||
|
||||
execution_results.append(
|
||||
{
|
||||
"iteration": iteration,
|
||||
"success": exec_result["success"],
|
||||
"output": exec_result["output"],
|
||||
"error": exec_result["error"],
|
||||
"files": [f["name"] for f in exec_result["files"]],
|
||||
"files": [f["name"] for f in sandbox_files],
|
||||
}
|
||||
)
|
||||
|
||||
|
|
@ -216,9 +264,9 @@ class CodeExecutionHandler:
|
|||
tool_result = exec_result["output"] or ""
|
||||
|
||||
# Collect generated files (returned directly, no storage)
|
||||
if exec_result["files"]:
|
||||
if sandbox_files:
|
||||
tool_result += "\n\nGenerated files:"
|
||||
for f in exec_result["files"]:
|
||||
for f in sandbox_files:
|
||||
file_content = base64.b64decode(f["content_base64"])
|
||||
# Add to generated files list (returned in response)
|
||||
generated_files.append(
|
||||
|
|
|
|||
|
|
@ -4,16 +4,32 @@ Ollama /chat/completion calls handled in llm_http_handler.py
|
|||
[TODO]: migrate embeddings to a base handler as well.
|
||||
"""
|
||||
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, Protocol, TypedDict
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
|
||||
class TokenEncoder(Protocol):
|
||||
"""The tokenizer surface used to estimate prompt tokens."""
|
||||
|
||||
def encode(self, text: str, /) -> Sequence[int]: ...
|
||||
|
||||
|
||||
class OllamaEmbeddingResponse(TypedDict):
|
||||
"""Body of an Ollama ``/api/embed`` response."""
|
||||
|
||||
embeddings: ReadOnly[list[list[float]]]
|
||||
prompt_eval_count: ReadOnly[NotRequired[int]]
|
||||
|
||||
|
||||
def _prepare_ollama_embedding_payload(
|
||||
model: str, prompts: list[str], optional_params: dict[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
data: Final[dict[str, Any]] = {"model": model, "input": prompts}
|
||||
model: str, prompts: list[str], optional_params: Mapping[str, object]
|
||||
) -> dict[str, object]:
|
||||
data: Final[dict[str, object]] = {"model": model, "input": prompts}
|
||||
special_optional_params: Final = ["truncate", "options", "keep_alive", "dimensions"]
|
||||
|
||||
for k, v in optional_params.items():
|
||||
|
|
@ -27,12 +43,12 @@ def _prepare_ollama_embedding_payload(
|
|||
|
||||
|
||||
def _process_ollama_embedding_response(
|
||||
response_json: dict,
|
||||
response_json: OllamaEmbeddingResponse,
|
||||
prompts: list[str],
|
||||
model: str,
|
||||
model_response: EmbeddingResponse,
|
||||
logging_obj: Any,
|
||||
encoding: Any,
|
||||
encoding: TokenEncoder | None,
|
||||
) -> EmbeddingResponse:
|
||||
output_data: Final = []
|
||||
embeddings: Final[list[list[float]]] = response_json["embeddings"]
|
||||
|
|
@ -72,7 +88,7 @@ async def ollama_aembeddings(
|
|||
model_response: EmbeddingResponse,
|
||||
optional_params: dict,
|
||||
logging_obj: Any,
|
||||
encoding: Any,
|
||||
encoding: TokenEncoder | None,
|
||||
):
|
||||
if not api_base.endswith("/api/embed"):
|
||||
api_base += "/api/embed"
|
||||
|
|
@ -80,7 +96,7 @@ async def ollama_aembeddings(
|
|||
data: Final = _prepare_ollama_embedding_payload(model, prompts, optional_params)
|
||||
|
||||
response: Final = await litellm.module_level_aclient.post(url=api_base, json=data)
|
||||
response_json: Final = response.json()
|
||||
response_json: Final[OllamaEmbeddingResponse] = response.json()
|
||||
|
||||
return _process_ollama_embedding_response(
|
||||
response_json=response_json,
|
||||
|
|
@ -99,7 +115,7 @@ def ollama_embeddings(
|
|||
optional_params: dict,
|
||||
model_response: EmbeddingResponse,
|
||||
logging_obj: Any,
|
||||
encoding: Any = None,
|
||||
encoding: TokenEncoder | None = None,
|
||||
):
|
||||
if not api_base.endswith("/api/embed"):
|
||||
api_base += "/api/embed"
|
||||
|
|
@ -107,7 +123,7 @@ def ollama_embeddings(
|
|||
data: Final = _prepare_ollama_embedding_payload(model, prompts, optional_params)
|
||||
|
||||
response: Final = litellm.module_level_client.post(url=api_base, json=data)
|
||||
response_json: Final = response.json()
|
||||
response_json: Final[OllamaEmbeddingResponse] = response.json()
|
||||
|
||||
return _process_ollama_embedding_response(
|
||||
response_json=response_json,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
import httpx
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
|
||||
|
|
@ -11,9 +13,11 @@ from litellm.secret_managers.main import get_secret_str
|
|||
from litellm.types.containers.main import (
|
||||
ContainerCreateOptionalRequestParams,
|
||||
ContainerFileListResponse,
|
||||
ContainerFileObject,
|
||||
ContainerListResponse,
|
||||
ContainerObject,
|
||||
DeleteContainerResult,
|
||||
ExpiresAfter,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
|
@ -32,6 +36,46 @@ else:
|
|||
BaseLLMException = Any
|
||||
|
||||
|
||||
class OpenAIContainerPayload(TypedDict):
|
||||
"""The JSON body OpenAI returns for a single container."""
|
||||
|
||||
id: ReadOnly[str]
|
||||
object: ReadOnly[Literal["container"]]
|
||||
created_at: ReadOnly[int]
|
||||
status: ReadOnly[str]
|
||||
expires_after: ReadOnly[ExpiresAfter | None]
|
||||
last_active_at: ReadOnly[int | None]
|
||||
name: ReadOnly[str | None]
|
||||
|
||||
|
||||
class OpenAIContainerListPayload(TypedDict):
|
||||
"""The JSON body OpenAI returns for a page of containers."""
|
||||
|
||||
object: ReadOnly[Literal["list"]]
|
||||
data: ReadOnly[list[ContainerObject]]
|
||||
first_id: ReadOnly[str | None]
|
||||
last_id: ReadOnly[str | None]
|
||||
has_more: ReadOnly[bool]
|
||||
|
||||
|
||||
class OpenAIContainerDeletedPayload(TypedDict):
|
||||
"""The JSON body OpenAI returns for a deleted container."""
|
||||
|
||||
id: ReadOnly[str]
|
||||
object: ReadOnly[Literal["container.deleted"]]
|
||||
deleted: ReadOnly[bool]
|
||||
|
||||
|
||||
class OpenAIContainerFileListPayload(TypedDict):
|
||||
"""The JSON body OpenAI returns for a page of container files."""
|
||||
|
||||
object: ReadOnly[Literal["list"]]
|
||||
data: ReadOnly[list[ContainerFileObject]]
|
||||
first_id: ReadOnly[str | None]
|
||||
last_id: ReadOnly[str | None]
|
||||
has_more: ReadOnly[bool]
|
||||
|
||||
|
||||
class OpenAIContainerConfig(BaseContainerConfig):
|
||||
"""Configuration class for OpenAI container API."""
|
||||
|
||||
|
|
@ -87,7 +131,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
def transform_container_create_request(
|
||||
self,
|
||||
name: str,
|
||||
container_create_optional_request_params: dict,
|
||||
container_create_optional_request_params: Mapping[str, object],
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
|
|
@ -111,7 +155,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ContainerObject:
|
||||
"""Transform the OpenAI container creation response."""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final[OpenAIContainerPayload] = raw_response.json()
|
||||
|
||||
# Transform the response data
|
||||
container_obj: Final = ContainerObject(**response_data)
|
||||
|
|
@ -140,7 +184,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
order: str | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_query: Mapping[str, object] | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform the container list request for OpenAI API.
|
||||
|
||||
|
|
@ -151,7 +195,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
url: Final = api_base
|
||||
|
||||
# Prepare query parameters
|
||||
params: Final = {}
|
||||
params: Final[dict[str, object]] = {}
|
||||
if after is not None:
|
||||
params["after"] = after
|
||||
if limit is not None:
|
||||
|
|
@ -171,7 +215,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ContainerListResponse:
|
||||
"""Transform the OpenAI container list response."""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final[OpenAIContainerListPayload] = raw_response.json()
|
||||
|
||||
# Transform the response data
|
||||
container_list: Final = ContainerListResponse(**response_data)
|
||||
|
|
@ -191,7 +235,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}")
|
||||
|
||||
# No additional data needed for GET request
|
||||
data: Final[dict[str, Any]] = {}
|
||||
data: Final[dict[str, object]] = {}
|
||||
|
||||
return url, data
|
||||
|
||||
|
|
@ -201,7 +245,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ContainerObject:
|
||||
"""Transform the OpenAI container retrieve response."""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final[OpenAIContainerPayload] = raw_response.json()
|
||||
# Transform the response data
|
||||
container_obj: Final = ContainerObject(**response_data)
|
||||
|
||||
|
|
@ -224,7 +268,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}")
|
||||
|
||||
# No data needed for DELETE request
|
||||
data: Final[dict[str, Any]] = {}
|
||||
data: Final[dict[str, object]] = {}
|
||||
|
||||
return url, data
|
||||
|
||||
|
|
@ -234,7 +278,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> DeleteContainerResult:
|
||||
"""Transform the OpenAI container delete response."""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final[OpenAIContainerDeletedPayload] = raw_response.json()
|
||||
|
||||
# Transform the response data
|
||||
delete_result: Final = DeleteContainerResult(**response_data)
|
||||
|
|
@ -250,7 +294,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
order: str | None = None,
|
||||
extra_query: dict[str, Any] | None = None,
|
||||
extra_query: Mapping[str, object] | None = None,
|
||||
) -> tuple[str, dict]:
|
||||
"""Transform the container file list request for OpenAI API.
|
||||
|
||||
|
|
@ -262,7 +306,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}/files")
|
||||
|
||||
# Prepare query parameters
|
||||
params: Final[dict[str, Any]] = {}
|
||||
params: Final[dict[str, object]] = {}
|
||||
if after is not None:
|
||||
params["after"] = after
|
||||
if limit is not None:
|
||||
|
|
@ -282,7 +326,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
logging_obj: LiteLLMLoggingObj,
|
||||
) -> ContainerFileListResponse:
|
||||
"""Transform the OpenAI container file list response."""
|
||||
response_data: Final = raw_response.json()
|
||||
response_data: Final[OpenAIContainerFileListPayload] = raw_response.json()
|
||||
|
||||
# Transform the response data
|
||||
file_list: Final = ContainerFileListResponse(**response_data)
|
||||
|
|
@ -308,7 +352,7 @@ class OpenAIContainerConfig(BaseContainerConfig):
|
|||
url: Final = join_container_api_base_path(api_base, f"/{encoded_container_id}/files/{encoded_file_id}/content")
|
||||
|
||||
# No query parameters needed
|
||||
params: Final[dict[str, Any]] = {}
|
||||
params: Final[dict[str, object]] = {}
|
||||
|
||||
return url, params
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from typing import Final, cast
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -86,13 +86,10 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
|
|||
secure_access=secure_access,
|
||||
)
|
||||
|
||||
response: Final = cast(
|
||||
httpx.Response,
|
||||
await self._http(client).post(
|
||||
url=f"{base}/sandboxes",
|
||||
headers=self._lifecycle_headers(key),
|
||||
json=body,
|
||||
),
|
||||
response: Final = await self._http(client).post(
|
||||
url=f"{base}/sandboxes",
|
||||
headers=self._lifecycle_headers(key),
|
||||
json=body,
|
||||
)
|
||||
data: Final = response.json()
|
||||
sandbox_id: Final = str(data["id"])
|
||||
|
|
@ -182,12 +179,9 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
|
|||
base: Final = str(handle._hidden_params.get("api_base") or self._api_base(api_base))
|
||||
key: Final = self._api_key(api_key=api_key, handle=handle)
|
||||
try:
|
||||
response: Final = cast(
|
||||
httpx.Response,
|
||||
await self._http(client).delete(
|
||||
url=f"{base}/sandboxes/{handle.id}",
|
||||
headers=self._lifecycle_headers(key),
|
||||
),
|
||||
response: Final = await self._http(client).delete(
|
||||
url=f"{base}/sandboxes/{handle.id}",
|
||||
headers=self._lifecycle_headers(key),
|
||||
)
|
||||
except httpx.HTTPStatusError as e:
|
||||
if e.response.status_code == 404:
|
||||
|
|
@ -245,12 +239,9 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
|
|||
) -> None:
|
||||
deadline: Final = time.monotonic() + ready_timeout
|
||||
while True:
|
||||
response = cast(
|
||||
httpx.Response,
|
||||
await self._http(client).get(
|
||||
url=f"{api_base}/sandboxes/{sandbox_id}",
|
||||
headers=headers,
|
||||
),
|
||||
response = await self._http(client).get(
|
||||
url=f"{api_base}/sandboxes/{sandbox_id}",
|
||||
headers=headers,
|
||||
)
|
||||
data = response.json()
|
||||
state = self._sandbox_state(data)
|
||||
|
|
@ -306,13 +297,10 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
|
|||
use_server_proxy: bool,
|
||||
client: AsyncHTTPHandler | None,
|
||||
) -> tuple[str, dict[str, str]]:
|
||||
response: Final = cast(
|
||||
httpx.Response,
|
||||
await self._http(client).get(
|
||||
url=f"{api_base}/sandboxes/{sandbox_id}/endpoints/{OPEN_SANDBOX_EXECD_PORT}",
|
||||
headers=headers,
|
||||
params={"use_server_proxy": use_server_proxy},
|
||||
),
|
||||
response: Final = await self._http(client).get(
|
||||
url=f"{api_base}/sandboxes/{sandbox_id}/endpoints/{OPEN_SANDBOX_EXECD_PORT}",
|
||||
headers=headers,
|
||||
params={"use_server_proxy": use_server_proxy},
|
||||
)
|
||||
data: Final = response.json()
|
||||
endpoint: Final = data.get("endpoint")
|
||||
|
|
@ -329,15 +317,12 @@ class OpenSandboxSandboxConfig(BaseSandboxConfig):
|
|||
client: AsyncHTTPHandler | None,
|
||||
) -> list[str]:
|
||||
timeout: Final = httpx.Timeout(connect=30.0, read=None, write=30.0, pool=None)
|
||||
response: Final = cast(
|
||||
httpx.Response,
|
||||
await self._http(client).post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
json=body,
|
||||
stream=True,
|
||||
),
|
||||
response: Final = await self._http(client).post(
|
||||
url=url,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
json=body,
|
||||
stream=True,
|
||||
)
|
||||
return await self._read_capped_lines(response)
|
||||
|
||||
|
|
|
|||
|
|
@ -97,7 +97,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
|
|||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def get_auth_credentials(self, litellm_params: dict) -> BaseVectorStoreAuthCredentials:
|
||||
def get_auth_credentials(self, litellm_params: Mapping[str, object]) -> BaseVectorStoreAuthCredentials:
|
||||
# Get credentials and project info
|
||||
vertex_credentials: Final = self.get_vertex_ai_credentials(dict(litellm_params))
|
||||
vertex_project: Final = self.get_vertex_ai_project(dict(litellm_params))
|
||||
|
|
@ -122,7 +122,9 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
|
|||
"write": [("POST", "/ragCorpora")],
|
||||
}
|
||||
|
||||
def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict:
|
||||
def validate_environment(
|
||||
self, headers: dict[str, str], litellm_params: GenericLiteLLMParams | None
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
Validate and set up authentication for Vertex AI RAG API
|
||||
"""
|
||||
|
|
@ -135,7 +137,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
|
|||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
litellm_params: dict,
|
||||
litellm_params: dict[str, object],
|
||||
) -> str:
|
||||
"""
|
||||
Get the Base endpoint for Vertex AI RAG API
|
||||
|
|
|
|||
|
|
@ -33,6 +33,6 @@ class DomainModel(BaseModel):
|
|||
return cls(**record.dict())
|
||||
return cls(**dict(record))
|
||||
|
||||
def to_db_dict(self, exclude_unset: bool = False) -> dict[str, Any]:
|
||||
def to_db_dict(self, exclude_unset: bool = False) -> dict[str, object]:
|
||||
"""Convert domain model to a dictionary for database operations."""
|
||||
return self.model_dump(exclude_none=True, exclude_unset=exclude_unset)
|
||||
|
|
|
|||
|
|
@ -1096,8 +1096,7 @@ async def exchange_token_with_server(
|
|||
headers={"Accept": "application/json", **token_request.headers},
|
||||
data=token_data,
|
||||
)
|
||||
if response is not None:
|
||||
response.raise_for_status()
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
fault: Final = classify_upstream_token_rejection(
|
||||
exc.response,
|
||||
|
|
@ -1119,11 +1118,6 @@ async def exchange_token_with_server(
|
|||
)
|
||||
return _bridge_mint_error_response("invalid_refresh")
|
||||
return render_token_fault(fault)
|
||||
if response is None:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail="MCP upstream token endpoint returned no response",
|
||||
)
|
||||
token_response = response.json()
|
||||
|
||||
# Validate token response against server-configured rules before any storage.
|
||||
|
|
@ -1536,16 +1530,10 @@ async def _post_dcr_registration(
|
|||
headers=headers,
|
||||
json=register_data,
|
||||
)
|
||||
if response is not None:
|
||||
response.raise_for_status()
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as exc:
|
||||
status_code, detail = dcr_fault_detail(classify_upstream_dcr_rejection(exc.response, log_context=server_id))
|
||||
raise HTTPException(status_code=status_code, detail=detail) from exc
|
||||
if response is None:
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail="MCP upstream registration endpoint returned no response",
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -36,7 +36,10 @@ a healed fleet has no null rows and the backfill exits after one query.
|
|||
|
||||
import json
|
||||
from collections import Counter
|
||||
from typing import Any, Final, Literal
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final, Literal, Protocol
|
||||
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._experimental.mcp_server.db import _decode_oauth_payload, decrypt_credentials
|
||||
|
|
@ -55,9 +58,59 @@ BackfillRule = Literal[
|
|||
_BACKFILL_AUDIT_ACTOR: Final = "oauth2_flow_backfill"
|
||||
|
||||
|
||||
def _decrypted_credentials(raw_credentials: Any) -> MCPCredentials | None:
|
||||
class _MCPServerRow(Protocol):
|
||||
"""The ``LiteLLM_MCPServerTable`` columns this backfill reads."""
|
||||
|
||||
@property
|
||||
def server_id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def authorization_url(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def registration_url(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def token_url(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def credentials(self) -> str | Mapping[str, JsonValue] | None: ...
|
||||
|
||||
|
||||
class _MCPUserCredentialRow(Protocol):
|
||||
"""The ``LiteLLM_MCPUserCredentials`` columns this backfill reads."""
|
||||
|
||||
@property
|
||||
def server_id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def credential_b64(self) -> str: ...
|
||||
|
||||
|
||||
class _MCPServerTable(Protocol):
|
||||
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_MCPServerRow]: ...
|
||||
|
||||
async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, str]) -> object: ...
|
||||
|
||||
|
||||
class _MCPUserCredentialsTable(Protocol):
|
||||
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_MCPUserCredentialRow]: ...
|
||||
|
||||
|
||||
def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable:
|
||||
"""The MCP server table, typed so the untyped prisma client surface stops here."""
|
||||
return prisma_client.db.litellm_mcpservertable
|
||||
|
||||
|
||||
def _mcp_user_credentials_table(prisma_client: PrismaClient) -> _MCPUserCredentialsTable:
|
||||
"""The per-user MCP credential table, typed so the untyped prisma client surface stops here."""
|
||||
return prisma_client.db.litellm_mcpusercredentials
|
||||
|
||||
|
||||
def _decrypted_credentials(raw_credentials: str | Mapping[str, JsonValue] | None) -> MCPCredentials | None:
|
||||
if raw_credentials is None:
|
||||
return None
|
||||
parsed: JsonValue | Mapping[str, JsonValue]
|
||||
if isinstance(raw_credentials, str):
|
||||
try:
|
||||
parsed = json.loads(raw_credentials)
|
||||
|
|
@ -92,14 +145,14 @@ def classify_null_flow_row(
|
|||
async def backfill_null_oauth2_flows(prisma_client: PrismaClient) -> dict[BackfillRule, int]:
|
||||
"""Classify every ``auth_type=oauth2`` row whose ``oauth2_flow`` is null; stamp the provable
|
||||
ones, warn on the ambiguous ones, and return counts per rule."""
|
||||
null_rows: Final[list[Any]] = await prisma_client.db.litellm_mcpservertable.find_many(
|
||||
null_rows: Final[Sequence[_MCPServerRow]] = await _mcp_server_table(prisma_client).find_many(
|
||||
where={"auth_type": "oauth2", "oauth2_flow": None},
|
||||
)
|
||||
if not null_rows:
|
||||
return {}
|
||||
|
||||
server_ids: Final = [row.server_id for row in null_rows]
|
||||
token_rows: Final[list[Any]] = await prisma_client.db.litellm_mcpusercredentials.find_many(
|
||||
token_rows: Final[Sequence[_MCPUserCredentialRow]] = await _mcp_user_credentials_table(prisma_client).find_many(
|
||||
where={"server_id": {"in": server_ids}},
|
||||
)
|
||||
server_ids_with_oauth_tokens: Final[set[str]] = {
|
||||
|
|
@ -141,7 +194,7 @@ async def backfill_null_oauth2_flows(prisma_client: PrismaClient) -> dict[Backfi
|
|||
stamped_flows: Final = {flow for _, (flow, _) in classified if flow is not None}
|
||||
for stamped_flow in stamped_flows:
|
||||
server_ids_for_flow = [row.server_id for row, (row_flow, _) in classified if row_flow == stamped_flow]
|
||||
await prisma_client.db.litellm_mcpservertable.update_many(
|
||||
await _mcp_server_table(prisma_client).update_many(
|
||||
where={"server_id": {"in": server_ids_for_flow}, "oauth2_flow": None},
|
||||
data={"oauth2_flow": stamped_flow, "updated_by": _BACKFILL_AUDIT_ACTOR},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -19,9 +19,9 @@ Implements the client-credentials behavior contract for the v2 resolver:
|
|||
identity.
|
||||
|
||||
The token-endpoint POST is injected (``M2MTokenEndpointPost``) so the grant orchestration is
|
||||
testable without a live IdP; ``post_client_credentials_grant`` is the httpx edge and the one
|
||||
place the untyped response boundary is contained. Failures are values: the source returns
|
||||
``Result[OAuthToken, CredError]``; only the httpx edge touches exceptions.
|
||||
testable without a live IdP; ``post_client_credentials_grant`` is the httpx edge. Failures are
|
||||
values: the source returns ``Result[OAuthToken, CredError]``; only the httpx edge touches
|
||||
exceptions.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -95,18 +95,17 @@ async def post_client_credentials_grant(
|
|||
) -> TokenEndpointOutcome:
|
||||
"""POST the grant to the token endpoint and classify the transport outcome.
|
||||
|
||||
The httpx edge: litellm's handler is partially typed (and raises ``HTTPStatusError`` itself on
|
||||
a 4xx/5xx), so the untyped boundary is contained here and every field the caller reads comes
|
||||
out of a validated ``TokenEndpointOutcome``.
|
||||
The httpx edge: litellm's handler raises ``HTTPStatusError`` itself on a 4xx/5xx, and every
|
||||
field the caller reads comes out of a validated ``TokenEndpointOutcome``.
|
||||
"""
|
||||
from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 # defer heavy handler import to call time
|
||||
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # handler is partially typed
|
||||
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # handler factory params are coarsely typed
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider # noqa: PLC0415 # deferred with the handler import
|
||||
|
||||
try:
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response = await client.post( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # handler is partially typed
|
||||
response: Final = await client.post( # pyright: ignore[reportUnknownMemberType] # handler params are coarsely typed
|
||||
url, headers={"Accept": "application/json", **headers}, data=form
|
||||
)
|
||||
except httpx.HTTPStatusError as status_err:
|
||||
|
|
@ -114,8 +113,6 @@ async def post_client_credentials_grant(
|
|||
return TokenEndpointDenied(status_code=status_code, detail=f"token endpoint returned HTTP {status_code}")
|
||||
except Exception as exc: # noqa: BLE001 # any transport failure is the same outcome: unreachable
|
||||
return TokenEndpointUnreachable(detail=str(exc))
|
||||
if not isinstance(response, httpx.Response):
|
||||
return TokenEndpointUnreachable(detail="token endpoint returned no response")
|
||||
try:
|
||||
body: Final = _TOKEN_BODY_ADAPTER.validate_json(response.content)
|
||||
except ValidationError:
|
||||
|
|
|
|||
|
|
@ -111,9 +111,6 @@ class TokenEndpointClient:
|
|||
return Error(
|
||||
CredError.of_upstream_unavailable("token exchange failed: token endpoint returned a non-JSON response")
|
||||
)
|
||||
if raw is None:
|
||||
verbose_proxy_logger.warning("MCP token endpoint %s returned no response", endpoint)
|
||||
return Error(CredError.of_upstream_unavailable("token exchange failed: no response from token endpoint"))
|
||||
try:
|
||||
parsed: Final = _TokenEndpointResponse.model_validate(raw)
|
||||
except ValidationError:
|
||||
|
|
@ -199,7 +196,7 @@ def _cache_ttl_seconds(expires_in: int | None) -> int:
|
|||
)
|
||||
|
||||
|
||||
async def _post_form(endpoint: str, data: dict[str, str]) -> object | None:
|
||||
async def _post_form(endpoint: str, data: dict[str, str]) -> object:
|
||||
# litellm's httpx handler and httpx.Response are only partially typed; the token endpoint
|
||||
# returns a JSON object that `_TokenEndpointResponse` validates, so the untyped boundary is
|
||||
# contained here. A non-2xx raises `httpx.HTTPStatusError`, an unreachable endpoint raises
|
||||
|
|
@ -208,8 +205,6 @@ async def _post_form(endpoint: str, data: dict[str, str]) -> object | None:
|
|||
# each to a CredError.
|
||||
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP) # pyright: ignore[reportUnknownVariableType] # litellm http handler is untyped
|
||||
response = await client.post(endpoint, data=data) # pyright: ignore[reportUnknownMemberType,reportUnknownVariableType] # litellm http handler is untyped
|
||||
if response is None:
|
||||
return None
|
||||
response.raise_for_status()
|
||||
return response.json() # pyright: ignore[reportAny] # untyped JSON; validated by _TokenEndpointResponse in fetch
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,9 @@
|
|||
import json
|
||||
from typing import Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Final, Protocol
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
|
|
@ -7,18 +11,73 @@ from litellm.proxy.utils import PrismaClient
|
|||
from litellm.repositories.table_repositories import MCPToolsetRepository
|
||||
from litellm.types.mcp_server.mcp_toolset import (
|
||||
MCPToolset,
|
||||
MCPToolsetTool,
|
||||
NewMCPToolsetRequest,
|
||||
UpdateMCPToolsetRequest,
|
||||
)
|
||||
|
||||
|
||||
def _toolset_from_row(row) -> MCPToolset:
|
||||
class MCPToolsetFields(TypedDict):
|
||||
"""The ``MCPToolset`` constructor keywords a toolset row expands into."""
|
||||
|
||||
toolset_id: ReadOnly[str]
|
||||
toolset_name: ReadOnly[str]
|
||||
description: NotRequired[ReadOnly[str | None]]
|
||||
tools: NotRequired[ReadOnly[list[MCPToolsetTool]]]
|
||||
created_at: NotRequired[ReadOnly[datetime | None]]
|
||||
created_by: NotRequired[ReadOnly[str | None]]
|
||||
updated_at: NotRequired[ReadOnly[datetime | None]]
|
||||
updated_by: NotRequired[ReadOnly[str | None]]
|
||||
|
||||
|
||||
class MCPToolsetRowData(TypedDict):
|
||||
"""A toolset table row, whose ``tools`` column is stored as JSON."""
|
||||
|
||||
toolset_id: ReadOnly[str]
|
||||
toolset_name: ReadOnly[str]
|
||||
description: NotRequired[ReadOnly[str | None]]
|
||||
tools: NotRequired[ReadOnly[str | list[MCPToolsetTool]]]
|
||||
created_at: NotRequired[ReadOnly[datetime | None]]
|
||||
created_by: NotRequired[ReadOnly[str | None]]
|
||||
updated_at: NotRequired[ReadOnly[datetime | None]]
|
||||
updated_by: NotRequired[ReadOnly[str | None]]
|
||||
|
||||
|
||||
class MCPToolsetRow(Protocol):
|
||||
"""A row of the toolset table, as the prisma client returns it."""
|
||||
|
||||
def model_dump(self) -> MCPToolsetRowData: ...
|
||||
|
||||
|
||||
class MCPToolsetTable(Protocol):
|
||||
"""The prisma table actions this module runs against the toolset table."""
|
||||
|
||||
async def create(self, data: Mapping[str, object]) -> MCPToolsetRow: ...
|
||||
|
||||
async def find_unique(self, where: Mapping[str, object]) -> MCPToolsetRow | None: ...
|
||||
|
||||
async def find_first(self, where: Mapping[str, object]) -> MCPToolsetRow | None: ...
|
||||
|
||||
async def find_many(self, where: Mapping[str, object]) -> Sequence[MCPToolsetRow]: ...
|
||||
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> MCPToolsetRow: ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> MCPToolsetRow: ...
|
||||
|
||||
|
||||
def _toolset_table(prisma_client: PrismaClient) -> MCPToolsetTable:
|
||||
"""The toolset table actions of the prisma client."""
|
||||
return MCPToolsetRepository(prisma_client).table
|
||||
|
||||
|
||||
def _toolset_from_row(row: MCPToolsetRow) -> MCPToolset:
|
||||
data: Final = row.model_dump()
|
||||
tools = data.get("tools") or []
|
||||
if isinstance(tools, str):
|
||||
tools = json.loads(tools)
|
||||
data["tools"] = tools
|
||||
return MCPToolset(**data)
|
||||
tools: Final = data.get("tools") or []
|
||||
resolved: Final[MCPToolsetFields] = {
|
||||
**data,
|
||||
"tools": json.loads(tools) if isinstance(tools, str) else tools,
|
||||
}
|
||||
return MCPToolset(**resolved)
|
||||
|
||||
|
||||
async def create_mcp_toolset(
|
||||
|
|
@ -31,7 +90,7 @@ async def create_mcp_toolset(
|
|||
data_dict["tools"] = json.dumps(data_dict.get("tools", []))
|
||||
data_dict["created_by"] = touched_by
|
||||
data_dict["updated_by"] = touched_by
|
||||
row: Final = await MCPToolsetRepository(prisma_client).table.create(data=data_dict)
|
||||
row: Final = await _toolset_table(prisma_client).create(data=data_dict)
|
||||
return _toolset_from_row(row)
|
||||
|
||||
|
||||
|
|
@ -39,7 +98,7 @@ async def get_mcp_toolset(
|
|||
prisma_client: PrismaClient,
|
||||
toolset_id: str,
|
||||
) -> MCPToolset | None:
|
||||
row: Final = await MCPToolsetRepository(prisma_client).table.find_unique(where={"toolset_id": toolset_id})
|
||||
row: Final = await _toolset_table(prisma_client).find_unique(where={"toolset_id": toolset_id})
|
||||
if row is None:
|
||||
return None
|
||||
return _toolset_from_row(row)
|
||||
|
|
@ -47,13 +106,11 @@ async def get_mcp_toolset(
|
|||
|
||||
async def list_mcp_toolsets(
|
||||
prisma_client: PrismaClient,
|
||||
toolset_ids: list[str] | None = None,
|
||||
) -> list[MCPToolset]:
|
||||
toolset_ids: Sequence[str] | None = None,
|
||||
) -> Sequence[MCPToolset]:
|
||||
try:
|
||||
where = {}
|
||||
if toolset_ids is not None:
|
||||
where = {"toolset_id": {"in": toolset_ids}}
|
||||
rows: Final = await MCPToolsetRepository(prisma_client).table.find_many(where=where)
|
||||
where: Final[Mapping[str, object]] = {} if toolset_ids is None else {"toolset_id": {"in": toolset_ids}}
|
||||
rows: Final = await _toolset_table(prisma_client).find_many(where=where)
|
||||
return [_toolset_from_row(r) for r in rows]
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("litellm.proxy._experimental.mcp_server.toolset_db::list_mcp_toolsets - %s", e)
|
||||
|
|
@ -64,7 +121,7 @@ async def get_mcp_toolset_by_name(
|
|||
prisma_client: PrismaClient,
|
||||
toolset_name: str,
|
||||
) -> MCPToolset | None:
|
||||
row: Final = await MCPToolsetRepository(prisma_client).table.find_first(where={"toolset_name": toolset_name})
|
||||
row: Final = await _toolset_table(prisma_client).find_first(where={"toolset_name": toolset_name})
|
||||
if row is None:
|
||||
return None
|
||||
return _toolset_from_row(row)
|
||||
|
|
@ -80,7 +137,7 @@ async def update_mcp_toolset(
|
|||
data_dict["tools"] = json.dumps(data_dict["tools"])
|
||||
data_dict["updated_by"] = touched_by
|
||||
try:
|
||||
row: Final = await MCPToolsetRepository(prisma_client).table.update(
|
||||
row: Final = await _toolset_table(prisma_client).update(
|
||||
where={"toolset_id": data.toolset_id},
|
||||
data=data_dict,
|
||||
)
|
||||
|
|
@ -98,7 +155,7 @@ async def delete_mcp_toolset(
|
|||
toolset_id: str,
|
||||
) -> MCPToolset | None:
|
||||
try:
|
||||
row: Final = await MCPToolsetRepository(prisma_client).table.delete(where={"toolset_id": toolset_id})
|
||||
row: Final = await _toolset_table(prisma_client).delete(where={"toolset_id": toolset_id})
|
||||
except Exception as e:
|
||||
from prisma.errors import RecordNotFoundError
|
||||
|
||||
|
|
|
|||
|
|
@ -401,12 +401,10 @@ class AgentRegistry:
|
|||
The patched agent
|
||||
"""
|
||||
try:
|
||||
existing_row: Final = await AgentsRepository(prisma_client).table.find_unique(
|
||||
where={"agent_id": agent_id} # mutable-ok: prisma filters are plain dicts
|
||||
)
|
||||
if existing_row is None:
|
||||
existing_record: Final = await agents_table(prisma_client).find_unique(where={"agent_id": agent_id})
|
||||
if existing_record is None:
|
||||
raise Exception(f"Agent with ID {agent_id} not found")
|
||||
existing_agent: Final = dict(existing_row)
|
||||
existing_agent: Final[Mapping[str, object]] = dict(existing_record)
|
||||
|
||||
augment_agent: Final = {**existing_agent, **agent}
|
||||
update_data: Final[dict[str, object]] = {}
|
||||
|
|
@ -433,7 +431,7 @@ class AgentRegistry:
|
|||
update_data["extra_headers"] = extra_headers_value if extra_headers_value is not None else []
|
||||
if agent.get("object_permission") is not None:
|
||||
agent_copy: Final = dict(augment_agent)
|
||||
existing_object_permission_id: Final = existing_agent.get("object_permission_id")
|
||||
existing_object_permission_id: Final = existing_record.object_permission_id
|
||||
object_permission_id: Final = await handle_update_object_permission_common(
|
||||
agent_copy,
|
||||
existing_object_permission_id,
|
||||
|
|
|
|||
|
|
@ -34,7 +34,7 @@ def teams():
|
|||
"""Manage teams and team assignments"""
|
||||
|
||||
|
||||
def display_teams_table(teams: list[dict[str, Any]]) -> None:
|
||||
def display_teams_table(teams: Sequence[dict[str, Any]]) -> None:
|
||||
"""Display teams in a formatted table"""
|
||||
console: Final = Console()
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import copy
|
||||
import os
|
||||
from collections.abc import Callable, Iterable
|
||||
from collections.abc import Callable, Iterable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Optional, TypeAlias
|
||||
|
||||
|
|
@ -525,8 +525,8 @@ LITELLM_PROXY_INTERNAL_METADATA_KEYS: Final = frozenset(
|
|||
|
||||
|
||||
def sanitize_openai_provider_metadata(
|
||||
metadata: dict[str, Any] | None,
|
||||
) -> dict[str, str] | None:
|
||||
metadata: Mapping[str, object] | None,
|
||||
) -> Mapping[str, object] | None:
|
||||
"""
|
||||
Keep only provider-safe OpenAI metadata entries (string keys -> string values).
|
||||
|
||||
|
|
@ -644,7 +644,7 @@ def process_callback(_callback: str, callback_type: str, environment_variables:
|
|||
return {"name": _callback, "variables": env_vars_dict, "type": callback_type}
|
||||
|
||||
|
||||
def normalize_callback_names(callbacks: Iterable[Any]) -> list[Any]:
|
||||
def normalize_callback_names(callbacks: Iterable[object] | None) -> list[object]:
|
||||
if callbacks is None:
|
||||
return []
|
||||
return [c.lower() if isinstance(c, str) else c for c in callbacks]
|
||||
|
|
@ -674,7 +674,7 @@ def decrypt_callback_vars(metadata: Any) -> Any:
|
|||
return _transform_callback_vars(metadata, _decrypt_or_passthrough)
|
||||
|
||||
|
||||
def _transform_callback_vars(metadata: Any, transform: Callable[[str, Any], Any]) -> Any:
|
||||
def _transform_callback_vars(metadata: object, transform: Callable[[str, Any], Any]) -> object:
|
||||
if not isinstance(metadata, dict):
|
||||
return metadata
|
||||
out: Final = copy.deepcopy(metadata)
|
||||
|
|
@ -704,7 +704,7 @@ def is_sensitive_callback_key(
|
|||
return _CALLBACK_VAR_MASKER.is_sensitive_key(key)
|
||||
|
||||
|
||||
def _encrypt_if_plaintext(key: str, value: Any) -> Any:
|
||||
def _encrypt_if_plaintext(key: str, value: object) -> object:
|
||||
if not isinstance(value, str) or not value:
|
||||
return value
|
||||
if not is_sensitive_callback_key(key):
|
||||
|
|
@ -725,7 +725,7 @@ def _encrypt_if_plaintext(key: str, value: Any) -> Any:
|
|||
return value
|
||||
|
||||
|
||||
def _decrypt_or_passthrough(key: str, value: Any) -> Any:
|
||||
def _decrypt_or_passthrough(key: str, value: object) -> object:
|
||||
if not isinstance(value, str) or not value:
|
||||
return value
|
||||
if not value.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX):
|
||||
|
|
|
|||
|
|
@ -40,30 +40,31 @@ class UserApiKeyCache(DualCache):
|
|||
@overload
|
||||
def get_cache(
|
||||
self,
|
||||
key: Any,
|
||||
parent_otel_span: Any = None,
|
||||
key: object,
|
||||
parent_otel_span: object = None,
|
||||
local_only: bool = False,
|
||||
*,
|
||||
model_type: type[T],
|
||||
**kwargs: Any,
|
||||
**kwargs: object,
|
||||
) -> T | None: ...
|
||||
|
||||
@overload
|
||||
def get_cache(
|
||||
self,
|
||||
key: Any,
|
||||
parent_otel_span: Any = None,
|
||||
key: object,
|
||||
parent_otel_span: object = None,
|
||||
local_only: bool = False,
|
||||
**kwargs: Any,
|
||||
model_type: None = None,
|
||||
**kwargs: object,
|
||||
) -> Any: ...
|
||||
|
||||
def get_cache(
|
||||
self,
|
||||
key,
|
||||
parent_otel_span=None,
|
||||
key: object,
|
||||
parent_otel_span: object = None,
|
||||
local_only: bool = False,
|
||||
model_type: type[BaseModel] | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
) -> Any | BaseModel | None:
|
||||
if model_type is None and "model_type" in kwargs:
|
||||
model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
|
||||
|
|
@ -85,30 +86,31 @@ class UserApiKeyCache(DualCache):
|
|||
@overload
|
||||
async def async_get_cache(
|
||||
self,
|
||||
key: Any,
|
||||
parent_otel_span: Any = None,
|
||||
key: object,
|
||||
parent_otel_span: object = None,
|
||||
local_only: bool = False,
|
||||
*,
|
||||
model_type: type[T],
|
||||
**kwargs: Any,
|
||||
**kwargs: object,
|
||||
) -> T | None: ...
|
||||
|
||||
@overload
|
||||
async def async_get_cache(
|
||||
self,
|
||||
key: Any,
|
||||
parent_otel_span: Any = None,
|
||||
key: object,
|
||||
parent_otel_span: object = None,
|
||||
local_only: bool = False,
|
||||
**kwargs: Any,
|
||||
model_type: None = None,
|
||||
**kwargs: object,
|
||||
) -> Any: ...
|
||||
|
||||
async def async_get_cache(
|
||||
self,
|
||||
key,
|
||||
parent_otel_span=None,
|
||||
key: object,
|
||||
parent_otel_span: object = None,
|
||||
local_only: bool = False,
|
||||
model_type: type[BaseModel] | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
) -> Any | BaseModel | None:
|
||||
if model_type is None and "model_type" in kwargs:
|
||||
model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
|
||||
|
|
@ -129,17 +131,17 @@ class UserApiKeyCache(DualCache):
|
|||
return None
|
||||
return decoded
|
||||
|
||||
def set_cache(self, key, value, local_only: bool = False, **kwargs):
|
||||
def set_cache(self, key: object, value: object, local_only: bool = False, **kwargs: object):
|
||||
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
|
||||
payload: Final = CacheCodec.serialize(value, model_type=model_type)
|
||||
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
|
||||
return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs)
|
||||
|
||||
async def async_set_cache(self, key, value, local_only: bool = False, **kwargs):
|
||||
async def async_set_cache(self, key: object, value: object, local_only: bool = False, **kwargs: object):
|
||||
model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None))
|
||||
payload: Final = CacheCodec.serialize(value, model_type=model_type)
|
||||
payload: Final[object] = CacheCodec.serialize(value, model_type=model_type)
|
||||
return await super().async_set_cache(key=key, value=payload, local_only=local_only, **kwargs)
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs) -> None:
|
||||
async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs: object) -> None:
|
||||
"""
|
||||
Batch writes with the same Codec boundary as ``async_set_cache`` without
|
||||
``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged.
|
||||
|
|
|
|||
|
|
@ -748,7 +748,7 @@ class CompresrGuardrail(CustomGuardrail):
|
|||
}
|
||||
|
||||
try:
|
||||
raw_response: HttpxResponse | None = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped
|
||||
raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped
|
||||
url=url,
|
||||
json=payload,
|
||||
headers=self._request_headers(),
|
||||
|
|
@ -778,11 +778,11 @@ class CompresrGuardrail(CustomGuardrail):
|
|||
{"detail": str(e)},
|
||||
)
|
||||
return None
|
||||
if raw_response is None or not 200 <= raw_response.status_code < 300:
|
||||
if not 200 <= raw_response.status_code < 300:
|
||||
self._handle_compress_failure(
|
||||
"Compresr compression service returned an error",
|
||||
{
|
||||
"status_code": getattr(raw_response, "status_code", None),
|
||||
"status_code": raw_response.status_code,
|
||||
"body": _safe_response_text(raw_response),
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -433,7 +433,7 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
payload["model"] = model
|
||||
|
||||
try:
|
||||
raw_response: HttpxResponse | None = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
|
||||
raw_response: HttpxResponse = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
|
||||
url=f"{self.headroom_api_base}/v1/compress",
|
||||
json=payload,
|
||||
headers=self._request_headers(),
|
||||
|
|
@ -458,16 +458,6 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
False,
|
||||
{},
|
||||
)
|
||||
if raw_response is None:
|
||||
return (
|
||||
self._handle_compress_failure(
|
||||
messages,
|
||||
"Headroom compression service returned no response",
|
||||
{},
|
||||
),
|
||||
False,
|
||||
{},
|
||||
)
|
||||
response: Final[HttpxResponse] = raw_response
|
||||
|
||||
if response.status_code != 200:
|
||||
|
|
@ -580,7 +570,7 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
params["query"] = query
|
||||
|
||||
try:
|
||||
raw_response: HttpxResponse | None = await self.async_handler.get( # pyright: ignore[reportUnknownMemberType]
|
||||
raw_response: HttpxResponse = await self.async_handler.get( # pyright: ignore[reportUnknownMemberType]
|
||||
url=f"{self.headroom_api_base}/v1/retrieve/{hash_value}",
|
||||
params=params,
|
||||
headers=self._request_headers(),
|
||||
|
|
@ -589,7 +579,7 @@ class HeadroomGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.warning("Headroom: retrieve failed for hash=%s: %s", hash_value, e)
|
||||
return f"[Headroom: retrieval failed for hash={hash_value}]"
|
||||
|
||||
if raw_response is None or raw_response.status_code == 404:
|
||||
if raw_response.status_code == 404:
|
||||
return f"[Headroom: hash={hash_value} not found or expired]"
|
||||
|
||||
if raw_response.status_code != 200:
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ from __future__ import annotations
|
|||
|
||||
import os
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypedDict
|
||||
from typing import TYPE_CHECKING, Final, Literal, Protocol
|
||||
from urllib.parse import urlparse
|
||||
from uuid import uuid4
|
||||
|
||||
|
|
@ -11,7 +11,7 @@ import requests
|
|||
from fastapi import HTTPException
|
||||
from httpx import HTTPStatusError
|
||||
from requests.auth import HTTPBasicAuth
|
||||
from typing_extensions import ReadOnly
|
||||
from typing_extensions import ReadOnly, TypedDict, Unpack
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
|
|
@ -40,14 +40,21 @@ if TYPE_CHECKING:
|
|||
_AUTH_TIMEOUT_SECONDS: Final[float] = 30.0
|
||||
|
||||
|
||||
class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object):
|
||||
"""Base-class constructor options carried by this guardrail's forwarded keyword arguments."""
|
||||
|
||||
guardrail_name: ReadOnly[str | None]
|
||||
supported_event_hooks: list[GuardrailEventHooks] | None
|
||||
|
||||
|
||||
class _HiddenlayerEvaluation(TypedDict, total=False):
|
||||
action: str
|
||||
threat_level: str
|
||||
action: ReadOnly[str]
|
||||
threat_level: ReadOnly[str]
|
||||
|
||||
|
||||
class _HiddenlayerAnalysisEntry(TypedDict, total=False):
|
||||
name: str
|
||||
detected: bool
|
||||
name: ReadOnly[str]
|
||||
detected: ReadOnly[bool]
|
||||
|
||||
|
||||
class _HiddenlayerModifiedMessage(TypedDict):
|
||||
|
|
@ -59,9 +66,9 @@ class _HiddenlayerModifiedSide(TypedDict):
|
|||
|
||||
|
||||
class _HiddenlayerResponse(TypedDict, total=False):
|
||||
evaluation: _HiddenlayerEvaluation
|
||||
analysis: Sequence[_HiddenlayerAnalysisEntry]
|
||||
modified_data: Mapping[str, _HiddenlayerModifiedSide]
|
||||
evaluation: ReadOnly[_HiddenlayerEvaluation]
|
||||
analysis: ReadOnly[Sequence[_HiddenlayerAnalysisEntry]]
|
||||
modified_data: ReadOnly[Mapping[str, _HiddenlayerModifiedSide]]
|
||||
|
||||
|
||||
class _ProxyServerRequest(TypedDict, total=False):
|
||||
|
|
@ -194,7 +201,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
auth_url: str | None = None,
|
||||
**kwargs: Any,
|
||||
**kwargs: Unpack[_CustomGuardrailOptions],
|
||||
) -> None:
|
||||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
self.hiddenlayer_client_id = api_id or os.getenv("HIDDENLAYER_CLIENT_ID")
|
||||
|
|
@ -399,7 +406,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
auth_url: str | None = None,
|
||||
**kwargs: Any,
|
||||
**kwargs: Unpack[_CustomGuardrailOptions],
|
||||
) -> None:
|
||||
self.hiddenlayer_client_id = api_id or os.getenv("HIDDENLAYER_CLIENT_ID")
|
||||
self.hiddenlayer_client_secret = api_key or os.getenv("HIDDENLAYER_CLIENT_SECRET")
|
||||
|
|
@ -530,7 +537,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
|
|||
self,
|
||||
payload: _HiddenlayerV2Payload,
|
||||
input_type: Literal["request", "response"],
|
||||
hl_headers: dict[str, str],
|
||||
hl_headers: Mapping[str, str],
|
||||
) -> httpx.Response:
|
||||
if input_type == "request":
|
||||
path = "detection/v2/request-evaluations"
|
||||
|
|
|
|||
|
|
@ -7,8 +7,9 @@
|
|||
import enum
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Callable, Mapping
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -36,6 +37,8 @@ _AIDR_SCAN_ENDPOINT: Final = "/litellm/guardrail"
|
|||
_INTERVENED_INPUT_FIELDS: Final = ("texts", "images", "tools", "tool_calls")
|
||||
_DEFAULT_API_BASE_HOSTNAME: Final = urlparse(_DEFAULT_API_BASE).hostname
|
||||
|
||||
_GuardrailJsonResponse: TypeAlias = Exception | str | dict[str, object]
|
||||
|
||||
_KEYS_DUPLICATING_SCAN_INPUTS: Final = ("messages", "input")
|
||||
_LOGGING_KEYS_DUPLICATING_SCAN_INPUTS: Final = _KEYS_DUPLICATING_SCAN_INPUTS + (
|
||||
"additional_args",
|
||||
|
|
@ -119,7 +122,7 @@ class NomaV2Guardrail(CustomGuardrail):
|
|||
|
||||
def _resolve_action_from_response(
|
||||
self,
|
||||
response_json: dict,
|
||||
response_json: Mapping[str, object],
|
||||
) -> _Action:
|
||||
action: Final = response_json.get("action")
|
||||
if isinstance(action, str):
|
||||
|
|
@ -165,10 +168,11 @@ class NomaV2Guardrail(CustomGuardrail):
|
|||
|
||||
@staticmethod
|
||||
def _sanitize_payload_for_transport(payload: dict) -> dict:
|
||||
def _default(obj: Any) -> Any:
|
||||
if hasattr(obj, "model_dump"):
|
||||
def _default(obj: object) -> object:
|
||||
model_dump: Final[Callable[[], Mapping[str, object]] | None] = getattr(obj, "model_dump", None)
|
||||
if model_dump is not None:
|
||||
try:
|
||||
return obj.model_dump()
|
||||
return model_dump()
|
||||
except Exception:
|
||||
pass
|
||||
return str(obj)
|
||||
|
|
@ -178,7 +182,7 @@ class NomaV2Guardrail(CustomGuardrail):
|
|||
except (ValueError, TypeError):
|
||||
json_str = safe_dumps(payload)
|
||||
|
||||
safe_payload: Final = safe_json_loads(json_str, default={})
|
||||
safe_payload: Final[object] = safe_json_loads(json_str, default={})
|
||||
if safe_payload == {} and payload:
|
||||
verbose_proxy_logger.warning(
|
||||
"Noma v2 guardrail: payload serialization failed, falling back to empty payload"
|
||||
|
|
@ -196,7 +200,7 @@ class NomaV2Guardrail(CustomGuardrail):
|
|||
async def _call_noma_scan(
|
||||
self,
|
||||
payload: dict,
|
||||
) -> dict:
|
||||
) -> dict[str, object]:
|
||||
headers: Final[dict[str, str]] = {"Content-Type": "application/json"}
|
||||
authorization_header: Final = self._get_authorization_header()
|
||||
if authorization_header:
|
||||
|
|
@ -215,7 +219,7 @@ class NomaV2Guardrail(CustomGuardrail):
|
|||
response.text,
|
||||
)
|
||||
response.raise_for_status()
|
||||
response_json: Final = response.json()
|
||||
response_json: Final[dict[str, object]] = response.json()
|
||||
verbose_proxy_logger.debug(
|
||||
"Noma v2 AIDR response parsed: %s",
|
||||
json.dumps(response_json, default=str),
|
||||
|
|
@ -227,7 +231,7 @@ class NomaV2Guardrail(CustomGuardrail):
|
|||
request_data: dict,
|
||||
start_time: datetime,
|
||||
guardrail_status: GuardrailStatus,
|
||||
guardrail_json_response: Any,
|
||||
guardrail_json_response: _GuardrailJsonResponse,
|
||||
) -> None:
|
||||
end_time: Final = datetime.now()
|
||||
duration: Final = (end_time - start_time).total_seconds()
|
||||
|
|
@ -270,11 +274,11 @@ class NomaV2Guardrail(CustomGuardrail):
|
|||
) -> GenericGuardrailAPIInputs:
|
||||
start_time: Final = datetime.now()
|
||||
guardrail_status: GuardrailStatus = "success"
|
||||
guardrail_json_response: Any = {}
|
||||
guardrail_json_response: _GuardrailJsonResponse = {}
|
||||
dynamic_params = self.get_guardrail_dynamic_request_body_params(request_data)
|
||||
if not isinstance(dynamic_params, dict):
|
||||
dynamic_params = {}
|
||||
response_json: dict | None = None
|
||||
response_json: dict[str, object] | None = None
|
||||
|
||||
# Per-request dynamic params can override configured application context.
|
||||
application_id = self._get_non_empty_str(dynamic_params.get("application_id"))
|
||||
|
|
|
|||
|
|
@ -197,14 +197,11 @@ class RepelloAIGuardrail(CustomGuardrail):
|
|||
repelloai_response: RepelloAIAnalyzeResponse | None = None
|
||||
try:
|
||||
verbose_proxy_logger.debug("RepelloAI Argus request: %s", request)
|
||||
raw_response: HttpxResponse | None = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
|
||||
response: Final[HttpxResponse] = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
|
||||
url=endpoint,
|
||||
headers={"X-API-Key": self.repelloai_api_key},
|
||||
json=request,
|
||||
)
|
||||
if raw_response is None:
|
||||
raise ValueError("RepelloAI Argus returned no response")
|
||||
response: Final[HttpxResponse] = raw_response
|
||||
self._raise_for_config_error(response)
|
||||
response.raise_for_status()
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
from typing import Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import Final, Protocol
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
|
|
@ -18,7 +20,59 @@ from litellm.repositories.table_repositories import JWTKeyMappingRepository
|
|||
router: Final = APIRouter()
|
||||
|
||||
|
||||
def _to_response(mapping) -> JWTKeyMappingResponse:
|
||||
class _JWTKeyMappingRecord(Protocol):
|
||||
"""A ``LiteLLM_JWTKeyMapping`` row, viewed through the columns these endpoints read."""
|
||||
|
||||
@property
|
||||
def id(self) -> str: ...
|
||||
|
||||
@property
|
||||
def jwt_claim_name(self) -> str: ...
|
||||
|
||||
@property
|
||||
def jwt_claim_value(self) -> str: ...
|
||||
|
||||
@property
|
||||
def description(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def is_active(self) -> bool: ...
|
||||
|
||||
@property
|
||||
def created_at(self) -> datetime: ...
|
||||
|
||||
@property
|
||||
def updated_at(self) -> datetime: ...
|
||||
|
||||
@property
|
||||
def created_by(self) -> str | None: ...
|
||||
|
||||
@property
|
||||
def updated_by(self) -> str | None: ...
|
||||
|
||||
|
||||
class _JWTKeyMappingTable(Protocol):
|
||||
"""The Prisma table actions these endpoints issue against the JWT key mapping table."""
|
||||
|
||||
async def create(self, *, data: Mapping[str, object]) -> _JWTKeyMappingRecord: ...
|
||||
|
||||
async def find_unique(self, *, where: Mapping[str, object]) -> _JWTKeyMappingRecord | None: ...
|
||||
|
||||
async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> _JWTKeyMappingRecord: ...
|
||||
|
||||
async def delete(self, *, where: Mapping[str, object]) -> _JWTKeyMappingRecord | None: ...
|
||||
|
||||
async def find_many(self, *, skip: int, take: int, order: Mapping[str, str]) -> Sequence[_JWTKeyMappingRecord]: ...
|
||||
|
||||
async def count(self) -> int: ...
|
||||
|
||||
|
||||
def _mapping_table(prisma_client: object) -> _JWTKeyMappingTable:
|
||||
"""View the JWT key mapping repository's untyped Prisma table through the actions used here."""
|
||||
return JWTKeyMappingRepository(prisma_client).table
|
||||
|
||||
|
||||
def _to_response(mapping: _JWTKeyMappingRecord) -> JWTKeyMappingResponse:
|
||||
"""Convert a Prisma mapping object to a safe response (no hashed token)."""
|
||||
return JWTKeyMappingResponse(
|
||||
id=mapping.id,
|
||||
|
|
@ -62,7 +116,7 @@ async def create_jwt_key_mapping(
|
|||
if data.description is not None:
|
||||
create_data["description"] = data.description
|
||||
|
||||
new_mapping: Final = await JWTKeyMappingRepository(prisma_client).table.create(data=create_data)
|
||||
new_mapping: Final = await _mapping_table(prisma_client).create(data=create_data)
|
||||
|
||||
# Invalidate cache
|
||||
cache_key: Final = f"jwt_key_mapping:{data.jwt_claim_name}:{data.jwt_claim_value}"
|
||||
|
|
@ -110,7 +164,7 @@ async def update_jwt_key_mapping(
|
|||
|
||||
try:
|
||||
# Get old mapping for cache invalidation
|
||||
old_mapping: Final = await JWTKeyMappingRepository(prisma_client).table.find_unique(where={"id": data.id})
|
||||
old_mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": data.id})
|
||||
|
||||
if old_mapping is None:
|
||||
raise HTTPException(status_code=404, detail="Mapping not found")
|
||||
|
|
@ -118,9 +172,7 @@ async def update_jwt_key_mapping(
|
|||
cache_key = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}"
|
||||
await user_api_key_cache.async_delete_cache(cache_key)
|
||||
|
||||
updated_mapping: Final = await JWTKeyMappingRepository(prisma_client).table.update(
|
||||
where={"id": data.id}, data=update_data
|
||||
)
|
||||
updated_mapping: Final = await _mapping_table(prisma_client).update(where={"id": data.id}, data=update_data)
|
||||
|
||||
if updated_mapping is None:
|
||||
raise HTTPException(status_code=404, detail="Mapping not found")
|
||||
|
|
@ -162,7 +214,7 @@ async def delete_jwt_key_mapping(
|
|||
|
||||
try:
|
||||
# Get old mapping for cache invalidation
|
||||
old_mapping: Final = await JWTKeyMappingRepository(prisma_client).table.find_unique(where={"id": data.id})
|
||||
old_mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": data.id})
|
||||
|
||||
if old_mapping is None:
|
||||
raise HTTPException(status_code=404, detail="Mapping not found")
|
||||
|
|
@ -170,7 +222,7 @@ async def delete_jwt_key_mapping(
|
|||
cache_key: Final = f"jwt_key_mapping:{old_mapping.jwt_claim_name}:{old_mapping.jwt_claim_value}"
|
||||
await user_api_key_cache.async_delete_cache(cache_key)
|
||||
|
||||
await JWTKeyMappingRepository(prisma_client).table.delete(where={"id": data.id})
|
||||
await _mapping_table(prisma_client).delete(where={"id": data.id})
|
||||
return {"status": "success"}
|
||||
except HTTPException:
|
||||
raise
|
||||
|
|
@ -198,12 +250,12 @@ async def list_jwt_key_mappings(
|
|||
|
||||
try:
|
||||
skip: Final = (page - 1) * size
|
||||
mappings: Final = await JWTKeyMappingRepository(prisma_client).table.find_many(
|
||||
mappings: Final = await _mapping_table(prisma_client).find_many(
|
||||
skip=skip,
|
||||
take=size,
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
total_count: Final = await JWTKeyMappingRepository(prisma_client).table.count()
|
||||
total_count: Final = await _mapping_table(prisma_client).count()
|
||||
return {
|
||||
"mappings": [_to_response(m) for m in mappings],
|
||||
"total_count": total_count,
|
||||
|
|
@ -235,7 +287,7 @@ async def info_jwt_key_mapping(
|
|||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
mapping: Final = await JWTKeyMappingRepository(prisma_client).table.find_unique(where={"id": id})
|
||||
mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": id})
|
||||
if mapping is None:
|
||||
raise HTTPException(status_code=404, detail="Mapping not found")
|
||||
return _to_response(mapping)
|
||||
|
|
|
|||
|
|
@ -2,8 +2,9 @@
|
|||
Search Tool Registry for managing search tool configurations.
|
||||
"""
|
||||
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from typing import Final
|
||||
from typing import Final, Protocol
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
|
@ -13,6 +14,40 @@ from litellm.repositories.table_repositories import SearchToolsRepository
|
|||
from litellm.types.search import SearchTool
|
||||
|
||||
|
||||
class SearchToolRecord(Protocol):
|
||||
search_tool_id: str
|
||||
search_tool_name: str
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
def __iter__(self) -> Iterator[tuple[str, object]]: ...
|
||||
|
||||
|
||||
class SearchToolTableClient(Protocol):
|
||||
async def create(self, data: Mapping[str, object]) -> SearchToolRecord: ...
|
||||
|
||||
async def find_unique(self, where: Mapping[str, object]) -> SearchToolRecord | None: ...
|
||||
|
||||
async def find_many(self, order: Mapping[str, str] | None = None) -> Sequence[SearchToolRecord]: ...
|
||||
|
||||
async def update(self, where: Mapping[str, object], data: Mapping[str, object]) -> SearchToolRecord: ...
|
||||
|
||||
async def delete(self, where: Mapping[str, object]) -> SearchToolRecord: ...
|
||||
|
||||
|
||||
class _SearchToolsRepositoryView(Protocol):
|
||||
@property
|
||||
def table(self) -> SearchToolTableClient: ...
|
||||
|
||||
|
||||
def _search_tools_table_of(repository: _SearchToolsRepositoryView) -> SearchToolTableClient:
|
||||
return repository.table
|
||||
|
||||
|
||||
def _search_tools_table(prisma_client: PrismaClient) -> SearchToolTableClient:
|
||||
return _search_tools_table_of(SearchToolsRepository(prisma_client))
|
||||
|
||||
|
||||
class SearchToolRegistry:
|
||||
"""
|
||||
Handles adding, removing, and getting search tools in DB + in memory.
|
||||
|
|
@ -22,7 +57,7 @@ class SearchToolRegistry:
|
|||
pass
|
||||
|
||||
@staticmethod
|
||||
def _convert_prisma_to_dict(prisma_obj) -> dict:
|
||||
def _convert_prisma_to_dict(prisma_obj: SearchToolRecord) -> dict:
|
||||
"""
|
||||
Convert Prisma result to dict with datetime objects as ISO format strings.
|
||||
|
||||
|
|
@ -35,9 +70,9 @@ class SearchToolRegistry:
|
|||
result: Final = dict(prisma_obj)
|
||||
# Convert datetime objects to ISO format strings
|
||||
if "created_at" in result and result["created_at"]:
|
||||
result["created_at"] = result["created_at"].isoformat()
|
||||
result["created_at"] = prisma_obj.created_at.isoformat()
|
||||
if "updated_at" in result and result["updated_at"]:
|
||||
result["updated_at"] = result["updated_at"].isoformat()
|
||||
result["updated_at"] = prisma_obj.updated_at.isoformat()
|
||||
return result
|
||||
|
||||
###########################################################
|
||||
|
|
@ -61,7 +96,7 @@ class SearchToolRegistry:
|
|||
search_tool_info: Final[str] = safe_dumps(search_tool.get("search_tool_info", {}))
|
||||
|
||||
# Create search tool in DB
|
||||
created_search_tool: Final = await SearchToolsRepository(prisma_client).table.create(
|
||||
created_search_tool: Final = await _search_tools_table(prisma_client).create(
|
||||
data={
|
||||
"search_tool_name": search_tool_name,
|
||||
"litellm_params": litellm_params,
|
||||
|
|
@ -95,7 +130,7 @@ class SearchToolRegistry:
|
|||
"""
|
||||
try:
|
||||
# Get search tool before deletion for response
|
||||
existing_tool: Final = await SearchToolsRepository(prisma_client).table.find_unique(
|
||||
existing_tool: Final = await _search_tools_table(prisma_client).find_unique(
|
||||
where={"search_tool_id": search_tool_id}
|
||||
)
|
||||
|
||||
|
|
@ -103,7 +138,7 @@ class SearchToolRegistry:
|
|||
raise Exception(f"Search tool with ID {search_tool_id} not found")
|
||||
|
||||
# Delete from DB
|
||||
await SearchToolsRepository(prisma_client).table.delete(where={"search_tool_id": search_tool_id})
|
||||
await _search_tools_table(prisma_client).delete(where={"search_tool_id": search_tool_id})
|
||||
|
||||
return {
|
||||
"message": f"Search tool {search_tool_id} deleted successfully",
|
||||
|
|
@ -131,7 +166,7 @@ class SearchToolRegistry:
|
|||
search_tool_info: Final[str] = safe_dumps(search_tool.get("search_tool_info", {}))
|
||||
|
||||
# Update in DB
|
||||
updated_search_tool: Final = await SearchToolsRepository(prisma_client).table.update(
|
||||
updated_search_tool: Final = await _search_tools_table(prisma_client).update(
|
||||
where={"search_tool_id": search_tool_id},
|
||||
data={
|
||||
"search_tool_name": search_tool_name,
|
||||
|
|
@ -163,7 +198,7 @@ class SearchToolRegistry:
|
|||
try:
|
||||
search_tools_from_db: Final = await call_with_db_reconnect_retry(
|
||||
prisma_client,
|
||||
lambda: SearchToolsRepository(prisma_client).table.find_many(
|
||||
lambda: _search_tools_table(prisma_client).find_many(
|
||||
order={"created_at": "desc"},
|
||||
),
|
||||
reason="get_all_search_tools_from_db_lookup_failure",
|
||||
|
|
@ -194,7 +229,7 @@ class SearchToolRegistry:
|
|||
Search tool configuration or None if not found
|
||||
"""
|
||||
try:
|
||||
search_tool: Final = await SearchToolsRepository(prisma_client).table.find_unique(
|
||||
search_tool: Final = await _search_tools_table(prisma_client).find_unique(
|
||||
where={"search_tool_id": search_tool_id}
|
||||
)
|
||||
|
||||
|
|
@ -222,7 +257,7 @@ class SearchToolRegistry:
|
|||
Search tool configuration or None if not found
|
||||
"""
|
||||
try:
|
||||
search_tool: Final = await SearchToolsRepository(prisma_client).table.find_unique(
|
||||
search_tool: Final = await _search_tools_table(prisma_client).find_unique(
|
||||
where={"search_tool_name": search_tool_name}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,11 +1,45 @@
|
|||
from collections.abc import Sequence
|
||||
from typing import Any, Final
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.types.vector_stores import (
|
||||
VectorStoreResultContent,
|
||||
VectorStoreSearchResponse,
|
||||
)
|
||||
from litellm.types.vector_stores import VectorStoreSearchResponse
|
||||
|
||||
|
||||
class _ResultContentView(TypedDict):
|
||||
"""Content entry carried by a vector store search result."""
|
||||
|
||||
type: ReadOnly[NotRequired[str]]
|
||||
text: ReadOnly[str]
|
||||
|
||||
|
||||
class _SearchResultView(TypedDict):
|
||||
"""Vector store search result, as far as :class:`RAGQuery` reads it."""
|
||||
|
||||
content: ReadOnly[NotRequired[Sequence[_ResultContentView]]]
|
||||
text: ReadOnly[NotRequired[str]]
|
||||
|
||||
|
||||
class _SearchDataView(TypedDict):
|
||||
results: ReadOnly[Sequence[_SearchResultView]]
|
||||
|
||||
|
||||
class _ContextChunksView(TypedDict):
|
||||
chunks: ReadOnly[Sequence[_SearchResultView | str | None]]
|
||||
|
||||
|
||||
class _RerankResultView(TypedDict):
|
||||
index: ReadOnly[NotRequired[int]]
|
||||
|
||||
|
||||
class _RerankResultsView(TypedDict):
|
||||
results: ReadOnly[Sequence[_RerankResultView]]
|
||||
|
||||
|
||||
class _MessageView(TypedDict):
|
||||
message: ReadOnly[object]
|
||||
|
||||
|
||||
class RAGQuery:
|
||||
|
|
@ -42,9 +76,10 @@ class RAGQuery:
|
|||
"""
|
||||
context_content = RAGQuery.CONTENT_PREFIX_STRING
|
||||
|
||||
for chunk in context_chunks:
|
||||
chunks: Final[_ContextChunksView] = {"chunks": context_chunks}
|
||||
for chunk in chunks["chunks"]:
|
||||
if isinstance(chunk, dict):
|
||||
result_content: list[VectorStoreResultContent] | None = chunk.get("content")
|
||||
result_content: Sequence[_ResultContentView] | None = chunk.get("content")
|
||||
if result_content:
|
||||
for content_item in result_content:
|
||||
content_text: str | None = content_item.get("text")
|
||||
|
|
@ -64,14 +99,15 @@ class RAGQuery:
|
|||
def add_search_results_to_response(
|
||||
response: ModelResponse,
|
||||
search_results: VectorStoreSearchResponse,
|
||||
rerank_results: Any | None = None,
|
||||
rerank_results: object = None,
|
||||
) -> ModelResponse:
|
||||
"""
|
||||
Add search results to the response choices.
|
||||
"""
|
||||
if hasattr(response, "choices") and response.choices:
|
||||
for choice in response.choices:
|
||||
message = getattr(choice, "message", None)
|
||||
message_view: _MessageView = {"message": getattr(choice, "message", None)}
|
||||
message = message_view["message"]
|
||||
if message is not None:
|
||||
# Get existing provider_specific_fields or create new dict
|
||||
provider_fields = getattr(message, "provider_specific_fields", None) or {}
|
||||
|
|
@ -91,7 +127,8 @@ class RAGQuery:
|
|||
) -> list[str | dict[str, Any]]:
|
||||
"""Extract text documents from vector store search response."""
|
||||
documents: Final[list[str | dict[str, Any]]] = []
|
||||
for result in search_response.get("data", []):
|
||||
search_data: Final[_SearchDataView] = {"results": search_response.get("data", [])}
|
||||
for result in search_data["results"]:
|
||||
content_list = result.get("content", [])
|
||||
for content in content_list:
|
||||
if content.get("type") == "text" and content.get("text"):
|
||||
|
|
@ -99,11 +136,13 @@ class RAGQuery:
|
|||
return documents
|
||||
|
||||
@staticmethod
|
||||
def get_top_chunks_from_rerank(search_response: Any, rerank_response: Any) -> list[Any]:
|
||||
def get_top_chunks_from_rerank(search_response: Any, rerank_response: Any) -> list[_SearchResultView]:
|
||||
"""Get the original search results corresponding to the top reranked results."""
|
||||
top_chunks: Final = []
|
||||
original_results: Final = search_response.get("data", [])
|
||||
for result in rerank_response.get("results", []):
|
||||
top_chunks: Final[list[_SearchResultView]] = []
|
||||
search_data: Final[_SearchDataView] = {"results": search_response.get("data", [])}
|
||||
original_results: Final = search_data["results"]
|
||||
reranked: Final[_RerankResultsView] = {"results": rerank_response.get("results", [])}
|
||||
for result in reranked["results"]:
|
||||
index = result.get("index")
|
||||
if index is not None and index < len(original_results):
|
||||
top_chunks.append(original_results[index])
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ def record_to_dict(record: DbRecord) -> Mapping[str, object]:
|
|||
class BaseRepository(ABC, Generic[T]):
|
||||
"""Abstract base class for all repositories."""
|
||||
|
||||
def __init__(self, prisma_client: Any): # any-ok: PrismaClient is an untyped runtime wrapper
|
||||
def __init__(self, prisma_client: object):
|
||||
self._prisma_client = prisma_client
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -3,9 +3,9 @@ Team repository for database operations on LiteLLM_TeamTable.
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
|
|
@ -21,6 +21,25 @@ if TYPE_CHECKING:
|
|||
from prisma import Prisma
|
||||
from prisma import models as prisma_models
|
||||
|
||||
|
||||
class _TeamArrays(Protocol):
|
||||
"""The string array columns of a team row, which the domain model leaves untyped."""
|
||||
|
||||
@property
|
||||
def members(self) -> Sequence[str]: ...
|
||||
|
||||
@property
|
||||
def admins(self) -> Sequence[str]: ...
|
||||
|
||||
@property
|
||||
def models(self) -> Sequence[str]: ...
|
||||
|
||||
|
||||
def _team_arrays(team: LiteLLM_TeamTable) -> _TeamArrays:
|
||||
"""View a team's untyped list columns as sequences of ids."""
|
||||
return team
|
||||
|
||||
|
||||
_MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member])
|
||||
_JSON_ENCODED_TEAM_FIELDS: Final = (
|
||||
"metadata",
|
||||
|
|
@ -80,8 +99,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
|
|||
)
|
||||
if not rows:
|
||||
return None
|
||||
raw_value: Final = rows[0]["members_with_roles"]
|
||||
parsed: Final = json.loads(raw_value) if isinstance(raw_value, str) else raw_value
|
||||
raw_value: Final[object] = rows[0]["members_with_roles"]
|
||||
parsed: Final[object] = json.loads(raw_value) if isinstance(raw_value, str) else raw_value
|
||||
if not parsed:
|
||||
return []
|
||||
return _MEMBERS_WITH_ROLES_ADAPTER.validate_python(parsed)
|
||||
|
|
@ -315,7 +334,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
|
|||
if team is None:
|
||||
return None
|
||||
|
||||
members: Final = [m for m in team.members if m != user_id]
|
||||
members: Final = [m for m in _team_arrays(team).members if m != user_id]
|
||||
return await self.update(team_id, {"members": members}, id_field="team_id")
|
||||
|
||||
async def add_admin(self, team_id: str, user_id: str) -> LiteLLM_TeamTable | None:
|
||||
|
|
@ -340,7 +359,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
|
|||
if team is None:
|
||||
return None
|
||||
|
||||
admins: Final = [a for a in team.admins if a != user_id]
|
||||
admins: Final = [a for a in _team_arrays(team).admins if a != user_id]
|
||||
return await self.update(team_id, {"admins": admins}, id_field="team_id")
|
||||
|
||||
async def add_models(self, team_id: str, models: list[str]) -> LiteLLM_TeamTable | None:
|
||||
|
|
@ -365,5 +384,5 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
|
|||
if team is None:
|
||||
return None
|
||||
|
||||
current_models: Final = [m for m in team.models if m not in models]
|
||||
current_models: Final = [m for m in _team_arrays(team).models if m not in models]
|
||||
return await self.update(team_id, {"models": current_models}, id_field="team_id")
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 3009
|
||||
"limit": 2995
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 71
|
||||
},
|
||||
"ANN003": {
|
||||
"limit": 827
|
||||
"limit": 809
|
||||
},
|
||||
"ANN201": {
|
||||
"limit": 2002
|
||||
|
|
@ -15,7 +15,7 @@
|
|||
"limit": 845
|
||||
},
|
||||
"ANN204": {
|
||||
"limit": 702
|
||||
"limit": 698
|
||||
},
|
||||
"ANN205": {
|
||||
"limit": 112
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 133
|
||||
},
|
||||
"ANN401": {
|
||||
"limit": 644
|
||||
"limit": 587
|
||||
},
|
||||
"ASYNC230": {
|
||||
"limit": 11
|
||||
|
|
@ -195,7 +195,7 @@
|
|||
"limit": 22
|
||||
},
|
||||
"SIM101": {
|
||||
"limit": 58
|
||||
"limit": 56
|
||||
},
|
||||
"SIM102": {
|
||||
"limit": 314
|
||||
|
|
@ -231,7 +231,7 @@
|
|||
"limit": 5
|
||||
},
|
||||
"TID251": {
|
||||
"limit": 1112
|
||||
"limit": 1108
|
||||
},
|
||||
"TRY002": {
|
||||
"limit": 524
|
||||
|
|
@ -246,7 +246,7 @@
|
|||
"limit": 113
|
||||
},
|
||||
"TRY300": {
|
||||
"limit": 857
|
||||
"limit": 855
|
||||
},
|
||||
"UP028": {
|
||||
"limit": 2
|
||||
|
|
|
|||
|
|
@ -310,6 +310,25 @@ class TestGDCGeminiConfig:
|
|||
api_base=TEST_API_BASE,
|
||||
)
|
||||
|
||||
def test_validate_environment_credentials_missing_audience_binding_are_named(self):
|
||||
config = GDCGeminiConfig()
|
||||
creds_without_audience_binding = MagicMock(spec=[])
|
||||
|
||||
with patch(
|
||||
"google.auth.load_credentials_from_dict",
|
||||
return_value=(creds_without_audience_binding, None),
|
||||
):
|
||||
with pytest.raises(AttributeError, match="must expose with_gdch_audience"):
|
||||
config.validate_environment(
|
||||
headers={},
|
||||
model=TEST_MODEL,
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={"vertex_project": TEST_PROJECT},
|
||||
api_key=TEST_API_KEY,
|
||||
api_base=TEST_API_BASE,
|
||||
)
|
||||
|
||||
def test_validate_environment_string_false_disables_token_caching(self):
|
||||
config = GDCGeminiConfig()
|
||||
mock_creds = MagicMock()
|
||||
|
|
|
|||
|
|
@ -224,20 +224,6 @@ async def test_fetch_invalid_json_maps_to_upstream_unavailable():
|
|||
assert "idp.example.com" not in result.error.summary
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_none_response_is_upstream_unavailable():
|
||||
with patch(_PATCH_TARGET, return_value=_client(None)):
|
||||
result = await TokenEndpointClient().fetch(
|
||||
_ENDPOINT,
|
||||
_CLIENT_ID,
|
||||
{"grant_type": "g"},
|
||||
ClientSecretAuth(client_secret=SecretStr("s")),
|
||||
)
|
||||
|
||||
assert isinstance(result, Error)
|
||||
assert result.error.tag == "upstream_unavailable"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_missing_access_token_is_upstream_unavailable():
|
||||
bad = MagicMock()
|
||||
|
|
@ -275,21 +261,6 @@ async def test_fetch_http_error_does_not_leak_endpoint_url():
|
|||
assert "idp.example.com" not in result.error.summary
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_none_response_does_not_leak_endpoint_url():
|
||||
with patch(_PATCH_TARGET, return_value=_client(None)):
|
||||
result = await TokenEndpointClient().fetch(
|
||||
_ENDPOINT,
|
||||
_CLIENT_ID,
|
||||
{"grant_type": "g"},
|
||||
ClientSecretAuth(client_secret=SecretStr("s")),
|
||||
)
|
||||
|
||||
assert isinstance(result, Error)
|
||||
assert _ENDPOINT not in result.error.summary
|
||||
assert "idp.example.com" not in result.error.summary
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_missing_access_token_does_not_leak_endpoint_url():
|
||||
bad = MagicMock()
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22583
|
||||
"limit": 22521
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26840
|
||||
"limit": 26820
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 269
|
||||
|
|
@ -15,7 +15,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT006": {
|
||||
"limit": 1060
|
||||
"limit": 1039
|
||||
},
|
||||
"LIT007": {
|
||||
"limit": 0
|
||||
|
|
@ -27,12 +27,12 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16560
|
||||
"limit": 16546
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5576
|
||||
"limit": 5575
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4505
|
||||
"limit": 4495
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue