mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge pull request #36054 from BerriAI/litellm_reduce_any_types
refactor(types): cut 653 implicit and explicit Any diagnostics across 11 modules
This commit is contained in:
commit
9e7b05731d
15 changed files with 558 additions and 247 deletions
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"limit": 29204
|
||||
"limit": 28842
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"limit": 2634
|
||||
|
|
@ -24,7 +24,7 @@
|
|||
"limit": 19
|
||||
},
|
||||
"reportExplicitAny": {
|
||||
"limit": 9225
|
||||
"limit": 9103
|
||||
},
|
||||
"reportFunctionMemberAccess": {
|
||||
"limit": 7
|
||||
|
|
@ -54,10 +54,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportMissingParameterType": {
|
||||
"limit": 5850
|
||||
"limit": 5843
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15833
|
||||
"limit": 15816
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -99,19 +99,19 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 45145
|
||||
"limit": 45110
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 113
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 39881
|
||||
"limit": 39838
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 20258
|
||||
"limit": 20237
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 31429
|
||||
"limit": 31383
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 122
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ import ast
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Callable, Mapping
|
||||
from typing import Any, Final, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -47,7 +48,7 @@ class RedisSemanticCache(BaseCache):
|
|||
similarity_threshold: float | None = None,
|
||||
embedding_model: str = "text-embedding-ada-002",
|
||||
index_name: str | None = None,
|
||||
**kwargs,
|
||||
**kwargs: object,
|
||||
):
|
||||
"""
|
||||
Initialize the Redis Semantic Cache.
|
||||
|
|
@ -150,11 +151,11 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
def _init_semantic_cache(
|
||||
self,
|
||||
semantic_cache_cls: Any,
|
||||
semantic_cache_cls: Callable[..., object],
|
||||
index_name: str,
|
||||
redis_url: str,
|
||||
cache_vectorizer: Any,
|
||||
) -> Any:
|
||||
cache_vectorizer: object,
|
||||
) -> object:
|
||||
def _is_schema_mismatch(exc: ValueError) -> bool:
|
||||
error_message: Final = str(exc).lower()
|
||||
return any(phrase in error_message for phrase in ("schema does not match", "index schema"))
|
||||
|
|
@ -206,12 +207,12 @@ class RedisSemanticCache(BaseCache):
|
|||
def _get_cache_filters(self, key: str) -> dict[str, str]:
|
||||
return {self.CACHE_KEY_FIELD_NAME: str(key)}
|
||||
|
||||
def _get_cache_key_filter_expression(self, key: str) -> Any:
|
||||
def _get_cache_key_filter_expression(self, key: str) -> object:
|
||||
from redisvl.query.filter import Tag
|
||||
|
||||
return Tag(self.CACHE_KEY_FIELD_NAME) == str(key)
|
||||
|
||||
def _cache_hit_matches_key(self, cache_hit: dict[str, Any], key: str) -> bool:
|
||||
def _cache_hit_matches_key(self, cache_hit: Mapping[str, object], key: str) -> bool:
|
||||
# Pre-isolation entries with no ``litellm_cache_key`` field cannot be
|
||||
# safely reassigned to a caller's scope and are treated as misses.
|
||||
cached_key = cache_hit.get(self.CACHE_KEY_FIELD_NAME)
|
||||
|
|
@ -297,7 +298,7 @@ class RedisSemanticCache(BaseCache):
|
|||
return
|
||||
|
||||
@staticmethod
|
||||
def _coerce_response_input_value(value: Any) -> Any:
|
||||
def _coerce_response_input_value(value: object) -> object:
|
||||
model_dump: Final = getattr(value, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
return model_dump()
|
||||
|
|
@ -340,7 +341,7 @@ class RedisSemanticCache(BaseCache):
|
|||
)
|
||||
return embedding_response["data"][0]["embedding"]
|
||||
|
||||
def _get_cache_logic(self, cached_response: Any) -> Any:
|
||||
def _get_cache_logic(self, cached_response: Any) -> object:
|
||||
"""
|
||||
Process the cached response to prepare it for use.
|
||||
|
||||
|
|
@ -369,7 +370,7 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
return cached_response
|
||||
|
||||
def set_cache(self, key: str, value: Any, **kwargs) -> None:
|
||||
def set_cache(self, key: str, value: object, **kwargs) -> None:
|
||||
"""
|
||||
Store a value in the semantic cache.
|
||||
|
||||
|
|
@ -405,7 +406,7 @@ class RedisSemanticCache(BaseCache):
|
|||
except Exception as e:
|
||||
print_verbose(f"Error setting {value_str or value} in the Redis semantic cache: {e}")
|
||||
|
||||
def get_cache(self, key: str, **kwargs) -> Any:
|
||||
def get_cache(self, key: str, **kwargs) -> object:
|
||||
"""
|
||||
Retrieve a semantically similar cached response.
|
||||
|
||||
|
|
@ -428,7 +429,7 @@ class RedisSemanticCache(BaseCache):
|
|||
# Check the cache for semantically similar prompts in this exact
|
||||
# LiteLLM cache-key scope.
|
||||
prompt_embedding: Final = self._get_embedding(prompt, metadata=kwargs.get("metadata"))
|
||||
check_kwargs: Final[dict[str, Any]] = {
|
||||
check_kwargs: Final[Mapping[str, object]] = {
|
||||
"prompt": prompt,
|
||||
"vector": prompt_embedding,
|
||||
"filter_expression": self._get_cache_key_filter_expression(key),
|
||||
|
|
@ -508,7 +509,7 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Error generating async embedding: {e}")
|
||||
raise ValueError(f"Failed to generate embedding: {e}") from e
|
||||
|
||||
async def async_set_cache(self, key: str, value: Any, **kwargs) -> None:
|
||||
async def async_set_cache(self, key: str, value: object, **kwargs) -> None:
|
||||
"""
|
||||
Asynchronously store a value in the semantic cache.
|
||||
|
||||
|
|
@ -548,7 +549,7 @@ class RedisSemanticCache(BaseCache):
|
|||
except Exception as e:
|
||||
print_verbose(f"Error in async_set_cache: {e}")
|
||||
|
||||
async def async_get_cache(self, key: str, **kwargs) -> Any:
|
||||
async def async_get_cache(self, key: str, **kwargs) -> object:
|
||||
"""
|
||||
Asynchronously retrieve a semantically similar cached response.
|
||||
|
||||
|
|
@ -573,7 +574,7 @@ class RedisSemanticCache(BaseCache):
|
|||
|
||||
# Check the cache for semantically similar prompts in this exact
|
||||
# LiteLLM cache-key scope.
|
||||
check_kwargs: Final[dict[str, Any]] = {
|
||||
check_kwargs: Final[Mapping[str, object]] = {
|
||||
"prompt": prompt,
|
||||
"vector": prompt_embedding,
|
||||
"filter_expression": self._get_cache_key_filter_expression(key),
|
||||
|
|
@ -615,7 +616,7 @@ class RedisSemanticCache(BaseCache):
|
|||
print_verbose(f"Error in async_get_cache: {e}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
||||
async def _index_info(self) -> dict[str, Any]:
|
||||
async def _index_info(self) -> Mapping[str, object]:
|
||||
"""
|
||||
Get information about the Redis index.
|
||||
|
||||
|
|
@ -625,7 +626,7 @@ class RedisSemanticCache(BaseCache):
|
|||
aindex: Final = await self.llmcache._get_async_index()
|
||||
return await aindex.info()
|
||||
|
||||
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs) -> None:
|
||||
async def async_set_cache_pipeline(self, cache_list: list[tuple[str, Any]], **kwargs: object) -> None:
|
||||
"""
|
||||
Asynchronously store multiple values in the semantic cache.
|
||||
|
||||
|
|
|
|||
|
|
@ -4,8 +4,9 @@ import json
|
|||
import os
|
||||
import re
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Final, cast
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timezone, tzinfo
|
||||
from typing import Any, Final, TypedDict, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, Field
|
||||
|
|
@ -34,6 +35,17 @@ GALILEO_CLOUD_API_BASE_URL: Final = "https://api.galileo.ai"
|
|||
GALILEO_MAX_IN_MEMORY_RECORDS: Final = 1000
|
||||
|
||||
|
||||
class GalileoStandardLoggingFields(TypedDict, total=False):
|
||||
call_type: str
|
||||
model: str
|
||||
prompt_tokens: int
|
||||
completion_tokens: int
|
||||
total_tokens: int
|
||||
response_cost: float
|
||||
startTime: float
|
||||
endTime: float
|
||||
|
||||
|
||||
class LLMResponse(BaseModel):
|
||||
latency_ms: int
|
||||
status_code: int
|
||||
|
|
@ -59,7 +71,7 @@ class LLMResponse(BaseModel):
|
|||
|
||||
class GalileoObserve(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
self.in_memory_records: list[dict] = []
|
||||
self.in_memory_records: list[Mapping[str, object]] = []
|
||||
self.batch_size = 1
|
||||
self.api_key = os.getenv("GALILEO_API_KEY")
|
||||
self.project_id = os.getenv("GALILEO_PROJECT_ID")
|
||||
|
|
@ -176,7 +188,7 @@ class GalileoObserve(CustomLogger):
|
|||
return False
|
||||
|
||||
@staticmethod
|
||||
def _galileo_input_messages(messages: Any | None, input_text: str) -> list[dict[str, str]]:
|
||||
def _galileo_input_messages(messages: object, input_text: str) -> list[dict[str, str]]:
|
||||
if isinstance(messages, dict):
|
||||
messages = messages.get("messages")
|
||||
if not messages:
|
||||
|
|
@ -203,11 +215,11 @@ class GalileoObserve(CustomLogger):
|
|||
return [{"role": "user", "content": input_text}]
|
||||
|
||||
@staticmethod
|
||||
def _local_timezone():
|
||||
def _local_timezone() -> tzinfo:
|
||||
return datetime.now().astimezone().tzinfo or timezone.utc
|
||||
|
||||
@staticmethod
|
||||
def _format_created_at(dt: datetime | Any) -> str:
|
||||
def _format_created_at(dt: object) -> str:
|
||||
"""Serialize timestamps as UTC ISO-8601 for Galileo."""
|
||||
if not isinstance(dt, datetime):
|
||||
return str(dt)
|
||||
|
|
@ -226,7 +238,7 @@ class GalileoObserve(CustomLogger):
|
|||
return created_at
|
||||
|
||||
@staticmethod
|
||||
def _token_metrics_from_record(record: dict[str, Any]) -> dict[str, Any]:
|
||||
def _token_metrics_from_record(record: Mapping[str, Any]) -> dict[str, Any]:
|
||||
num_input_tokens: Final = int(record.get("num_input_tokens") or 0)
|
||||
num_output_tokens: Final = int(record.get("num_output_tokens") or 0)
|
||||
num_total_tokens = int(record.get("num_total_tokens") or 0)
|
||||
|
|
@ -244,7 +256,7 @@ class GalileoObserve(CustomLogger):
|
|||
|
||||
@staticmethod
|
||||
def _record_to_v2_span(
|
||||
record: dict[str, Any],
|
||||
record: Mapping[str, Any],
|
||||
*,
|
||||
trace_id: str,
|
||||
span_id: str,
|
||||
|
|
@ -275,7 +287,7 @@ class GalileoObserve(CustomLogger):
|
|||
return span
|
||||
|
||||
@staticmethod
|
||||
def _record_to_v2_trace(record: dict[str, Any]) -> dict[str, Any]:
|
||||
def _record_to_v2_trace(record: Mapping[str, Any]) -> dict[str, Any]:
|
||||
trace_id: Final = str(uuid.uuid4())
|
||||
span_id: Final = str(uuid.uuid4())
|
||||
created_at: Final = GalileoObserve._normalize_created_at(record.get("created_at", ""))
|
||||
|
|
@ -295,7 +307,7 @@ class GalileoObserve(CustomLogger):
|
|||
"spans": [GalileoObserve._record_to_v2_span(record, trace_id=trace_id, span_id=span_id)],
|
||||
}
|
||||
|
||||
def _build_traces_payload(self, records: list[dict]) -> dict[str, Any]:
|
||||
def _build_traces_payload(self, records: Sequence[Mapping[str, Any]]) -> dict[str, Any]:
|
||||
payload: Final[dict[str, Any]] = {
|
||||
"traces": [self._record_to_v2_trace(record) for record in records],
|
||||
"logging_method": "api_direct",
|
||||
|
|
@ -357,7 +369,7 @@ class GalileoObserve(CustomLogger):
|
|||
@staticmethod
|
||||
def _log_v2_payload_validation(payload: dict[str, Any]) -> None:
|
||||
missing_fields: Final[list[str]] = []
|
||||
traces: Final = payload.get("traces", [])
|
||||
traces: Final[Sequence[object]] = payload.get("traces", [])
|
||||
if not traces:
|
||||
missing_fields.append("traces")
|
||||
|
||||
|
|
@ -385,7 +397,7 @@ class GalileoObserve(CustomLogger):
|
|||
)
|
||||
|
||||
def _log_flush_payload(self, url: str, payload: dict[str, Any]) -> None:
|
||||
traces: Final = payload.get("traces", [])
|
||||
traces: Final[Sequence[object]] = payload.get("traces", [])
|
||||
verbose_logger.debug(
|
||||
"Galileo Logger flush URL: %s trace_count=%s",
|
||||
url,
|
||||
|
|
@ -415,8 +427,8 @@ class GalileoObserve(CustomLogger):
|
|||
pass
|
||||
|
||||
@staticmethod
|
||||
def _build_prompt(kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
optional_params: Final = kwargs.get("optional_params", {}) or {}
|
||||
def _build_prompt(kwargs: Mapping[str, Any]) -> dict[str, Any]:
|
||||
optional_params: Final[Mapping[str, object]] = kwargs.get("optional_params", {}) or {}
|
||||
prompt: Final[dict[str, Any]] = {"messages": kwargs.get("messages")}
|
||||
if optional_params.get("functions") is not None:
|
||||
prompt["functions"] = optional_params["functions"]
|
||||
|
|
@ -425,13 +437,13 @@ class GalileoObserve(CustomLogger):
|
|||
return prompt
|
||||
|
||||
@staticmethod
|
||||
def _serialize_galileo_output(value: Any) -> str:
|
||||
def _serialize_galileo_output(value: object) -> str:
|
||||
if value is None:
|
||||
return ""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
|
||||
def _json_default(obj: Any) -> Any:
|
||||
def _json_default(obj: Any) -> object:
|
||||
if hasattr(obj, "model_dump"):
|
||||
return obj.model_dump()
|
||||
return str(obj)
|
||||
|
|
@ -439,8 +451,8 @@ class GalileoObserve(CustomLogger):
|
|||
return json.dumps(value, default=_json_default)
|
||||
|
||||
@staticmethod
|
||||
def _prompt_to_input_text(prompt: dict[str, Any]) -> str:
|
||||
messages: Final = prompt.get("messages")
|
||||
def _prompt_to_input_text(prompt: Mapping[str, Any]) -> str:
|
||||
messages: Final[object] = prompt.get("messages")
|
||||
if messages is not None:
|
||||
text: Final = GalileoObserve._input_text_from_messages(messages)
|
||||
if text:
|
||||
|
|
@ -448,7 +460,7 @@ class GalileoObserve(CustomLogger):
|
|||
return json.dumps(prompt, default=str)
|
||||
|
||||
@staticmethod
|
||||
def _get_chat_content_for_galileo(response_obj: litellm.ModelResponse) -> Any:
|
||||
def _get_chat_content_for_galileo(response_obj: litellm.ModelResponse) -> object:
|
||||
if response_obj.choices and len(response_obj.choices) > 0:
|
||||
message: Final = response_obj["choices"][0]["message"]
|
||||
if hasattr(message, "json"):
|
||||
|
|
@ -470,23 +482,23 @@ class GalileoObserve(CustomLogger):
|
|||
@staticmethod
|
||||
def _get_responses_api_content_for_galileo(
|
||||
response_obj: ResponsesAPIResponse,
|
||||
) -> Any:
|
||||
) -> object:
|
||||
if hasattr(response_obj, "output") and response_obj.output:
|
||||
return response_obj.output
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _langfuse_style_rerank_prompt(kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
def _langfuse_style_rerank_prompt(kwargs: Mapping[str, object]) -> dict[str, Any]:
|
||||
"""Match Langfuse rerank input: prompt = {"messages": kwargs.get("messages")}."""
|
||||
return {"messages": kwargs.get("messages")}
|
||||
|
||||
def _get_galileo_input_output_content(
|
||||
self,
|
||||
kwargs: dict[str, Any],
|
||||
response_obj: Any,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
level: str = "DEFAULT",
|
||||
status_message: str | None = None,
|
||||
) -> tuple[str, str, Any]:
|
||||
) -> tuple[str, str, object]:
|
||||
"""
|
||||
Mirror Langfuse _get_langfuse_input_output_content for Galileo ingest.
|
||||
|
||||
|
|
@ -582,12 +594,12 @@ class GalileoObserve(CustomLogger):
|
|||
|
||||
return self._prompt_to_input_text(prompt), "", kwargs.get("messages") or []
|
||||
|
||||
def get_output_str_from_response(self, response_obj: Any, kwargs: dict[str, Any]) -> str:
|
||||
def get_output_str_from_response(self, response_obj: object, kwargs: Mapping[str, object]) -> str:
|
||||
_, output_text, _ = self._get_galileo_input_output_content(kwargs=kwargs, response_obj=response_obj)
|
||||
return output_text
|
||||
|
||||
@staticmethod
|
||||
def _input_text_from_messages(messages: Any) -> str:
|
||||
def _input_text_from_messages(messages: object) -> str:
|
||||
"""Return a plain-string summary of the input suitable for the trace-level input field."""
|
||||
if isinstance(messages, str):
|
||||
return messages
|
||||
|
|
@ -613,7 +625,13 @@ class GalileoObserve(CustomLogger):
|
|||
return str(content)
|
||||
return ""
|
||||
|
||||
async def async_log_success_event(self, kwargs: Any, response_obj: Any, start_time: Any, end_time: Any):
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
verbose_logger.debug("On Async Success")
|
||||
try:
|
||||
await self._async_log_success_event_impl(
|
||||
|
|
@ -625,7 +643,13 @@ class GalileoObserve(CustomLogger):
|
|||
except Exception:
|
||||
verbose_logger.exception("Galileo Logger: unexpected error in async_log_success_event")
|
||||
|
||||
async def _async_log_success_event_impl(self, kwargs: Any, response_obj: Any, start_time: Any, end_time: Any):
|
||||
async def _async_log_success_event_impl(
|
||||
self,
|
||||
kwargs: Mapping[str, Any],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
if not self._is_configured():
|
||||
verbose_logger.debug(
|
||||
"Galileo Logger: skipping — GALILEO_PROJECT_ID=%s GALILEO_API_KEY=%s GALILEO_BASE_URL=%s",
|
||||
|
|
@ -635,7 +659,7 @@ class GalileoObserve(CustomLogger):
|
|||
)
|
||||
return
|
||||
|
||||
slo: Final[dict[str, Any] | None] = kwargs.get("standard_logging_object")
|
||||
slo: Final[GalileoStandardLoggingFields | None] = kwargs.get("standard_logging_object")
|
||||
if slo is None:
|
||||
verbose_logger.debug("Galileo Logger: no standard_logging_object in kwargs, skipping")
|
||||
return
|
||||
|
|
@ -646,8 +670,8 @@ class GalileoObserve(CustomLogger):
|
|||
kwargs=kwargs, response_obj=response_obj
|
||||
)
|
||||
|
||||
raw_start: Final = slo.get("startTime")
|
||||
raw_end: Final = slo.get("endTime")
|
||||
raw_start: Final[float | None] = slo.get("startTime")
|
||||
raw_end: Final[float | None] = slo.get("endTime")
|
||||
if raw_start is None or raw_end is None:
|
||||
verbose_logger.debug(
|
||||
"Galileo Logger: standard_logging_object missing startTime/endTime, "
|
||||
|
|
@ -710,7 +734,7 @@ class GalileoObserve(CustomLogger):
|
|||
if len(self.in_memory_records) >= self.batch_size:
|
||||
await self.flush_in_memory_records()
|
||||
|
||||
async def flush_in_memory_records(self):
|
||||
async def flush_in_memory_records(self) -> None:
|
||||
if not self.in_memory_records:
|
||||
return
|
||||
|
||||
|
|
@ -774,5 +798,11 @@ class GalileoObserve(CustomLogger):
|
|||
if not self.use_v2_api and response.status_code in (401, 403):
|
||||
self.headers = None
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
async def async_log_failure_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
verbose_logger.debug("On Async Failure")
|
||||
|
|
|
|||
|
|
@ -7,10 +7,13 @@ import contextvars
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Mapping, Sequence
|
||||
from pathlib import PurePosixPath
|
||||
from typing import Any, Final
|
||||
from typing import Any, Final, TypeAlias, TypedDict
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
|
||||
# Tool names emitted from OpenAPI specs must work across all major LLM providers.
|
||||
# OpenAI/Anthropic/Bedrock all enforce a character class roughly equivalent to
|
||||
# ^[a-zA-Z0-9_-]+$ on tool names. Many specs (notably GitHub's REST API) use
|
||||
|
|
@ -44,6 +47,41 @@ from litellm.proxy._experimental.mcp_server.tool_registry import (
|
|||
global_mcp_tool_registry,
|
||||
)
|
||||
|
||||
_OpenAPIParameter: TypeAlias = Mapping[str, Any]
|
||||
|
||||
|
||||
class _OpenAPIJSONSchema(TypedDict, total=False):
|
||||
properties: Mapping[str, object]
|
||||
|
||||
|
||||
class _OpenAPIMediaType(TypedDict, total=False):
|
||||
schema: _OpenAPIJSONSchema
|
||||
|
||||
|
||||
class _OpenAPIRequestBody(TypedDict, total=False):
|
||||
description: str
|
||||
required: bool
|
||||
content: Mapping[str, _OpenAPIMediaType]
|
||||
|
||||
|
||||
class _OpenAPIOperation(TypedDict, total=False):
|
||||
operationId: str
|
||||
summary: str
|
||||
description: str
|
||||
parameters: Sequence[_OpenAPIParameter]
|
||||
requestBody: _OpenAPIRequestBody
|
||||
|
||||
|
||||
class _OpenAPIPathItem(TypedDict, total=False):
|
||||
summary: str
|
||||
description: str
|
||||
parameters: Sequence[_OpenAPIParameter]
|
||||
|
||||
|
||||
class _OpenAPIComponents(TypedDict, total=False):
|
||||
parameters: Mapping[str, _OpenAPIParameter]
|
||||
|
||||
|
||||
# Store the base URL and headers globally
|
||||
BASE_URL: Final = ""
|
||||
HEADERS: Final[dict[str, str]] = {}
|
||||
|
|
@ -69,7 +107,7 @@ _request_resolved_auth_headers: Final[contextvars.ContextVar[dict[str, str] | No
|
|||
)
|
||||
|
||||
|
||||
def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str:
|
||||
def _sanitize_path_parameter_value(param_value: object, param_name: str) -> str:
|
||||
"""Ensure path params cannot introduce directory traversal."""
|
||||
if param_value is None:
|
||||
return ""
|
||||
|
|
@ -109,7 +147,7 @@ def load_openapi_spec(filepath: str) -> dict[str, Any]:
|
|||
async def load_openapi_spec_async(filepath: str) -> dict[str, Any]:
|
||||
if filepath.startswith("http://") or filepath.startswith("https://"):
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
|
||||
r: Final = await async_safe_get(client, filepath)
|
||||
r: Final[httpx.Response] = await async_safe_get(client, filepath)
|
||||
r.raise_for_status()
|
||||
return r.json()
|
||||
|
||||
|
|
@ -121,11 +159,11 @@ async def load_openapi_spec_async(filepath: str) -> dict[str, Any]:
|
|||
return json.load(f)
|
||||
|
||||
|
||||
def get_base_url(spec: dict[str, Any], spec_path: str | None = None) -> str:
|
||||
def get_base_url(spec: Mapping[str, Any], spec_path: str | None = None) -> str:
|
||||
"""Extract base URL from OpenAPI spec."""
|
||||
# OpenAPI 3.x
|
||||
if "servers" in spec and spec["servers"]:
|
||||
server_url: Final = spec["servers"][0]["url"]
|
||||
server_url: Final[str] = spec["servers"][0]["url"]
|
||||
|
||||
# If the server URL is relative (starts with /), derive base from spec_path
|
||||
if server_url.startswith("/") and spec_path:
|
||||
|
|
@ -147,8 +185,8 @@ def get_base_url(spec: dict[str, Any], spec_path: str | None = None) -> str:
|
|||
return server_url
|
||||
# OpenAPI 2.x (Swagger)
|
||||
elif "host" in spec:
|
||||
scheme: Final = spec.get("schemes", ["https"])[0]
|
||||
base_path: Final = spec.get("basePath", "")
|
||||
scheme: Final[str] = spec.get("schemes", ["https"])[0]
|
||||
base_path: Final[str] = spec.get("basePath", "")
|
||||
return f"{scheme}://{spec['host']}{base_path}"
|
||||
|
||||
# Fallback: derive base URL from spec_path if it's a URL
|
||||
|
|
@ -172,20 +210,24 @@ def get_base_url(spec: dict[str, Any], spec_path: str | None = None) -> str:
|
|||
return ""
|
||||
|
||||
|
||||
def _resolve_ref(param: dict[str, Any], component_params: dict[str, Any]) -> dict[str, Any] | None:
|
||||
def _resolve_ref(
|
||||
param: _OpenAPIParameter, component_params: Mapping[str, _OpenAPIParameter]
|
||||
) -> _OpenAPIParameter | None:
|
||||
"""Resolve a single parameter, following a $ref if present.
|
||||
|
||||
Returns the resolved param dict, or None if the $ref target is absent from
|
||||
components (so callers can skip/filter it rather than propagating a stub
|
||||
with name=None that would corrupt deduplication).
|
||||
"""
|
||||
ref: Final = param.get("$ref", "")
|
||||
ref: Final[str] = param.get("$ref", "")
|
||||
if not ref.startswith("#/components/parameters/"):
|
||||
return param
|
||||
return component_params.get(ref.split("/")[-1])
|
||||
|
||||
|
||||
def _resolve_param_list(raw: list[dict[str, Any]], component_params: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
def _resolve_param_list(
|
||||
raw: Sequence[_OpenAPIParameter], component_params: Mapping[str, _OpenAPIParameter]
|
||||
) -> list[_OpenAPIParameter]:
|
||||
"""Resolve $refs in a parameter list, dropping any unresolvable entries."""
|
||||
result: Final = []
|
||||
for p in raw:
|
||||
|
|
@ -196,9 +238,9 @@ def _resolve_param_list(raw: list[dict[str, Any]], component_params: dict[str, A
|
|||
|
||||
|
||||
def resolve_operation_params(
|
||||
operation: dict[str, Any],
|
||||
path_item: dict[str, Any],
|
||||
components: dict[str, Any],
|
||||
operation: _OpenAPIOperation,
|
||||
path_item: _OpenAPIPathItem,
|
||||
components: _OpenAPIComponents,
|
||||
) -> dict[str, Any]:
|
||||
"""Return a copy of *operation* with fully-resolved, merged parameters.
|
||||
|
||||
|
|
@ -214,7 +256,7 @@ def resolve_operation_params(
|
|||
merged with the operation-level params; operation-level wins when the
|
||||
same ``name`` + ``in`` combination appears in both.
|
||||
"""
|
||||
component_params: Final = components.get("parameters", {})
|
||||
component_params: Final[Mapping[str, _OpenAPIParameter]] = components.get("parameters", {})
|
||||
path_level: Final = _resolve_param_list(path_item.get("parameters", []), component_params)
|
||||
op_level: Final = _resolve_param_list(operation.get("parameters", []), component_params)
|
||||
op_keys: Final = {(p["name"], p.get("in")) for p in op_level}
|
||||
|
|
@ -224,7 +266,7 @@ def resolve_operation_params(
|
|||
return result
|
||||
|
||||
|
||||
def extract_parameters(operation: dict[str, Any]) -> tuple:
|
||||
def extract_parameters(operation: Mapping[str, Any]) -> tuple[Sequence[str], Sequence[str], Sequence[str]]:
|
||||
"""Extract parameter names from OpenAPI operation."""
|
||||
path_params: Final = []
|
||||
query_params: Final = []
|
||||
|
|
@ -250,7 +292,7 @@ def extract_parameters(operation: dict[str, Any]) -> tuple:
|
|||
return path_params, query_params, body_params
|
||||
|
||||
|
||||
def build_input_schema(operation: dict[str, Any]) -> dict[str, Any]:
|
||||
def build_input_schema(operation: Mapping[str, Any]) -> dict[str, Any]:
|
||||
"""Build MCP input schema from OpenAPI operation."""
|
||||
properties: Final = {}
|
||||
required: Final = []
|
||||
|
|
@ -274,12 +316,12 @@ def build_input_schema(operation: dict[str, Any]) -> dict[str, Any]:
|
|||
|
||||
# Process requestBody (OpenAPI 3.x)
|
||||
if "requestBody" in operation:
|
||||
request_body: Final = operation["requestBody"]
|
||||
content: Final = request_body.get("content", {})
|
||||
request_body: Final[_OpenAPIRequestBody] = operation["requestBody"]
|
||||
content: Final[Mapping[str, _OpenAPIMediaType]] = request_body.get("content", {})
|
||||
|
||||
# Try to get JSON schema
|
||||
if "application/json" in content:
|
||||
schema: Final = content["application/json"].get("schema", {})
|
||||
schema: Final[_OpenAPIJSONSchema] = content["application/json"].get("schema", {})
|
||||
properties["body"] = {
|
||||
"type": "object",
|
||||
"description": request_body.get("description", "Request body"),
|
||||
|
|
@ -347,7 +389,7 @@ def _merge_openapi_tool_request_headers(
|
|||
def create_tool_function(
|
||||
path: str,
|
||||
method: str,
|
||||
operation: dict[str, Any],
|
||||
operation: Mapping[str, Any],
|
||||
base_url: str,
|
||||
headers: dict[str, str] | None = None,
|
||||
):
|
||||
|
|
@ -373,7 +415,7 @@ def create_tool_function(
|
|||
path_params, query_params, body_params = extract_parameters(operation)
|
||||
original_method: Final = method.lower()
|
||||
|
||||
async def tool_function(**kwargs: Any) -> str:
|
||||
async def tool_function(**kwargs: object) -> str:
|
||||
"""
|
||||
Dynamically generated tool function.
|
||||
|
||||
|
|
@ -448,10 +490,10 @@ def create_tool_function(
|
|||
return tool_function
|
||||
|
||||
|
||||
def register_tools_from_openapi(spec: dict[str, Any], base_url: str):
|
||||
def register_tools_from_openapi(spec: Mapping[str, Any], base_url: str) -> None:
|
||||
"""Register MCP tools from OpenAPI specification."""
|
||||
paths: Final = spec.get("paths", {})
|
||||
used_names: Final[set] = set()
|
||||
paths: Final[Mapping[str, Mapping[str, Any]]] = spec.get("paths", {})
|
||||
used_names: Final = set()
|
||||
|
||||
for path, path_item in paths.items():
|
||||
for method in ["get", "post", "put", "delete", "patch"]:
|
||||
|
|
|
|||
|
|
@ -18,8 +18,9 @@ Endpoints:
|
|||
|
||||
import json
|
||||
import re
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Final
|
||||
from typing import Final, Protocol, TypedDict
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi.responses import JSONResponse
|
||||
|
|
@ -41,7 +42,30 @@ from litellm.types.proxy.claude_code_endpoints import (
|
|||
router: Final = APIRouter()
|
||||
|
||||
|
||||
async def _get_prisma_client():
|
||||
class _PluginRecord(Protocol):
|
||||
id: str
|
||||
name: str
|
||||
version: str | None
|
||||
description: str | None
|
||||
manifest_json: str | None
|
||||
enabled: bool
|
||||
created_at: datetime | None
|
||||
updated_at: datetime | None
|
||||
created_by: str | None
|
||||
|
||||
|
||||
class _MarketplaceEntry(TypedDict, total=False):
|
||||
name: str
|
||||
source: object
|
||||
version: str
|
||||
description: str
|
||||
author: object
|
||||
homepage: object
|
||||
keywords: object
|
||||
category: object
|
||||
|
||||
|
||||
async def _get_prisma_client() -> object:
|
||||
"""Get the prisma client from proxy_server."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
|
|
@ -77,12 +101,14 @@ async def get_marketplace():
|
|||
try:
|
||||
prisma_client: Final = await _get_prisma_client()
|
||||
|
||||
plugins: Final = await ClaudeCodePluginRepository(prisma_client).table.find_many(where={"enabled": True})
|
||||
plugins: Final[Sequence[_PluginRecord]] = await ClaudeCodePluginRepository(prisma_client).table.find_many(
|
||||
where={"enabled": True}
|
||||
)
|
||||
|
||||
plugin_list: Final = []
|
||||
for plugin in plugins:
|
||||
try:
|
||||
manifest = json.loads(plugin.manifest_json)
|
||||
manifest: Mapping[str, object] = json.loads(plugin.manifest_json or "{}")
|
||||
except json.JSONDecodeError:
|
||||
verbose_proxy_logger.warning("Plugin %s has invalid manifest JSON, skipping", plugin.name)
|
||||
continue
|
||||
|
|
@ -92,7 +118,7 @@ async def get_marketplace():
|
|||
verbose_proxy_logger.warning("Plugin %s has no source field, skipping", plugin.name)
|
||||
continue
|
||||
|
||||
entry: dict[str, Any] = {
|
||||
entry: _MarketplaceEntry = {
|
||||
"name": plugin.name,
|
||||
"source": manifest["source"],
|
||||
}
|
||||
|
|
@ -137,7 +163,7 @@ async def get_marketplace():
|
|||
_VALID_GIT_SUBDIR_PATH_RE: Final = re.compile(r"^[a-zA-Z0-9][a-zA-Z0-9._-]*(/[a-zA-Z0-9][a-zA-Z0-9._-]*)*$")
|
||||
|
||||
|
||||
def _validate_plugin_source(source: dict[str, Any]) -> None:
|
||||
def _validate_plugin_source(source: Mapping[str, str]) -> None:
|
||||
"""Validate plugin source format, raising HTTPException on invalid input."""
|
||||
source_type: Final = source.get("source")
|
||||
if source_type == "github":
|
||||
|
|
@ -179,9 +205,9 @@ def _validate_plugin_source(source: dict[str, Any]) -> None:
|
|||
)
|
||||
|
||||
|
||||
def _build_plugin_manifest(name: str, spec: PluginSpec) -> dict[str, Any]:
|
||||
def _build_plugin_manifest(name: str, spec: PluginSpec) -> Mapping[str, object]:
|
||||
"""Build the stored manifest dict shared by plugin create and update."""
|
||||
dumped = spec.model_dump(exclude_none=True)
|
||||
dumped: Final[Mapping[str, object]] = spec.model_dump(exclude_none=True)
|
||||
return {"name": name, **{key: value for key, value in dumped.items() if value and key != "name"}}
|
||||
|
||||
|
||||
|
|
@ -255,14 +281,16 @@ async def register_plugin(
|
|||
|
||||
_validate_plugin_source(request.source)
|
||||
|
||||
existing = await ClaudeCodePluginRepository(prisma_client).table.find_unique(where={"name": request.name})
|
||||
existing: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
|
||||
where={"name": request.name}
|
||||
)
|
||||
if existing:
|
||||
raise _name_conflict_error(request.name)
|
||||
|
||||
manifest = _build_plugin_manifest(request.name, request)
|
||||
manifest: Final[Mapping[str, object]] = _build_plugin_manifest(request.name, request)
|
||||
|
||||
try:
|
||||
plugin = await ClaudeCodePluginRepository(prisma_client).table.create(
|
||||
plugin: Final[_PluginRecord] = await ClaudeCodePluginRepository(prisma_client).table.create(
|
||||
data={
|
||||
"name": request.name,
|
||||
"version": request.version,
|
||||
|
|
@ -326,7 +354,9 @@ async def list_plugins(
|
|||
prisma_client: Final = await _get_prisma_client()
|
||||
|
||||
where: Final = {"enabled": True} if enabled_only else {}
|
||||
plugins: Final = await ClaudeCodePluginRepository(prisma_client).table.find_many(where=where)
|
||||
plugins: Final[Sequence[_PluginRecord]] = await ClaudeCodePluginRepository(prisma_client).table.find_many(
|
||||
where=where
|
||||
)
|
||||
|
||||
plugin_list: Final = []
|
||||
for p in plugins:
|
||||
|
|
@ -391,7 +421,9 @@ async def get_plugin(
|
|||
try:
|
||||
prisma_client: Final = await _get_prisma_client()
|
||||
|
||||
plugin: Final = await ClaudeCodePluginRepository(prisma_client).table.find_unique(where={"name": plugin_name})
|
||||
plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
|
||||
where={"name": plugin_name}
|
||||
)
|
||||
|
||||
if not plugin:
|
||||
raise HTTPException(
|
||||
|
|
@ -399,7 +431,7 @@ async def get_plugin(
|
|||
detail={"error": f"Plugin '{plugin_name}' not found"},
|
||||
)
|
||||
|
||||
manifest: Final = json.loads(plugin.manifest_json) if plugin.manifest_json else {}
|
||||
manifest: Final[Mapping[str, object]] = json.loads(plugin.manifest_json or "{}") if plugin.manifest_json else {}
|
||||
|
||||
return {
|
||||
"id": plugin.id,
|
||||
|
|
@ -477,19 +509,19 @@ async def update_plugin(
|
|||
from prisma.errors import PrismaError
|
||||
|
||||
try:
|
||||
prisma_client = await _get_prisma_client()
|
||||
prisma_client: Final = await _get_prisma_client()
|
||||
|
||||
_validate_plugin_source(request.source)
|
||||
|
||||
existing = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
|
||||
existing: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
|
||||
where={"name": plugin_name} # mutable-ok: prisma query arguments must be plain dicts
|
||||
)
|
||||
if not existing:
|
||||
raise _error_response(404, f"Plugin '{plugin_name}' not found")
|
||||
|
||||
manifest = _build_plugin_manifest(plugin_name, request)
|
||||
manifest: Final[Mapping[str, object]] = _build_plugin_manifest(plugin_name, request)
|
||||
|
||||
plugin = await ClaudeCodePluginRepository(prisma_client).table.update(
|
||||
plugin: Final[_PluginRecord] = await ClaudeCodePluginRepository(prisma_client).table.update(
|
||||
where={"name": plugin_name}, # mutable-ok: prisma query arguments must be plain dicts
|
||||
data={ # mutable-ok: prisma query arguments must be plain dicts
|
||||
"version": request.version,
|
||||
|
|
@ -540,7 +572,9 @@ async def enable_plugin(
|
|||
try:
|
||||
prisma_client: Final = await _get_prisma_client()
|
||||
|
||||
plugin: Final = await ClaudeCodePluginRepository(prisma_client).table.find_unique(where={"name": plugin_name})
|
||||
plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
|
||||
where={"name": plugin_name}
|
||||
)
|
||||
if not plugin:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
|
|
@ -583,7 +617,9 @@ async def disable_plugin(
|
|||
try:
|
||||
prisma_client: Final = await _get_prisma_client()
|
||||
|
||||
plugin: Final = await ClaudeCodePluginRepository(prisma_client).table.find_unique(where={"name": plugin_name})
|
||||
plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
|
||||
where={"name": plugin_name}
|
||||
)
|
||||
if not plugin:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
|
|
@ -626,7 +662,9 @@ async def delete_plugin(
|
|||
try:
|
||||
prisma_client: Final = await _get_prisma_client()
|
||||
|
||||
plugin: Final = await ClaudeCodePluginRepository(prisma_client).table.find_unique(where={"name": plugin_name})
|
||||
plugin: Final[_PluginRecord | None] = await ClaudeCodePluginRepository(prisma_client).table.find_unique(
|
||||
where={"name": plugin_name}
|
||||
)
|
||||
if not plugin:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -26,7 +27,7 @@ class CustomOpenAPISpec:
|
|||
RESPONSES_API_PATHS = ["/v1/responses", "/responses"]
|
||||
|
||||
@staticmethod
|
||||
def get_pydantic_schema(model_class) -> dict[str, Any] | None:
|
||||
def get_pydantic_schema(model_class) -> Mapping[str, object] | None:
|
||||
"""
|
||||
Get JSON schema from a Pydantic model, handling both v1 and v2 APIs.
|
||||
|
||||
|
|
@ -53,7 +54,9 @@ class CustomOpenAPISpec:
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
def add_schema_to_components(openapi_schema: dict[str, Any], schema_name: str, schema_def: dict[str, Any]) -> None:
|
||||
def add_schema_to_components(
|
||||
openapi_schema: dict[str, Any], schema_name: str, schema_def: Mapping[str, object]
|
||||
) -> None:
|
||||
"""
|
||||
Add a schema definition to the OpenAPI components/schemas section.
|
||||
|
||||
|
|
@ -72,7 +75,7 @@ class CustomOpenAPISpec:
|
|||
CustomOpenAPISpec._move_defs_to_components(openapi_schema, {schema_name: schema_def})
|
||||
|
||||
@staticmethod
|
||||
def add_request_body_to_paths(openapi_schema: dict[str, Any], paths: list[str], schema_ref: str) -> None:
|
||||
def add_request_body_to_paths(openapi_schema: dict[str, Any], paths: Sequence[str], schema_ref: str) -> None:
|
||||
"""
|
||||
Add request body with expanded form fields for better Swagger UI display.
|
||||
This keeps the request body but expands it to show individual fields in the UI.
|
||||
|
|
@ -130,7 +133,7 @@ class CustomOpenAPISpec:
|
|||
openapi_schema["paths"][path]["post"]["parameters"] = filtered_params
|
||||
|
||||
@staticmethod
|
||||
def _move_defs_to_components(openapi_schema: dict[str, Any], defs: dict[str, Any]) -> None:
|
||||
def _move_defs_to_components(openapi_schema: dict[str, Any], defs: Mapping[str, Mapping[str, Any]]) -> None:
|
||||
"""
|
||||
Move $defs from Pydantic v2 schema to OpenAPI components/schemas.
|
||||
This makes the definitions resolvable in Swagger/OpenAPI viewers.
|
||||
|
|
@ -218,7 +221,7 @@ class CustomOpenAPISpec:
|
|||
return {"type": "string"}
|
||||
|
||||
@staticmethod
|
||||
def _expand_field_definition(field_def: dict[str, Any]) -> dict[str, Any]:
|
||||
def _expand_field_definition(field_def: dict[str, object]) -> dict[str, object]:
|
||||
"""
|
||||
Expand a Pydantic field definition for inline use in OpenAPI schema.
|
||||
This creates a full field definition that Swagger UI can render as individual form fields.
|
||||
|
|
@ -234,12 +237,12 @@ class CustomOpenAPISpec:
|
|||
|
||||
@staticmethod
|
||||
def add_request_schema(
|
||||
openapi_schema: dict[str, Any],
|
||||
openapi_schema: dict[str, object],
|
||||
model_class: type,
|
||||
schema_name: str,
|
||||
paths: list[str],
|
||||
paths: Sequence[str],
|
||||
operation_name: str,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Generic method to add a request schema to OpenAPI specification.
|
||||
|
||||
|
|
@ -279,8 +282,8 @@ class CustomOpenAPISpec:
|
|||
|
||||
@staticmethod
|
||||
def add_chat_completion_request_schema(
|
||||
openapi_schema: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
openapi_schema: dict[str, object],
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Add ProxyChatCompletionRequest schema to chat completion endpoints for documentation.
|
||||
This shows the request body in Swagger without runtime validation.
|
||||
|
|
@ -306,7 +309,7 @@ class CustomOpenAPISpec:
|
|||
return openapi_schema
|
||||
|
||||
@staticmethod
|
||||
def add_embedding_request_schema(openapi_schema: dict[str, Any]) -> dict[str, Any]:
|
||||
def add_embedding_request_schema(openapi_schema: dict[str, object]) -> dict[str, object]:
|
||||
"""
|
||||
Add EmbeddingRequest schema to embedding endpoints for documentation.
|
||||
This shows the request body in Swagger without runtime validation.
|
||||
|
|
@ -333,8 +336,8 @@ class CustomOpenAPISpec:
|
|||
|
||||
@staticmethod
|
||||
def add_responses_api_request_schema(
|
||||
openapi_schema: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
openapi_schema: dict[str, object],
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Add ResponsesAPIRequestParams schema to responses API endpoints for documentation.
|
||||
This shows the request body in Swagger without runtime validation.
|
||||
|
|
@ -361,8 +364,8 @@ class CustomOpenAPISpec:
|
|||
|
||||
@staticmethod
|
||||
def add_llm_api_request_schema_body(
|
||||
openapi_schema: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
openapi_schema: dict[str, object],
|
||||
) -> dict[str, object]:
|
||||
"""
|
||||
Add LLM API request schema bodies to OpenAPI specification for documentation.
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
import json
|
||||
from typing import Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Set as AbstractSet
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -15,6 +17,29 @@ from litellm.proxy.common_utils.resource_ownership import (
|
|||
from litellm.repositories.table_repositories import ManagedObjectRepository
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
class _ManagedObjectRow(Protocol):
|
||||
model_object_id: str
|
||||
unified_object_id: str | None
|
||||
file_purpose: str | None
|
||||
created_by: str | None
|
||||
|
||||
|
||||
class _ManagedObjectTable(Protocol):
|
||||
async def find_unique(self, *, where: Mapping[str, str]) -> _ManagedObjectRow | None: ...
|
||||
|
||||
async def find_first(self, *, where: Mapping[str, str]) -> _ManagedObjectRow | None: ...
|
||||
|
||||
async def find_many(self, *, where: Mapping[str, object]) -> Sequence[_ManagedObjectRow]: ...
|
||||
|
||||
async def create(self, *, data: Mapping[str, str]) -> _ManagedObjectRow: ...
|
||||
|
||||
async def update(self, *, where: Mapping[str, str], data: Mapping[str, str]) -> _ManagedObjectRow | None: ...
|
||||
|
||||
|
||||
CONTAINER_OBJECT_PURPOSE: Final = "container"
|
||||
|
||||
# 60s LRU/TTL cache absorbs every container access check before it reaches
|
||||
|
|
@ -39,7 +64,7 @@ _CONTAINER_STORED_ID_CACHE: Final = InMemoryCache(max_size_in_memory=10000, defa
|
|||
_ALLOWED_CONTAINER_IDS_CACHE: Final = InMemoryCache(max_size_in_memory=2048, default_ttl=60)
|
||||
|
||||
|
||||
def _allowed_container_ids_cache_key(owner_scopes: list[str]) -> str:
|
||||
def _allowed_container_ids_cache_key(owner_scopes: Sequence[str]) -> str:
|
||||
"""JSON-encode the sorted scope list — using a separator like ``|``
|
||||
would collide for any tenant whose user_id / team_id / org_id /
|
||||
api_key happens to contain the separator. JSON quoting escapes
|
||||
|
|
@ -86,7 +111,7 @@ async def get_container_forwarding_params(
|
|||
return params
|
||||
|
||||
|
||||
def _get_response_id(response: Any) -> str | None:
|
||||
def _get_response_id(response: object) -> str | None:
|
||||
if response is None:
|
||||
return None
|
||||
if isinstance(response, dict):
|
||||
|
|
@ -96,7 +121,7 @@ def _get_response_id(response: Any) -> str | None:
|
|||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
def _dump_response(response: Any) -> dict[str, Any]:
|
||||
def _dump_response(response: Any) -> dict[str, object]:
|
||||
if isinstance(response, dict):
|
||||
return dict(response)
|
||||
if hasattr(response, "model_dump"):
|
||||
|
|
@ -106,17 +131,17 @@ def _dump_response(response: Any) -> dict[str, Any]:
|
|||
return {"id": _get_response_id(response)}
|
||||
|
||||
|
||||
async def _get_prisma_client():
|
||||
async def _get_prisma_client() -> "PrismaClient | None":
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
return prisma_client
|
||||
|
||||
|
||||
def _custom_llm_provider_from_responses_response(
|
||||
response: Any,
|
||||
response: object,
|
||||
default: str = "openai",
|
||||
) -> str:
|
||||
hidden_params: dict[str, Any] = {}
|
||||
hidden_params: Mapping[str, object] = {}
|
||||
if isinstance(response, dict):
|
||||
hidden_params = response.get("_hidden_params") or {}
|
||||
else:
|
||||
|
|
@ -129,7 +154,7 @@ def _custom_llm_provider_from_responses_response(
|
|||
|
||||
|
||||
async def record_container_owners_from_responses_response(
|
||||
response: Any,
|
||||
response: object,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
custom_llm_provider: str | None = None,
|
||||
) -> None:
|
||||
|
|
@ -160,10 +185,10 @@ async def record_container_owners_from_responses_response(
|
|||
|
||||
|
||||
async def record_container_owner(
|
||||
response: Any,
|
||||
response: object,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
custom_llm_provider: str,
|
||||
) -> Any:
|
||||
) -> object:
|
||||
container_id: Final = _get_response_id(response)
|
||||
if container_id is None:
|
||||
verbose_proxy_logger.warning("Skipping container ownership tracking because provider response has no id")
|
||||
|
|
@ -195,7 +220,7 @@ async def record_container_owner(
|
|||
verbose_proxy_logger.warning("Skipping container ownership tracking because prisma_client is None")
|
||||
return response
|
||||
|
||||
table: Final = ManagedObjectRepository(prisma_client).table
|
||||
table: Final[_ManagedObjectTable] = ManagedObjectRepository(prisma_client).table
|
||||
existing: Final = await table.find_unique(where={"model_object_id": model_object_id})
|
||||
if existing is not None:
|
||||
if getattr(existing, "file_purpose", None) != CONTAINER_OBJECT_PURPOSE:
|
||||
|
|
@ -247,15 +272,16 @@ async def _get_container_owner(original_container_id: str, custom_llm_provider:
|
|||
if prisma_client is None:
|
||||
return None
|
||||
|
||||
row: Final = await ManagedObjectRepository(prisma_client).table.find_first(
|
||||
table: Final[_ManagedObjectTable] = ManagedObjectRepository(prisma_client).table
|
||||
row: Final[_ManagedObjectRow | None] = await table.find_first(
|
||||
where={
|
||||
"model_object_id": model_object_id,
|
||||
"file_purpose": CONTAINER_OBJECT_PURPOSE,
|
||||
}
|
||||
)
|
||||
owner: Final = getattr(row, "created_by", None) if row is not None else None
|
||||
owner: Final[str | None] = getattr(row, "created_by", None) if row is not None else None
|
||||
_CONTAINER_OWNER_CACHE.set_cache(model_object_id, owner if owner is not None else _NEGATIVE_OWNER_SENTINEL)
|
||||
stored_id: Final = getattr(row, "unified_object_id", None) if row is not None else None
|
||||
stored_id: Final[str | None] = getattr(row, "unified_object_id", None) if row is not None else None
|
||||
_CONTAINER_STORED_ID_CACHE.set_cache(
|
||||
model_object_id,
|
||||
(stored_id if isinstance(stored_id, str) and stored_id else _NEGATIVE_STORED_ID_SENTINEL),
|
||||
|
|
@ -283,13 +309,14 @@ async def _get_stored_container_id(original_container_id: str, custom_llm_provid
|
|||
if prisma_client is None:
|
||||
return None
|
||||
|
||||
row: Final = await ManagedObjectRepository(prisma_client).table.find_first(
|
||||
table: Final[_ManagedObjectTable] = ManagedObjectRepository(prisma_client).table
|
||||
row: Final[_ManagedObjectRow | None] = await table.find_first(
|
||||
where={
|
||||
"model_object_id": model_object_id,
|
||||
"file_purpose": CONTAINER_OBJECT_PURPOSE,
|
||||
}
|
||||
)
|
||||
stored_id: Final = getattr(row, "unified_object_id", None) if row is not None else None
|
||||
stored_id: Final[str | None] = getattr(row, "unified_object_id", None) if row is not None else None
|
||||
_CONTAINER_STORED_ID_CACHE.set_cache(
|
||||
model_object_id,
|
||||
(stored_id if isinstance(stored_id, str) and stored_id else _NEGATIVE_STORED_ID_SENTINEL),
|
||||
|
|
@ -317,7 +344,7 @@ async def assert_user_can_access_container(
|
|||
return original_container_id, resolved_provider
|
||||
|
||||
|
||||
def _get_container_list_data(response: Any) -> list[Any] | None:
|
||||
def _get_container_list_data(response: object) -> Sequence[object] | None:
|
||||
if response is None:
|
||||
return None
|
||||
if isinstance(response, dict):
|
||||
|
|
@ -327,7 +354,7 @@ def _get_container_list_data(response: Any) -> list[Any] | None:
|
|||
return data if isinstance(data, list) else None
|
||||
|
||||
|
||||
def _set_container_list_data(response: Any, data: list[Any], removed_filtered_items: bool = False) -> Any:
|
||||
def _set_container_list_data(response: Any, data: list[object], removed_filtered_items: bool = False) -> object:
|
||||
if isinstance(response, dict):
|
||||
response["data"] = data
|
||||
if data:
|
||||
|
|
@ -353,7 +380,7 @@ def _set_container_list_data(response: Any, data: list[Any], removed_filtered_it
|
|||
|
||||
async def _get_allowed_container_ids(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> set[str]:
|
||||
) -> AbstractSet[str]:
|
||||
owner_scopes: Final = get_resource_owner_scopes(user_api_key_dict)
|
||||
if not owner_scopes:
|
||||
return set()
|
||||
|
|
@ -367,7 +394,8 @@ async def _get_allowed_container_ids(
|
|||
if prisma_client is None:
|
||||
return set()
|
||||
|
||||
rows: Final = await ManagedObjectRepository(prisma_client).table.find_many(
|
||||
table: Final[_ManagedObjectTable] = ManagedObjectRepository(prisma_client).table
|
||||
rows: Final[Sequence[_ManagedObjectRow]] = await table.find_many(
|
||||
where={
|
||||
"file_purpose": CONTAINER_OBJECT_PURPOSE,
|
||||
"created_by": {"in": owner_scopes},
|
||||
|
|
@ -382,10 +410,10 @@ async def _get_allowed_container_ids(
|
|||
|
||||
|
||||
async def filter_container_list_response(
|
||||
response: Any,
|
||||
response: object,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
custom_llm_provider: str,
|
||||
) -> Any:
|
||||
) -> object:
|
||||
if is_proxy_admin(user_api_key_dict):
|
||||
return response
|
||||
|
||||
|
|
@ -394,7 +422,7 @@ async def filter_container_list_response(
|
|||
return response
|
||||
|
||||
allowed_container_ids: Final = await _get_allowed_container_ids(user_api_key_dict)
|
||||
filtered: Final[list[Any]] = []
|
||||
filtered: Final[list[object]] = []
|
||||
for item in data:
|
||||
container_id = _get_response_id(item)
|
||||
if container_id is None:
|
||||
|
|
|
|||
|
|
@ -9,11 +9,13 @@ import contextlib
|
|||
import json
|
||||
import os
|
||||
import ssl
|
||||
from collections.abc import AsyncGenerator
|
||||
from collections.abc import AsyncGenerator, AsyncIterable, Mapping, Sequence
|
||||
from ssl import SSLContext
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from fastapi import HTTPException
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
from websockets.asyncio.client import ClientConnection, connect
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
|
||||
|
|
@ -35,8 +37,8 @@ from litellm.types.guardrails import GuardrailEventHooks
|
|||
from litellm.types.utils import (
|
||||
CallTypesLiteral,
|
||||
Choices,
|
||||
EmbeddingResponse,
|
||||
ImageResponse,
|
||||
LLMResponseTypes,
|
||||
Message,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
ResponsesAPIResponse,
|
||||
|
|
@ -50,6 +52,44 @@ class CatoNetworksGuardrailMissingSecrets(Exception):
|
|||
pass
|
||||
|
||||
|
||||
class _WsSslKwargs(TypedDict, total=False):
|
||||
ssl: bool | str | SSLContext
|
||||
|
||||
|
||||
class _CatoRequiredAction(TypedDict, total=False):
|
||||
action_type: str
|
||||
detection_message: str
|
||||
|
||||
|
||||
class _CatoRedactedMessage(TypedDict):
|
||||
role: NotRequired[str]
|
||||
content: str | None
|
||||
|
||||
|
||||
class _CatoRedactedChat(TypedDict, total=False):
|
||||
all_redacted_messages: Sequence[_CatoRedactedMessage]
|
||||
|
||||
|
||||
class _CatoAnalysisResult(TypedDict, total=False):
|
||||
policy_drill_down: Mapping[str, object]
|
||||
|
||||
|
||||
class _CatoAnalyzeResponse(TypedDict):
|
||||
required_action: NotRequired[_CatoRequiredAction | None]
|
||||
analysis_result: NotRequired[_CatoAnalysisResult]
|
||||
redacted_chat: NotRequired[_CatoRedactedChat]
|
||||
|
||||
|
||||
class _CatoOutputRedaction(TypedDict):
|
||||
redacted_output: str
|
||||
|
||||
|
||||
class _CatoStreamMessage(TypedDict, total=False):
|
||||
verified_chunk: Mapping[str, object]
|
||||
done: bool
|
||||
blocking_message: str
|
||||
|
||||
|
||||
class CatoNetworksGuardrail(CustomGuardrail):
|
||||
@classmethod
|
||||
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
|
||||
|
|
@ -80,7 +120,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
super().__init__(**kwargs)
|
||||
|
||||
@staticmethod
|
||||
def _build_ws_ssl_kwargs(ssl_verify: bool | str | None, ws_api_base: str) -> dict:
|
||||
def _build_ws_ssl_kwargs(ssl_verify: bool | str | None, ws_api_base: str) -> _WsSslKwargs:
|
||||
"""Resolve the ``ssl`` argument for ``websockets.connect``. Mirrors the
|
||||
``ssl_verify`` handling applied to the HTTP handler so a custom Cato instance
|
||||
behind TLS honours the same verification settings for streaming."""
|
||||
|
|
@ -156,7 +196,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
return flattened
|
||||
|
||||
@staticmethod
|
||||
def _prompt_inspection_messages(prompt: Any) -> list:
|
||||
def _prompt_inspection_messages(prompt: object) -> Sequence[Mapping[str, str]]:
|
||||
"""Synthetic user messages for a legacy completion ``prompt`` (a string
|
||||
or a list of string prompts)."""
|
||||
if isinstance(prompt, str):
|
||||
|
|
@ -166,7 +206,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
return []
|
||||
|
||||
@staticmethod
|
||||
def _iter_schema_string_refs(data: dict):
|
||||
def _iter_schema_string_refs(data: Mapping[str, Any]):
|
||||
"""Yield ``(container, key)`` for every non-empty schema string the proxy
|
||||
forwards to the model inside tool/function and structured-output schemas:
|
||||
each ``tools[].function`` and legacy ``functions[]`` entry plus the
|
||||
|
|
@ -208,7 +248,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
stack.extend(reversed(node))
|
||||
|
||||
@classmethod
|
||||
def _extra_inspection_sources(cls, data: dict) -> list:
|
||||
def _extra_inspection_sources(cls, data: Mapping[str, Any]) -> Sequence[tuple[str, Sequence[Mapping[str, str]]]]:
|
||||
"""Text the proxy forwards to the model outside chat ``messages``:
|
||||
Responses-API ``input`` and ``instructions``, legacy completion
|
||||
``prompt`` and tool/function/``response_format`` schema strings. Returned
|
||||
|
|
@ -251,7 +291,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
json={"messages": self._inspection_messages(data)},
|
||||
)
|
||||
response.raise_for_status()
|
||||
res: Final = response.json()
|
||||
res: Final[_CatoAnalyzeResponse] = response.json()
|
||||
required_action: Final = res.get("required_action")
|
||||
action_type: Final = required_action and required_action.get("action_type", None)
|
||||
if action_type is None:
|
||||
|
|
@ -267,7 +307,11 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
verbose_proxy_logger.error("Cato: %s action", action_type)
|
||||
return data
|
||||
|
||||
def _handle_block_action(self, analysis_result: Any, required_action: Any) -> None:
|
||||
def _handle_block_action(
|
||||
self,
|
||||
analysis_result: _CatoAnalysisResult,
|
||||
required_action: Any,
|
||||
) -> None:
|
||||
detection_message: Final = required_action.get("detection_message", None)
|
||||
verbose_proxy_logger.info(
|
||||
"Cato: Violation detected enabled policies: {policies}".format(
|
||||
|
|
@ -348,7 +392,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
hook: str,
|
||||
key_alias: str | None,
|
||||
user_email: str | None = None,
|
||||
) -> dict | None:
|
||||
) -> _CatoOutputRedaction | None:
|
||||
call_id: Final = request_data.get("litellm_call_id")
|
||||
inspection_messages: Final = self._inspection_messages(request_data)
|
||||
assistant_index: Final = len(inspection_messages)
|
||||
|
|
@ -363,7 +407,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
json={"messages": inspection_messages + [{"role": "assistant", "content": output}]},
|
||||
)
|
||||
response.raise_for_status()
|
||||
res: Final = response.json()
|
||||
res: Final[_CatoAnalyzeResponse] = response.json()
|
||||
required_action: Final = res.get("required_action")
|
||||
action_type: Final = required_action and required_action.get("action_type", None)
|
||||
if action_type and action_type == "block_action":
|
||||
|
|
@ -378,7 +422,11 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
return {"redacted_output": redacted_output}
|
||||
return None
|
||||
|
||||
def _handle_block_action_on_output(self, analysis_result: Any, required_action: Any) -> None:
|
||||
def _handle_block_action_on_output(
|
||||
self,
|
||||
analysis_result: _CatoAnalysisResult,
|
||||
required_action: Any,
|
||||
) -> None:
|
||||
detection_message: Final = required_action.get("detection_message", None)
|
||||
verbose_proxy_logger.info(
|
||||
"Cato: detected: {detected}, enabled policies: {policies}".format(
|
||||
|
|
@ -422,7 +470,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
@staticmethod
|
||||
def _output_fragments(message: Any) -> list:
|
||||
def _output_fragments(message: Message) -> Sequence[tuple[tuple[str, int | None], str]]:
|
||||
"""Assistant text the proxy returns to the caller: ``content`` plus every
|
||||
``tool_calls[].function.arguments`` string, each tagged with where a
|
||||
redaction must be written back. ``content`` is only included when present
|
||||
|
|
@ -439,7 +487,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
return fragments
|
||||
|
||||
@staticmethod
|
||||
def _apply_output_fragment(message: Any, target: tuple, redacted: str) -> None:
|
||||
def _apply_output_fragment(message: Any, target: tuple[str, int | None], redacted: str) -> None:
|
||||
kind, idx = target
|
||||
if kind == "content":
|
||||
message.content = redacted
|
||||
|
|
@ -447,11 +495,11 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
message.tool_calls[idx].function.arguments = redacted
|
||||
|
||||
@staticmethod
|
||||
def _responses_output_field(item: Any, key: str) -> Any:
|
||||
def _responses_output_field(item: object, key: str) -> str | Sequence[object] | None:
|
||||
return item.get(key) if isinstance(item, dict) else getattr(item, key, None)
|
||||
|
||||
@classmethod
|
||||
def _responses_output_fragments(cls, response: ResponsesAPIResponse) -> list:
|
||||
def _responses_output_fragments(cls, response: ResponsesAPIResponse) -> Sequence[tuple[object, str, str]]:
|
||||
"""Assistant text the Responses API returns to the caller: every
|
||||
``output_text`` content block plus every function-call ``arguments``
|
||||
string, each paired with the ``(container, key)`` a Cato redaction is
|
||||
|
|
@ -474,7 +522,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
return fragments
|
||||
|
||||
@staticmethod
|
||||
def _apply_responses_output_fragment(container: Any, key: str, redacted: str) -> None:
|
||||
def _apply_responses_output_fragment(container: object, key: str, redacted: str) -> None:
|
||||
if isinstance(container, dict):
|
||||
container[key] = redacted
|
||||
else:
|
||||
|
|
@ -505,8 +553,8 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
self,
|
||||
data: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any | ModelResponse | EmbeddingResponse | ImageResponse,
|
||||
) -> Any:
|
||||
response: LLMResponseTypes,
|
||||
) -> LLMResponseTypes:
|
||||
user_email: Final = self._resolve_cato_user_email(user_api_key_dict)
|
||||
if isinstance(response, ModelResponse) and response.choices:
|
||||
for choice in response.choices:
|
||||
|
|
@ -526,7 +574,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response,
|
||||
response: AsyncIterable[object],
|
||||
request_data: dict,
|
||||
) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
from litellm.proxy.proxy_server import StreamingCallbackError
|
||||
|
|
@ -547,7 +595,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
try:
|
||||
while True:
|
||||
raw_message = await self._await_cato_message(websocket, sender)
|
||||
result = json.loads(raw_message)
|
||||
result: _CatoStreamMessage = json.loads(raw_message)
|
||||
if verified_chunk := result.get("verified_chunk"):
|
||||
yield ModelResponseStream.model_validate(verified_chunk)
|
||||
continue
|
||||
|
|
@ -560,7 +608,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
finally:
|
||||
await self._cancel_background_task(sender)
|
||||
|
||||
async def _await_cato_message(self, websocket: ClientConnection, sender: asyncio.Task) -> Any:
|
||||
async def _await_cato_message(self, websocket: ClientConnection, sender: asyncio.Task[None]) -> str | bytes:
|
||||
"""Wait for the next Cato message, surfacing a dead forwarding task instead of blocking."""
|
||||
from litellm.proxy.proxy_server import StreamingCallbackError
|
||||
|
||||
|
|
@ -578,7 +626,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
|
|||
async def forward_the_stream_to_cato(
|
||||
self,
|
||||
websocket: ClientConnection,
|
||||
response_iter: AsyncGenerator[Any, None],
|
||||
response_iter: AsyncIterable[object],
|
||||
) -> None:
|
||||
async for chunk in response_iter:
|
||||
if isinstance(chunk, BaseModel):
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
|
||||
from urllib.parse import urlparse
|
||||
from uuid import uuid4
|
||||
|
||||
|
|
@ -29,9 +30,31 @@ from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import (
|
|||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
||||
|
||||
class _HiddenlayerEvaluation(TypedDict, total=False):
|
||||
action: str
|
||||
threat_level: str
|
||||
|
||||
|
||||
class _HiddenlayerAnalysisEntry(TypedDict, total=False):
|
||||
name: str
|
||||
detected: bool
|
||||
|
||||
|
||||
class _HiddenlayerModifiedSide(TypedDict):
|
||||
messages: Any
|
||||
|
||||
|
||||
class _HiddenlayerResponse(TypedDict, total=False):
|
||||
evaluation: _HiddenlayerEvaluation
|
||||
analysis: Sequence[_HiddenlayerAnalysisEntry]
|
||||
modified_data: Mapping[str, _HiddenlayerModifiedSide]
|
||||
|
||||
|
||||
def is_saas(host: str) -> bool:
|
||||
"""Checks whether the connection is to the SaaS platform"""
|
||||
|
||||
|
|
@ -43,7 +66,7 @@ def is_saas(host: str) -> bool:
|
|||
return False
|
||||
|
||||
|
||||
def _get_jwt(auth_url, api_id, api_key):
|
||||
def _get_jwt(auth_url, api_id, api_key) -> str:
|
||||
token_url: Final = f"{auth_url}/oauth2/token?grant_type=client_credentials"
|
||||
|
||||
resp: Final = requests.post(token_url, auth=HTTPBasicAuth(api_id, api_key))
|
||||
|
|
@ -139,7 +162,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
|
|||
|
||||
if scan_params := inputs.get("structured_messages"):
|
||||
last_msg: Final = scan_params[-1]
|
||||
result = await self._call_hiddenlayer(
|
||||
result: _HiddenlayerResponse = await self._call_hiddenlayer(
|
||||
project_id,
|
||||
hl_request_metadata,
|
||||
{
|
||||
|
|
@ -205,11 +228,11 @@ class HiddenlayerGuardrail(CustomGuardrail):
|
|||
async def _call_hiddenlayer(
|
||||
self,
|
||||
project_id: str | None,
|
||||
metadata: dict[str, str],
|
||||
payload: dict[str, Any],
|
||||
metadata: Mapping[str, str],
|
||||
payload: Mapping[str, Sequence[Mapping[str, str]]],
|
||||
input_type: Literal["request", "response"],
|
||||
) -> dict[str, Any]:
|
||||
data: Final[dict[str, Any]] = {"metadata": metadata}
|
||||
) -> _HiddenlayerResponse:
|
||||
data: Final[dict[str, object]] = {"metadata": metadata}
|
||||
|
||||
if input_type == "request":
|
||||
data["input"] = payload
|
||||
|
|
@ -235,7 +258,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
|
|||
headers=headers,
|
||||
)
|
||||
response.raise_for_status()
|
||||
result = response.json()
|
||||
result: _HiddenlayerResponse = response.json()
|
||||
|
||||
verbose_proxy_logger.debug("Hiddenlayer reponse: %s", result)
|
||||
|
||||
|
|
@ -265,7 +288,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
|
|||
return result
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> type[GuardrailConfigModel] | None:
|
||||
def get_config_model() -> type[GuardrailConfigModel[BaseModel]] | None:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import (
|
||||
HiddenlayerGuardrailConfigModel,
|
||||
)
|
||||
|
|
@ -343,7 +366,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
|
|||
if "hl-requester-id" not in hl_headers:
|
||||
hl_headers["hl-requester-id"] = "LiteLLM"
|
||||
|
||||
payload: Any
|
||||
payload: object
|
||||
if input_type == "request":
|
||||
payload = {
|
||||
"messages": inputs.get("structured_messages"),
|
||||
|
|
@ -461,7 +484,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
|
|||
return response
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> type[GuardrailConfigModel] | None:
|
||||
def get_config_model() -> type[GuardrailConfigModel[BaseModel]] | None:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import (
|
||||
HiddenlayerGuardrailConfigModel,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2,8 +2,11 @@ import threading
|
|||
import time
|
||||
import uuid
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from typing_extensions import NotRequired, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
|
|
@ -15,6 +18,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
GRAPH_API_BASE: Final = "https://graph.microsoft.com/v1.0"
|
||||
|
|
@ -25,6 +29,11 @@ GRAPH_SCOPE: Final = "https://graph.microsoft.com/.default"
|
|||
SCOPE_CACHE_TTL_SECONDS: Final = 3600.0
|
||||
|
||||
|
||||
class GraphTokenResponse(TypedDict):
|
||||
access_token: str
|
||||
expires_in: NotRequired[int]
|
||||
|
||||
|
||||
class PurviewGuardrailBase:
|
||||
"""
|
||||
Base class for Microsoft Purview guardrails.
|
||||
|
|
@ -41,8 +50,8 @@ class PurviewGuardrailBase:
|
|||
client_secret: str,
|
||||
purview_app_name: str = "LiteLLM",
|
||||
user_id_field: str = "user_id",
|
||||
**kwargs: Any,
|
||||
):
|
||||
**kwargs: object,
|
||||
) -> None:
|
||||
# Forward remaining kwargs to the next class in the MRO
|
||||
# (typically CustomGuardrail).
|
||||
super().__init__(**kwargs)
|
||||
|
|
@ -59,7 +68,7 @@ class PurviewGuardrailBase:
|
|||
|
||||
# Protection scope cache: user_id -> (etag, scope_response, fetched_at)
|
||||
# Capped at 1000 entries (LRU eviction) to avoid unbounded growth.
|
||||
self._scope_cache: OrderedDict[str, tuple[str, dict[str, Any], float]] = OrderedDict()
|
||||
self._scope_cache: OrderedDict[str, tuple[str, Mapping[str, object], float]] = OrderedDict()
|
||||
self._scope_cache_maxsize = 1000
|
||||
# Use a threading.Lock (not asyncio.Lock) because this lock is acquired
|
||||
# from both the proxy's main asyncio event loop and from short-lived
|
||||
|
|
@ -100,7 +109,7 @@ class PurviewGuardrailBase:
|
|||
headers={"Content-Type": "application/x-www-form-urlencoded"},
|
||||
)
|
||||
response.raise_for_status()
|
||||
token_data: Final = response.json()
|
||||
token_data: Final[GraphTokenResponse] = response.json()
|
||||
access_token: Final = token_data["access_token"]
|
||||
expires_in: Final = int(token_data.get("expires_in", 3599))
|
||||
# Recompute ``now`` after the await so the expiry reflects when the
|
||||
|
|
@ -117,9 +126,9 @@ class PurviewGuardrailBase:
|
|||
async def _graph_post(
|
||||
self,
|
||||
url: str,
|
||||
json_body: dict[str, Any],
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> tuple[dict[str, Any], dict[str, str]]:
|
||||
json_body: dict[str, object],
|
||||
extra_headers: Mapping[str, str] | None = None,
|
||||
) -> tuple[dict[str, object], dict[str, str]]:
|
||||
"""POST to Graph API with bearer auth.
|
||||
|
||||
Returns:
|
||||
|
|
@ -136,7 +145,7 @@ class PurviewGuardrailBase:
|
|||
verbose_proxy_logger.debug("Purview Graph POST %s", url)
|
||||
response: Final = await self.async_handler.post(url=url, headers=headers, json=json_body)
|
||||
response.raise_for_status()
|
||||
response_json: Final[dict[str, Any]] = response.json()
|
||||
response_json: Final[dict[str, object]] = response.json()
|
||||
response_headers: Final = dict(response.headers)
|
||||
verbose_proxy_logger.debug("Purview Graph response: %s", response_json)
|
||||
return response_json, response_headers
|
||||
|
|
@ -145,7 +154,7 @@ class PurviewGuardrailBase:
|
|||
# Protection scopes
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _compute_protection_scopes(self, user_id: str) -> tuple[str, dict[str, Any]]:
|
||||
async def _compute_protection_scopes(self, user_id: str) -> tuple[str, Mapping[str, object]]:
|
||||
"""Call protectionScopes/compute and cache with ETag.
|
||||
|
||||
Returns:
|
||||
|
|
@ -161,7 +170,7 @@ class PurviewGuardrailBase:
|
|||
return cached[0], cached[1]
|
||||
|
||||
url: Final = f"{GRAPH_API_BASE}/users/{encoded_user_id}/dataSecurityAndGovernance/protectionScopes/compute"
|
||||
body: Final[dict[str, Any]] = {
|
||||
body: Final[dict[str, object]] = {
|
||||
"activities": "uploadText,downloadText",
|
||||
"locations": [
|
||||
{
|
||||
|
|
@ -199,7 +208,7 @@ class PurviewGuardrailBase:
|
|||
activity: str,
|
||||
etag: str,
|
||||
correlation_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
) -> dict[str, object]:
|
||||
"""Call processContent for DLP policy evaluation.
|
||||
|
||||
Args:
|
||||
|
|
@ -211,7 +220,7 @@ class PurviewGuardrailBase:
|
|||
"""
|
||||
encoded_user_id: Final = self._encode_graph_user_id(user_id)
|
||||
url: Final = f"{GRAPH_API_BASE}/users/{encoded_user_id}/dataSecurityAndGovernance/processContent"
|
||||
body: Final[dict[str, Any]] = {
|
||||
body: Final[dict[str, object]] = {
|
||||
"contentToProcess": {
|
||||
"contentEntries": [
|
||||
{
|
||||
|
|
@ -261,7 +270,7 @@ class PurviewGuardrailBase:
|
|||
# User ID resolution
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _resolve_user_id(self, data: dict[str, Any], user_api_key_dict: Any) -> str | None:
|
||||
def _resolve_user_id(self, data: Mapping[str, object], user_api_key_dict: "UserAPIKeyAuth") -> str | None:
|
||||
"""Resolve the Entra user object ID from request data or auth context.
|
||||
|
||||
Returns the strongest available identity walking down four sources, in
|
||||
|
|
@ -284,7 +293,10 @@ class PurviewGuardrailBase:
|
|||
if hasattr(user_api_key_dict, "end_user_id") and user_api_key_dict.end_user_id:
|
||||
return str(user_api_key_dict.end_user_id)
|
||||
|
||||
metadata: Final = data.get("metadata") or data.get("litellm_metadata") or {}
|
||||
metadata_value: Final[object] = data.get("metadata") or data.get("litellm_metadata") or {}
|
||||
if not isinstance(metadata_value, Mapping):
|
||||
return None
|
||||
metadata: Final[Mapping[str, object]] = metadata_value
|
||||
uid = metadata.get("user_api_key_user_id")
|
||||
if uid:
|
||||
return str(uid)
|
||||
|
|
@ -296,15 +308,15 @@ class PurviewGuardrailBase:
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
def _logging_kwargs_metadata(kwargs: dict[str, Any]) -> dict[str, Any]:
|
||||
def _logging_kwargs_metadata(kwargs: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""Metadata dict from ``model_call_details`` / logging kwargs."""
|
||||
litellm_params: Final = kwargs.get("litellm_params") or {}
|
||||
litellm_params: Final[object] = kwargs.get("litellm_params") or {}
|
||||
if not isinstance(litellm_params, dict):
|
||||
return {}
|
||||
md: Final = litellm_params.get("metadata")
|
||||
return md if isinstance(md, dict) else {}
|
||||
|
||||
def _resolve_trusted_user_id(self, data: dict[str, Any], user_api_key_dict: Any) -> str | None:
|
||||
def _resolve_trusted_user_id(self, data: Mapping[str, object], user_api_key_dict: "UserAPIKeyAuth") -> str | None:
|
||||
"""Resolve user ID from API-key/JWT-bound identity for blocking DLP.
|
||||
|
||||
Uses only ``UserAPIKeyAuth.user_id`` (bound on the LiteLLM key or JWT).
|
||||
|
|
@ -325,7 +337,7 @@ class PurviewGuardrailBase:
|
|||
|
||||
return None
|
||||
|
||||
def _resolve_user_id_from_logging_kwargs(self, kwargs: dict[str, Any]) -> str | None:
|
||||
def _resolve_user_id_from_logging_kwargs(self, kwargs: Mapping[str, object]) -> str | None:
|
||||
"""Trusted-identity-only resolver for logging-only hooks.
|
||||
|
||||
Uses only the proxy-injected ``user_api_key_user_id`` (populated from
|
||||
|
|
@ -365,7 +377,7 @@ class PurviewGuardrailBase:
|
|||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def is_token_id_prompt(prompt: Any) -> bool:
|
||||
def is_token_id_prompt(prompt: str | Sequence[object] | None) -> bool:
|
||||
"""Return True if ``prompt`` carries OpenAI completions token ids.
|
||||
|
||||
Covers every list shape that ``completion_prompt_to_str`` cannot decode
|
||||
|
|
@ -383,7 +395,7 @@ class PurviewGuardrailBase:
|
|||
return False
|
||||
|
||||
@staticmethod
|
||||
def completion_prompt_to_str(prompt: Any) -> str | None:
|
||||
def completion_prompt_to_str(prompt: str | Sequence[object] | None) -> str | None:
|
||||
"""Normalize OpenAI ``/v1/completions`` ``prompt`` for text DLP.
|
||||
|
||||
Supports string prompts and list-of-string prompts. List-of-token-id prompts
|
||||
|
|
@ -408,7 +420,7 @@ class PurviewGuardrailBase:
|
|||
return None
|
||||
|
||||
@staticmethod
|
||||
def _extract_tool_call_args_from_message(message: Any) -> list[str]:
|
||||
def _extract_tool_call_args_from_message(message: object) -> list[str]:
|
||||
"""Return plaintext arguments strings from tool_calls and function_call fields.
|
||||
|
||||
Covers both the request path (assistant messages in chat histories that
|
||||
|
|
@ -419,7 +431,9 @@ class PurviewGuardrailBase:
|
|||
args: Final[list[str]] = []
|
||||
|
||||
# tool_calls: [{"function": {"arguments": "..."}}]
|
||||
tool_calls = message.get("tool_calls") if isinstance(message, dict) else getattr(message, "tool_calls", None)
|
||||
tool_calls: Final[Sequence[object] | None] = (
|
||||
message.get("tool_calls") if isinstance(message, dict) else getattr(message, "tool_calls", None)
|
||||
)
|
||||
if tool_calls:
|
||||
for tc in tool_calls:
|
||||
fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None)
|
||||
|
|
|
|||
|
|
@ -14,7 +14,8 @@ GET /v1/workflows/runs/{run_id}/messages - Fetch conversation history
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Final, Literal
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final, Literal, Protocol, TypedDict
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
|
|
@ -43,7 +44,7 @@ router: Final = APIRouter()
|
|||
_MAX_SEQUENCE_RETRIES: Final = 5
|
||||
|
||||
|
||||
def _json(value: Any) -> str:
|
||||
def _json(value: object) -> str:
|
||||
"""Serialize a Python value for prisma-client-py Json fields (must be a string)."""
|
||||
return json.dumps(value)
|
||||
|
||||
|
|
@ -62,7 +63,7 @@ def _caller_key(user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
|||
|
||||
|
||||
# Status transitions driven by event_type
|
||||
_EVENT_STATUS_MAP: Final[dict[str, str]] = {
|
||||
_EVENT_STATUS_MAP: Final[Mapping[str, str]] = {
|
||||
"step.started": "running",
|
||||
"step.failed": "failed",
|
||||
"hook.waiting": "paused",
|
||||
|
|
@ -77,8 +78,8 @@ _EVENT_STATUS_MAP: Final[dict[str, str]] = {
|
|||
|
||||
class WorkflowRunCreateRequest(BaseModel):
|
||||
workflow_type: str
|
||||
input: dict[str, Any] | None = None
|
||||
metadata: dict[str, Any] | None = None
|
||||
input: Mapping[str, object] | None = None
|
||||
metadata: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
WorkflowRunStatus = Literal["pending", "running", "paused", "completed", "failed"]
|
||||
|
|
@ -86,14 +87,14 @@ WorkflowRunStatus = Literal["pending", "running", "paused", "completed", "failed
|
|||
|
||||
class WorkflowRunUpdateRequest(BaseModel):
|
||||
status: WorkflowRunStatus | None = None
|
||||
output: dict[str, Any] | None = None
|
||||
metadata: dict[str, Any] | None = None
|
||||
output: Mapping[str, object] | None = None
|
||||
metadata: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
class WorkflowEventCreateRequest(BaseModel):
|
||||
event_type: str
|
||||
step_name: str
|
||||
data: dict[str, Any] | None = None
|
||||
data: Mapping[str, object] | None = None
|
||||
|
||||
|
||||
class WorkflowMessageCreateRequest(BaseModel):
|
||||
|
|
@ -102,15 +103,60 @@ class WorkflowMessageCreateRequest(BaseModel):
|
|||
session_id: str | None = None
|
||||
|
||||
|
||||
class _RunRow(Protocol):
|
||||
@property
|
||||
def created_by(self) -> str | None: ...
|
||||
|
||||
|
||||
class _SeqRow(Protocol):
|
||||
@property
|
||||
def sequence_number(self) -> int: ...
|
||||
|
||||
|
||||
class _RunCreateData(TypedDict, total=False):
|
||||
workflow_type: str
|
||||
created_by: str | None
|
||||
input: str
|
||||
metadata: str
|
||||
|
||||
|
||||
class _RunWhere(TypedDict, total=False):
|
||||
workflow_type: str
|
||||
status: str | Mapping[str, Sequence[str]]
|
||||
created_by: str
|
||||
|
||||
|
||||
class _RunUpdateData(TypedDict, total=False):
|
||||
status: WorkflowRunStatus
|
||||
output: str
|
||||
metadata: str
|
||||
|
||||
|
||||
class _EventCreateData(TypedDict, total=False):
|
||||
run_id: str
|
||||
event_type: str
|
||||
step_name: str
|
||||
sequence_number: int
|
||||
data: str
|
||||
|
||||
|
||||
class _MessageCreateData(TypedDict, total=False):
|
||||
run_id: str
|
||||
role: str
|
||||
content: str
|
||||
sequence_number: int
|
||||
session_id: str
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _get_next_sequence_number(prisma_client: Any, run_id: str, table: str) -> int:
|
||||
async def _get_next_sequence_number(prisma_client: object, run_id: str, table: str) -> int:
|
||||
"""Return MAX(sequence_number) + 1 for the given run, for either events or messages."""
|
||||
if table == "events":
|
||||
rows = await WorkflowEventRepository(prisma_client).table.find_many(
|
||||
rows: Sequence[_SeqRow] = await WorkflowEventRepository(prisma_client).table.find_many(
|
||||
where={"run_id": run_id},
|
||||
order={"sequence_number": "desc"},
|
||||
take=1,
|
||||
|
|
@ -125,12 +171,12 @@ async def _get_next_sequence_number(prisma_client: Any, run_id: str, table: str)
|
|||
|
||||
|
||||
async def _require_run(
|
||||
prisma_client: Any,
|
||||
prisma_client: object,
|
||||
run_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth | None = None,
|
||||
) -> Any:
|
||||
) -> _RunRow:
|
||||
"""Return the run or raise 404. For non-admin callers, also enforce key ownership."""
|
||||
run: Final = await WorkflowRunRepository(prisma_client).table.find_unique(where={"run_id": run_id})
|
||||
run: Final[_RunRow | None] = await WorkflowRunRepository(prisma_client).table.find_unique(where={"run_id": run_id})
|
||||
if run is None:
|
||||
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
|
||||
if user_api_key_dict is not None and not _is_admin(user_api_key_dict):
|
||||
|
|
@ -165,7 +211,7 @@ async def create_workflow_run(
|
|||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
try:
|
||||
create_data: Final[dict[str, Any]] = {
|
||||
create_data: Final[_RunCreateData] = {
|
||||
"workflow_type": data.workflow_type,
|
||||
"created_by": _caller_key(user_api_key_dict),
|
||||
}
|
||||
|
|
@ -173,7 +219,7 @@ async def create_workflow_run(
|
|||
create_data["input"] = _json(data.input)
|
||||
if data.metadata is not None:
|
||||
create_data["metadata"] = _json(data.metadata)
|
||||
run: Final = await WorkflowRunRepository(prisma_client).table.create(data=create_data)
|
||||
run: Final[_RunRow] = await WorkflowRunRepository(prisma_client).table.create(data=create_data)
|
||||
return run
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error creating workflow run: %s", e)
|
||||
|
|
@ -200,7 +246,7 @@ async def list_workflow_runs(
|
|||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
where: Final[dict[str, Any]] = {}
|
||||
where: Final[_RunWhere] = {}
|
||||
if workflow_type:
|
||||
where["workflow_type"] = workflow_type
|
||||
if status:
|
||||
|
|
@ -214,7 +260,7 @@ async def list_workflow_runs(
|
|||
where["created_by"] = caller
|
||||
|
||||
try:
|
||||
runs: Final = await WorkflowRunRepository(prisma_client).table.find_many(
|
||||
runs: Final[Sequence[object]] = await WorkflowRunRepository(prisma_client).table.find_many(
|
||||
where=where,
|
||||
order={"created_at": "desc"},
|
||||
take=limit,
|
||||
|
|
@ -241,7 +287,7 @@ async def get_workflow_run(
|
|||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
try:
|
||||
run: Final = await WorkflowRunRepository(prisma_client).table.find_unique(
|
||||
run: Final[_RunRow | None] = await WorkflowRunRepository(prisma_client).table.find_unique(
|
||||
where={"run_id": run_id},
|
||||
include={"events": {"order_by": {"sequence_number": "desc"}, "take": 1}},
|
||||
)
|
||||
|
|
@ -275,7 +321,7 @@ async def update_workflow_run(
|
|||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value)
|
||||
|
||||
update: Final[dict[str, Any]] = {}
|
||||
update: Final[_RunUpdateData] = {}
|
||||
if data.status is not None:
|
||||
update["status"] = data.status
|
||||
if data.output is not None:
|
||||
|
|
@ -290,7 +336,7 @@ async def update_workflow_run(
|
|||
await _require_run(prisma_client, run_id, user_api_key_dict)
|
||||
|
||||
try:
|
||||
run: Final = await WorkflowRunRepository(prisma_client).table.update(
|
||||
run: Final[_RunRow | None] = await WorkflowRunRepository(prisma_client).table.update(
|
||||
where={"run_id": run_id},
|
||||
data=update,
|
||||
)
|
||||
|
|
@ -332,7 +378,7 @@ async def append_workflow_event(
|
|||
for attempt in range(_MAX_SEQUENCE_RETRIES):
|
||||
try:
|
||||
seq = await _get_next_sequence_number(prisma_client, run_id, "events")
|
||||
event_data: dict[str, Any] = {
|
||||
event_data: _EventCreateData = {
|
||||
"run_id": run_id,
|
||||
"event_type": data.event_type,
|
||||
"step_name": data.step_name,
|
||||
|
|
@ -342,7 +388,7 @@ async def append_workflow_event(
|
|||
event_data["data"] = _json(data.data)
|
||||
|
||||
async with prisma_client.db.tx() as tx:
|
||||
event = await tx.litellm_workflowevent.create(data=event_data)
|
||||
event: object = await tx.litellm_workflowevent.create(data=event_data)
|
||||
if new_status:
|
||||
await tx.litellm_workflowrun.update(
|
||||
where={"run_id": run_id},
|
||||
|
|
@ -389,7 +435,7 @@ async def list_workflow_events(
|
|||
await _require_run(prisma_client, run_id, _read_scope_caller(user_api_key_dict))
|
||||
|
||||
try:
|
||||
events: Final = await WorkflowEventRepository(prisma_client).table.find_many(
|
||||
events: Final[Sequence[object]] = await WorkflowEventRepository(prisma_client).table.find_many(
|
||||
where={"run_id": run_id},
|
||||
order={"sequence_number": "asc"},
|
||||
take=limit,
|
||||
|
|
@ -424,7 +470,7 @@ async def append_workflow_message(
|
|||
for attempt in range(_MAX_SEQUENCE_RETRIES):
|
||||
try:
|
||||
seq = await _get_next_sequence_number(prisma_client, run_id, "messages")
|
||||
msg_data: dict[str, Any] = {
|
||||
msg_data: _MessageCreateData = {
|
||||
"run_id": run_id,
|
||||
"role": data.role,
|
||||
"content": data.content,
|
||||
|
|
@ -432,7 +478,7 @@ async def append_workflow_message(
|
|||
}
|
||||
if data.session_id is not None:
|
||||
msg_data["session_id"] = data.session_id
|
||||
msg = await WorkflowMessageRepository(prisma_client).table.create(data=msg_data)
|
||||
msg: object = await WorkflowMessageRepository(prisma_client).table.create(data=msg_data)
|
||||
return msg
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -473,7 +519,7 @@ async def list_workflow_messages(
|
|||
await _require_run(prisma_client, run_id, _read_scope_caller(user_api_key_dict))
|
||||
|
||||
try:
|
||||
messages: Final = await WorkflowMessageRepository(prisma_client).table.find_many(
|
||||
messages: Final[Sequence[object]] = await WorkflowMessageRepository(prisma_client).table.find_many(
|
||||
where={"run_id": run_id},
|
||||
order={"sequence_number": "asc"},
|
||||
take=limit,
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ This allows the same policy to be attached to multiple scopes.
|
|||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, TypedDict
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.repositories.table_repositories import PolicyAttachmentRepository
|
||||
|
|
@ -18,9 +18,18 @@ from litellm.types.proxy.policy_engine import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Sequence
|
||||
|
||||
from prisma.models import LiteLLM_PolicyAttachmentTable
|
||||
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
|
||||
class PolicyAttachmentMatch(TypedDict):
|
||||
policy_name: str
|
||||
matched_via: str
|
||||
|
||||
|
||||
class AttachmentRegistry:
|
||||
"""
|
||||
In-memory registry for storing and managing policy attachments.
|
||||
|
|
@ -40,7 +49,7 @@ class AttachmentRegistry:
|
|||
```
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self) -> None:
|
||||
self._attachments: list[PolicyAttachment] = []
|
||||
self._config_attachments: tuple[PolicyAttachment, ...] = ()
|
||||
self._initialized: bool = False
|
||||
|
|
@ -98,7 +107,7 @@ class AttachmentRegistry:
|
|||
"""
|
||||
return [r["policy_name"] for r in self.get_attached_policies_with_reasons(context)]
|
||||
|
||||
def get_attached_policies_with_reasons(self, context: PolicyMatchContext) -> list[dict[str, Any]]:
|
||||
def get_attached_policies_with_reasons(self, context: PolicyMatchContext) -> list[PolicyAttachmentMatch]:
|
||||
"""
|
||||
Get list of policy names and match reasons for the given context.
|
||||
|
||||
|
|
@ -107,8 +116,8 @@ class AttachmentRegistry:
|
|||
"""
|
||||
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
|
||||
|
||||
results: Final[list[dict[str, Any]]] = []
|
||||
seen_policies: Final[set] = set()
|
||||
results: Final[list[PolicyAttachmentMatch]] = []
|
||||
seen_policies: Final[set[str]] = set()
|
||||
|
||||
for attachment in self._attachments:
|
||||
scope = attachment.to_policy_scope()
|
||||
|
|
@ -280,7 +289,9 @@ class AttachmentRegistry:
|
|||
PolicyAttachmentDBResponse with the created attachment
|
||||
"""
|
||||
try:
|
||||
created_attachment: Final = await PolicyAttachmentRepository(prisma_client).table.create(
|
||||
created_attachment: Final[LiteLLM_PolicyAttachmentTable] = await PolicyAttachmentRepository(
|
||||
prisma_client
|
||||
).table.create(
|
||||
data={
|
||||
"policy_name": attachment_request.policy_name,
|
||||
"scope": attachment_request.scope,
|
||||
|
|
@ -340,9 +351,9 @@ class AttachmentRegistry:
|
|||
"""
|
||||
try:
|
||||
# Get attachment before deleting
|
||||
attachment: Final = await PolicyAttachmentRepository(prisma_client).table.find_unique(
|
||||
where={"attachment_id": attachment_id}
|
||||
)
|
||||
attachment: Final[LiteLLM_PolicyAttachmentTable | None] = await PolicyAttachmentRepository(
|
||||
prisma_client
|
||||
).table.find_unique(where={"attachment_id": attachment_id})
|
||||
|
||||
if attachment is None:
|
||||
raise Exception(f"Attachment with ID {attachment_id} not found")
|
||||
|
|
@ -375,9 +386,9 @@ class AttachmentRegistry:
|
|||
PolicyAttachmentDBResponse if found, None otherwise
|
||||
"""
|
||||
try:
|
||||
attachment: Final = await PolicyAttachmentRepository(prisma_client).table.find_unique(
|
||||
where={"attachment_id": attachment_id}
|
||||
)
|
||||
attachment: Final[LiteLLM_PolicyAttachmentTable | None] = await PolicyAttachmentRepository(
|
||||
prisma_client
|
||||
).table.find_unique(where={"attachment_id": attachment_id})
|
||||
|
||||
if attachment is None:
|
||||
return None
|
||||
|
|
@ -413,7 +424,9 @@ class AttachmentRegistry:
|
|||
List of PolicyAttachmentDBResponse objects
|
||||
"""
|
||||
try:
|
||||
attachments: Final = await PolicyAttachmentRepository(prisma_client).table.find_many(
|
||||
attachments: Final[Sequence[LiteLLM_PolicyAttachmentTable]] = await PolicyAttachmentRepository(
|
||||
prisma_client
|
||||
).table.find_many(
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,30 +1,30 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 3126
|
||||
"limit": 3121
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 71
|
||||
},
|
||||
"ANN003": {
|
||||
"limit": 836
|
||||
"limit": 834
|
||||
},
|
||||
"ANN201": {
|
||||
"limit": 2037
|
||||
"limit": 2033
|
||||
},
|
||||
"ANN202": {
|
||||
"limit": 869
|
||||
"limit": 865
|
||||
},
|
||||
"ANN204": {
|
||||
"limit": 715
|
||||
"limit": 713
|
||||
},
|
||||
"ANN205": {
|
||||
"limit": 115
|
||||
"limit": 114
|
||||
},
|
||||
"ANN206": {
|
||||
"limit": 133
|
||||
},
|
||||
"ANN401": {
|
||||
"limit": 1689
|
||||
"limit": 1630
|
||||
},
|
||||
"ASYNC230": {
|
||||
"limit": 11
|
||||
|
|
@ -222,7 +222,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"RET504": {
|
||||
"limit": 178
|
||||
"limit": 177
|
||||
},
|
||||
"RUF010": {
|
||||
"limit": 0
|
||||
|
|
@ -306,7 +306,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"TID251": {
|
||||
"limit": 1242
|
||||
"limit": 1240
|
||||
},
|
||||
"TRY002": {
|
||||
"limit": 528
|
||||
|
|
|
|||
|
|
@ -18,13 +18,14 @@ from litellm.types.proxy.claude_code_endpoints import (
|
|||
UpdatePluginRequest,
|
||||
)
|
||||
from litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_marketplace import (
|
||||
get_marketplace,
|
||||
register_plugin,
|
||||
update_plugin,
|
||||
)
|
||||
|
||||
|
||||
def _make_mock_prisma():
|
||||
"""Stateful prisma mock that supports find_unique, create, and update."""
|
||||
"""Stateful prisma mock that supports find_unique, find_many, create, and update."""
|
||||
store: dict = {}
|
||||
|
||||
mock_client = MagicMock()
|
||||
|
|
@ -34,6 +35,12 @@ def _make_mock_prisma():
|
|||
async def _find_unique(where):
|
||||
return store.get(where.get("name"))
|
||||
|
||||
async def _find_many(where=None):
|
||||
records = list(store.values())
|
||||
if where and "enabled" in where:
|
||||
return [r for r in records if r.enabled == where["enabled"]]
|
||||
return records
|
||||
|
||||
async def _create(data):
|
||||
record = MagicMock()
|
||||
record.id = "test-id"
|
||||
|
|
@ -52,6 +59,7 @@ def _make_mock_prisma():
|
|||
return record
|
||||
|
||||
mock_table.find_unique = AsyncMock(side_effect=_find_unique)
|
||||
mock_table.find_many = AsyncMock(side_effect=_find_many)
|
||||
mock_table.create = AsyncMock(side_effect=_create)
|
||||
mock_table.update = AsyncMock(side_effect=_update)
|
||||
mock_client.db.litellm_claudecodeplugintable = mock_table
|
||||
|
|
@ -211,6 +219,23 @@ async def test_update_plugin_db_error_maps_to_structured_500():
|
|||
assert "connection lost" in exc_info.value.detail["error"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_marketplace_skips_plugin_with_null_manifest():
|
||||
await register_plugin(
|
||||
request=RegisterPluginRequest(name="good-plugin", source=_GIT_SUBDIR_SOURCE, version="1.0.0"),
|
||||
user_api_key_dict=_USER,
|
||||
)
|
||||
|
||||
table = litellm.proxy.proxy_server.prisma_client.db.litellm_claudecodeplugintable
|
||||
await table.create(data={"name": "null-manifest-plugin", "manifest_json": None, "enabled": True})
|
||||
|
||||
response = await get_marketplace()
|
||||
|
||||
assert response.status_code == 200
|
||||
body = json.loads(response.body)
|
||||
assert [plugin["name"] for plugin in body["plugins"]] == ["good-plugin"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_plugin_git_subdir_missing_url():
|
||||
"""git-subdir without url field raises HTTP 400."""
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 23332
|
||||
"limit": 23256
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 27213
|
||||
|
|
@ -27,7 +27,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16792
|
||||
"limit": 16783
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5602
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue