mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
refactor(types): replace Any with proven types in 16 files (#45029)
* refactor(types): replace Any with proven types in 16 files * fix(types): import TypedDict from typing_extensions for pydantic on 3.10 --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
498a3e1a67
commit
cb138ba92f
16 changed files with 226 additions and 65 deletions
|
|
@ -315,7 +315,7 @@ class LLMCachingHandler:
|
|||
args = args or ()
|
||||
final_embedding_cached_response: EmbeddingResponse | None = None
|
||||
embedding_all_elements_cache_hit: bool = False
|
||||
cached_result: Any | None = None
|
||||
cached_result: object | None = None
|
||||
kwargs = kwargs.copy()
|
||||
#########################################################
|
||||
# Init cache timing metrics
|
||||
|
|
@ -849,7 +849,7 @@ class LLMCachingHandler:
|
|||
if new_kwargs.get("stream") is True and "cache_key" not in new_kwargs:
|
||||
new_kwargs["cache_key"] = litellm.cache.get_cache_key(**new_kwargs)
|
||||
self.request_kwargs = _drop_logging_obj_from_kwargs(new_kwargs)
|
||||
cached_result: Any | None = None
|
||||
cached_result: object | None = None
|
||||
if call_type == CallTypes.aembedding.value:
|
||||
if isinstance(new_kwargs["input"], str):
|
||||
new_kwargs["input"] = [new_kwargs["input"]]
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ from __future__ import annotations
|
|||
|
||||
import asyncio
|
||||
import importlib
|
||||
from collections.abc import AsyncIterator, Mapping
|
||||
from collections.abc import AsyncIterator, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType, ModuleType
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
|
@ -59,10 +59,10 @@ _MODEL_NODE = "model"
|
|||
class DeepAgentsDeps:
|
||||
"""The optional-dependency entrypoints this handler uses."""
|
||||
|
||||
create_deep_agent: Any
|
||||
chat_litellm: Any
|
||||
checkpointer_cls: Any
|
||||
command_cls: Any
|
||||
create_deep_agent: Callable[..., CompiledStateGraph]
|
||||
chat_litellm: Callable[..., BaseChatModel]
|
||||
checkpointer_cls: Callable[[], BaseCheckpointSaver]
|
||||
command_cls: Callable[..., Command]
|
||||
subagent_defaults: Mapping[str, object]
|
||||
convert_to_openai_messages: Any
|
||||
backend: ModuleType
|
||||
|
|
@ -112,7 +112,7 @@ class DeepAgentsHandler(BaseHarnessHandler):
|
|||
def __init__(self, config: BaseHarnessConfig) -> None:
|
||||
super().__init__(config)
|
||||
self._deps: DeepAgentsDeps | None = None
|
||||
self._agent: Any = None
|
||||
self._agent: CompiledStateGraph | None = None
|
||||
self._thread_id: str | None = None
|
||||
self._skip_tools: frozenset[str] = frozenset()
|
||||
|
||||
|
|
@ -221,7 +221,7 @@ class DeepAgentsHandler(BaseHarnessHandler):
|
|||
if ctx.output is not None:
|
||||
ctx.output_json = structured_json(values.get("structured_response"))
|
||||
|
||||
def _require_agent(self) -> tuple[Any, DeepAgentsDeps]:
|
||||
def _require_agent(self) -> tuple[CompiledStateGraph, DeepAgentsDeps]:
|
||||
if self._agent is None or self._deps is None:
|
||||
raise HarnessError("Deep Agents session is not started")
|
||||
return self._agent, self._deps
|
||||
|
|
|
|||
|
|
@ -60,7 +60,7 @@ _GCHUNK_FIELDS: Final[frozenset] = frozenset(GChunk.__annotations__)
|
|||
_USAGE_COST_HEADER_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.OPENROUTER.value})
|
||||
|
||||
|
||||
def _next_sync_or_exhausted(it: Any) -> object:
|
||||
def _next_sync_or_exhausted(it: Iterator[object]) -> object:
|
||||
"""
|
||||
Call next(it) from a thread and return _SYNC_ITER_EXHAUSTED on StopIteration.
|
||||
|
||||
|
|
|
|||
|
|
@ -1,9 +1,11 @@
|
|||
import json
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from collections.abc import AsyncIterator, Iterator, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import ConfigDict, TypeAdapter, with_config
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import cohere_messages_pt_v2
|
||||
|
|
@ -23,6 +25,35 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _CohereToolCall(TypedDict):
|
||||
name: NotRequired[ReadOnly[object]]
|
||||
generation_id: NotRequired[ReadOnly[object]]
|
||||
parameters: NotRequired[ReadOnly[object]]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _CohereBilledUnits(TypedDict):
|
||||
input_tokens: NotRequired[ReadOnly[int | float]]
|
||||
output_tokens: NotRequired[ReadOnly[int | float]]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _CohereMeta(TypedDict):
|
||||
billed_units: NotRequired[ReadOnly[_CohereBilledUnits]]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _CohereChatResponse(TypedDict):
|
||||
text: ReadOnly[str | None]
|
||||
citations: NotRequired[ReadOnly[object]]
|
||||
tool_calls: NotRequired[ReadOnly[Sequence[_CohereToolCall] | None]]
|
||||
meta: NotRequired[ReadOnly[_CohereMeta]]
|
||||
|
||||
|
||||
_COHERE_CHAT_RESPONSE: Final = TypeAdapter(_CohereChatResponse)
|
||||
|
||||
|
||||
class CohereError(BaseLLMException):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -231,7 +262,7 @@ class CohereChatConfig(BaseConfig):
|
|||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
try:
|
||||
raw_response_json: Final = raw_response.json()
|
||||
raw_response_json: Final = _COHERE_CHAT_RESPONSE.validate_python(raw_response.json())
|
||||
model_response.choices[0].message.content = raw_response_json["text"]
|
||||
except Exception:
|
||||
raise CohereError(message=raw_response.text, status_code=raw_response.status_code)
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ from datetime import datetime
|
|||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import ConfigDict, TypeAdapter, with_config
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.asyncify import can_block_current_thread
|
||||
|
|
@ -25,6 +27,24 @@ DEFAULT_GITHUB_ACCESS_TOKEN_URL: Final = "https://github.com/login/oauth/access_
|
|||
DEFAULT_GITHUB_API_KEY_URL: Final = "https://api.github.com/copilot_internal/v2/token"
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _DeviceCode(TypedDict):
|
||||
device_code: ReadOnly[str]
|
||||
user_code: ReadOnly[object]
|
||||
verification_uri: ReadOnly[object]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _AccessTokenPoll(TypedDict):
|
||||
access_token: ReadOnly[NotRequired[str]]
|
||||
error: ReadOnly[NotRequired[object]]
|
||||
|
||||
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(strict=True))
|
||||
_DEVICE_CODE: Final = TypeAdapter(_DeviceCode)
|
||||
_ACCESS_TOKEN_POLL: Final = TypeAdapter(_AccessTokenPoll)
|
||||
|
||||
|
||||
class Authenticator:
|
||||
def __init__(self) -> None:
|
||||
"""Initialize the GitHub Copilot authenticator with configurable token paths."""
|
||||
|
|
@ -226,7 +246,7 @@ class Authenticator:
|
|||
|
||||
return headers
|
||||
|
||||
def _get_device_code(self) -> dict[str, str]:
|
||||
def _get_device_code(self) -> _DeviceCode:
|
||||
"""
|
||||
Get a device code for GitHub authentication.
|
||||
|
||||
|
|
@ -246,7 +266,7 @@ class Authenticator:
|
|||
json={"client_id": client_id, "scope": "read:user"},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
resp_json: Final = resp.json()
|
||||
resp_json: Final = _JSON_OBJECT.validate_python(resp.json())
|
||||
|
||||
required_fields: Final = ["device_code", "user_code", "verification_uri"]
|
||||
if not all(field in resp_json for field in required_fields):
|
||||
|
|
@ -256,7 +276,7 @@ class Authenticator:
|
|||
status_code=400,
|
||||
)
|
||||
|
||||
return resp_json
|
||||
return _DEVICE_CODE.validate_python(resp_json)
|
||||
except httpx.HTTPStatusError as e:
|
||||
verbose_logger.error("HTTP error getting device code: %s", e)
|
||||
raise GetDeviceCodeError(
|
||||
|
|
@ -307,12 +327,13 @@ class Authenticator:
|
|||
},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
resp_json = resp.json()
|
||||
resp_json = _JSON_OBJECT.validate_python(resp.json())
|
||||
poll_result = _ACCESS_TOKEN_POLL.validate_python(resp_json)
|
||||
|
||||
if "access_token" in resp_json:
|
||||
if "access_token" in poll_result:
|
||||
verbose_logger.info("Authentication successful!")
|
||||
return resp_json["access_token"]
|
||||
elif "error" in resp_json and resp_json.get("error") == "authorization_pending":
|
||||
return poll_result["access_token"]
|
||||
elif "error" in poll_result and poll_result.get("error") == "authorization_pending":
|
||||
verbose_logger.debug("Authorization pending (attempt %s/%s)", attempt + 1, max_attempts)
|
||||
else:
|
||||
verbose_logger.warning("Unexpected response: %s", resp_json)
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import uuid
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
import httpx
|
||||
from pydantic import ConfigDict, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -29,6 +30,7 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
MANUS_API_BASE: Final = "https://api.manus.im"
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(strict=True))
|
||||
|
||||
|
||||
class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
|
|
@ -177,7 +179,7 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
original_response=raw_response.text,
|
||||
additional_args={"complete_input_dict": {}},
|
||||
)
|
||||
raw_response_json: Final = raw_response.json()
|
||||
raw_response_json: Final = _JSON_OBJECT.validate_python(raw_response.json())
|
||||
|
||||
# Manus uses camelCase "createdAt" instead of snake_case "created_at"
|
||||
if "createdAt" in raw_response_json and "created_at" not in raw_response_json:
|
||||
|
|
@ -269,7 +271,7 @@ class ManusResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
original_response=raw_response.text,
|
||||
additional_args={"complete_input_dict": {}},
|
||||
)
|
||||
raw_response_json: Final = raw_response.json()
|
||||
raw_response_json: Final = _JSON_OBJECT.validate_python(raw_response.json())
|
||||
|
||||
# Manus uses camelCase "createdAt" instead of snake_case "created_at"
|
||||
if "createdAt" in raw_response_json and "created_at" not in raw_response_json:
|
||||
|
|
|
|||
|
|
@ -4,7 +4,8 @@ from collections.abc import AsyncIterator, Iterator
|
|||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from httpx._models import Headers, Response
|
||||
from pydantic import ConfigDict, ValidationError
|
||||
from pydantic import ConfigDict, TypeAdapter, ValidationError, with_config
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -76,6 +77,21 @@ class _OllamaGenerateReasoning(LiteLLMBaseModel):
|
|||
return parse_content_for_reasoning(self.response)
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _OllamaGenerateMessage(TypedDict):
|
||||
content: ReadOnly[NotRequired[str]]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _OllamaGenerateResponse(TypedDict):
|
||||
message: ReadOnly[NotRequired[_OllamaGenerateMessage]]
|
||||
prompt_eval_count: ReadOnly[NotRequired[int]]
|
||||
eval_count: ReadOnly[NotRequired[int]]
|
||||
|
||||
|
||||
_OLLAMA_GENERATE_RESPONSE: Final = TypeAdapter(_OllamaGenerateResponse)
|
||||
|
||||
|
||||
class OllamaConfig(BaseConfig):
|
||||
"""
|
||||
Reference: https://github.com/ollama/ollama/blob/main/docs/api.md#parameters
|
||||
|
|
@ -159,7 +175,7 @@ class OllamaConfig(BaseConfig):
|
|||
system: str | None = None,
|
||||
template: str | None = None,
|
||||
) -> None:
|
||||
locals_: Final = locals().copy()
|
||||
locals_: Final[dict[str, object]] = locals().copy()
|
||||
for key, value in locals_.items():
|
||||
if key != "self" and value is not None:
|
||||
setattr(self.__class__, key, value)
|
||||
|
|
@ -288,7 +304,7 @@ class OllamaConfig(BaseConfig):
|
|||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
response_json: Final = raw_response.json()
|
||||
response_json: Final = _OLLAMA_GENERATE_RESPONSE.validate_python(raw_response.json())
|
||||
## RESPONSE OBJECT
|
||||
model_response.choices[0].finish_reason = "stop"
|
||||
if request_data.get("format", "") == "json":
|
||||
|
|
@ -303,7 +319,7 @@ class OllamaConfig(BaseConfig):
|
|||
model_response.choices[0].finish_reason = "stop"
|
||||
else:
|
||||
try:
|
||||
response_content: Final = json.loads(response_text)
|
||||
response_content: Final[object] = json.loads(response_text)
|
||||
|
||||
# Check if this is a function call format with name/arguments structure
|
||||
if (
|
||||
|
|
@ -356,7 +372,7 @@ class OllamaConfig(BaseConfig):
|
|||
len(tokenizer.encode(_prompt, disallowed_special=())),
|
||||
)
|
||||
completion_tokens: Final = response_json.get(
|
||||
"eval_count", len(response_json.get("message", dict()).get("content", ""))
|
||||
"eval_count", len((response_json.get("message") or {}).get("content", ""))
|
||||
)
|
||||
setattr(
|
||||
model_response,
|
||||
|
|
|
|||
|
|
@ -4,6 +4,8 @@ import time
|
|||
from collections.abc import Callable
|
||||
from typing import Final
|
||||
|
||||
from pydantic import ConfigDict, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.constants import REPLICATE_POLLING_DELAY_SECONDS
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
|
@ -20,6 +22,7 @@ from ..common_utils import ReplicateError
|
|||
from .transformation import ReplicateConfig
|
||||
|
||||
replicate_config: Final = ReplicateConfig()
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, object], config=ConfigDict(strict=True))
|
||||
|
||||
|
||||
# Function to handle prediction response (streaming)
|
||||
|
|
@ -81,7 +84,7 @@ async def async_handle_prediction_response_streaming(
|
|||
await asyncio.sleep(REPLICATE_POLLING_DELAY_SECONDS) # prevent being rate limited by replicate
|
||||
response = await http_client.get(prediction_url, headers=headers)
|
||||
if response.status_code == 200:
|
||||
response_data = response.json()
|
||||
response_data = _JSON_OBJECT.validate_python(response.json())
|
||||
status = response_data.get("status", "")
|
||||
# Check that "output" exists and is not None or empty
|
||||
output_present = "output" in response_data and response_data["output"] is not None
|
||||
|
|
@ -211,7 +214,7 @@ def completion(
|
|||
litellm.DEFAULT_REPLICATE_POLLING_DELAY_SECONDS + 2 * retry
|
||||
) # wait to allow response to be generated by replicate - else partial output is generated with status=="processing"
|
||||
response = httpx_client.get(url=prediction_url, headers=headers)
|
||||
if response.status_code == 200 and response.json().get("status") in [
|
||||
if response.status_code == 200 and _JSON_OBJECT.validate_python(response.json()).get("status") in [
|
||||
"processing",
|
||||
"starting",
|
||||
]:
|
||||
|
|
@ -280,7 +283,7 @@ async def async_completion(
|
|||
litellm.DEFAULT_REPLICATE_POLLING_DELAY_SECONDS + 2 * retry
|
||||
) # wait to allow response to be generated by replicate - else partial output is generated with status=="processing"
|
||||
response = await async_handler.get(url=prediction_url, headers=headers)
|
||||
if response.status_code == 200 and response.json().get("status") in [
|
||||
if response.status_code == 200 and _JSON_OBJECT.validate_python(response.json()).get("status") in [
|
||||
"processing",
|
||||
"starting",
|
||||
]:
|
||||
|
|
|
|||
|
|
@ -3209,7 +3209,7 @@ class ModelResponseIterator:
|
|||
def _apply_stream_usage_metadata(
|
||||
self,
|
||||
processed_chunk: GenerateContentResponseBody,
|
||||
model_response: Any,
|
||||
model_response: "ModelResponseStream",
|
||||
grounding_metadata: list[dict],
|
||||
) -> Usage | None:
|
||||
if "usageMetadata" not in processed_chunk:
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import copy
|
|||
import functools
|
||||
import os
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator, Iterator, Mapping, Sequence
|
||||
from collections.abc import AsyncGenerator, AsyncIterable, Iterator, Mapping, Sequence
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any, # noqa: TID251 # **kwargs forwards verbatim to CustomGuardrail.__init__
|
||||
|
|
@ -10,9 +10,11 @@ from typing import (
|
|||
Final,
|
||||
Literal,
|
||||
Optional,
|
||||
Protocol,
|
||||
)
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
|
|
@ -80,6 +82,8 @@ _VAULT_PREFIX: Final = f"litellm-{uuid.uuid4().hex}"
|
|||
|
||||
_DEFAULT_TIMEOUT_SECONDS: Final = 10.0
|
||||
|
||||
_SHIELD_BODY: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
def _string_tool_arguments(calls: Sequence[object]) -> Iterator[tuple[object, str]]:
|
||||
for call in calls:
|
||||
|
|
@ -89,6 +93,15 @@ def _string_tool_arguments(calls: Sequence[object]) -> Iterator[tuple[object, st
|
|||
yield function, arguments
|
||||
|
||||
|
||||
class _ContentDelta(Protocol):
|
||||
content: object
|
||||
|
||||
|
||||
class _FinishingDelta(Protocol):
|
||||
content: object
|
||||
tool_calls: object
|
||||
|
||||
|
||||
class LLMShieldProxyGuardrail(CustomGuardrail):
|
||||
"""Redacts PII before it leaves the proxy and restores it in the response.
|
||||
|
||||
|
|
@ -215,7 +228,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
|||
timeout=_DEFAULT_TIMEOUT_SECONDS,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
return _SHIELD_BODY.validate_python(response.json(), strict=True)
|
||||
except httpx.HTTPStatusError as exc:
|
||||
verbose_proxy_logger.exception("LLM Shield Proxy returned %s for %s", exc.response.status_code, path)
|
||||
raise GuardrailRaisedException(
|
||||
|
|
@ -329,8 +342,8 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
|||
self,
|
||||
data: MutableRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
) -> Any:
|
||||
response: object,
|
||||
) -> object:
|
||||
"""Restores the original values in a copy of a non-streaming response.
|
||||
|
||||
The copy is what keeps plaintext out of the response cache. LiteLLM caches the
|
||||
|
|
@ -345,8 +358,9 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
|||
return response
|
||||
response = detached(response) # rebind-ok: everything below restores the copy.
|
||||
|
||||
if self._is_anthropic_message_response(response):
|
||||
return await self._restore_anthropic_response(response, data)
|
||||
anthropic_body: Final = as_object(response)
|
||||
if anthropic_body is not None and self._is_anthropic_message_response(anthropic_body):
|
||||
return await self._restore_anthropic_response(anthropic_body, data)
|
||||
|
||||
response_slots: Final = self._responses_api_slots(response)
|
||||
if response_slots:
|
||||
|
|
@ -448,7 +462,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
|||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
response: AsyncIterable[object],
|
||||
request_data: MutableRequest,
|
||||
) -> AsyncGenerator[Any, None]:
|
||||
"""Restores original values incrementally, without buffering the stream.
|
||||
|
|
@ -545,7 +559,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
|||
|
||||
async def _restore_content_window(
|
||||
self,
|
||||
delta: Any,
|
||||
delta: _ContentDelta,
|
||||
key: CarryKey,
|
||||
carries: CarryWindows,
|
||||
session_id: str,
|
||||
|
|
@ -597,7 +611,7 @@ class LLMShieldProxyGuardrail(CustomGuardrail):
|
|||
|
||||
async def _flush_finished_choice(
|
||||
self,
|
||||
delta: Any,
|
||||
delta: _FinishingDelta,
|
||||
choice_index: int,
|
||||
carries: CarryWindows,
|
||||
session_id: str,
|
||||
|
|
|
|||
|
|
@ -16,10 +16,11 @@ from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator, Awaita
|
|||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypedDict, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, cast
|
||||
|
||||
import aiohttp
|
||||
from typing_extensions import NotRequired, ReadOnly
|
||||
from pydantic import ConfigDict, TypeAdapter, with_config
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm import get_secret
|
||||
|
|
@ -68,15 +69,22 @@ from litellm.utils import (
|
|||
)
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _PresidioAnonymizeItem(TypedDict, total=False):
|
||||
entity_type: ReadOnly[str | None]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _PresidioAnonymizeResponse(TypedDict):
|
||||
text: ReadOnly[str]
|
||||
items: ReadOnly[NotRequired[list[_PresidioAnonymizeItem]]]
|
||||
|
||||
|
||||
_PRESIDIO_ANONYMIZE_ADAPTER: Final[TypeAdapter[_PresidioAnonymizeResponse | None]] = TypeAdapter(
|
||||
_PresidioAnonymizeResponse | None
|
||||
)
|
||||
|
||||
|
||||
class _JsonResponse(Protocol):
|
||||
def json(self) -> Awaitable[object]: ...
|
||||
|
||||
|
|
@ -769,7 +777,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
raise Exception(
|
||||
f"Presidio anonymizer returned non-JSON Content-Type '{content_type}'; body: '{error_body[:200]}'"
|
||||
)
|
||||
return await response.json()
|
||||
return _PRESIDIO_ANONYMIZE_ADAPTER.validate_python(await response.json())
|
||||
|
||||
def _finalize_presidio_anonymize_simple(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -28,6 +28,8 @@ from datetime import datetime
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
|
||||
from fastapi.exceptions import HTTPException
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict, Unpack
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
|
|
@ -76,6 +78,12 @@ _DEFAULT_POLICIES: Final = [
|
|||
"Default_Policy_GeneralPromptAttackProtection",
|
||||
]
|
||||
|
||||
_RESPONSE_BODY: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object):
|
||||
supported_event_hooks: ReadOnly[list[GuardrailEventHooks]]
|
||||
|
||||
|
||||
class XecGuardMissingCredentials(Exception):
|
||||
pass
|
||||
|
|
@ -90,7 +98,7 @@ class XecGuardGuardrail(CustomGuardrail):
|
|||
policy_names: list[str] | None = None,
|
||||
block_on_error: bool | None = None,
|
||||
grounding_strictness: str | None = None,
|
||||
**kwargs: Any,
|
||||
**kwargs: Unpack[_CustomGuardrailOptions],
|
||||
) -> None:
|
||||
self.api_key = api_key or os.environ.get("XECGUARD_API_KEY")
|
||||
if not self.api_key:
|
||||
|
|
@ -122,9 +130,12 @@ class XecGuardGuardrail(CustomGuardrail):
|
|||
llm_provider=httpxSpecialProvider.GuardrailCallback,
|
||||
)
|
||||
|
||||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
forwarded: Final[_CustomGuardrailOptions] = {
|
||||
"supported_event_hooks": list(self.get_supported_event_hooks()),
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
super().__init__(**kwargs)
|
||||
super().__init__(**forwarded)
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> type["GuardrailConfigModel"] | None:
|
||||
|
|
@ -345,7 +356,7 @@ class XecGuardGuardrail(CustomGuardrail):
|
|||
path: str,
|
||||
payload: dict,
|
||||
suppress_errors: bool = False,
|
||||
) -> dict | None:
|
||||
) -> dict[str, object] | None:
|
||||
endpoint: Final = f"{self.api_base}{path}"
|
||||
verbose_proxy_logger.debug(
|
||||
"XecGuard: POST %s payload_keys=%s",
|
||||
|
|
@ -363,7 +374,7 @@ class XecGuardGuardrail(CustomGuardrail):
|
|||
timeout=self.timeout if self.timeout is not None else 10.0,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
return _RESPONSE_BODY.validate_python(response.json(), strict=True)
|
||||
except Exception as exc:
|
||||
verbose_proxy_logger.error("XecGuard API error: %s", str(exc))
|
||||
if suppress_errors:
|
||||
|
|
|
|||
|
|
@ -29,6 +29,9 @@ import json
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
|
||||
from pydantic import ConfigDict, TypeAdapter, with_config
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
|
|
@ -81,6 +84,22 @@ class _ChatCompletion(Protocol):
|
|||
def choices(self) -> Sequence[_ChatChoice]: ...
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _CodeExecutionToolArgs(TypedDict, total=False):
|
||||
code: ReadOnly[str]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _GeneratedFileEntry(TypedDict):
|
||||
name: ReadOnly[str]
|
||||
mime_type: ReadOnly[str]
|
||||
content_base64: ReadOnly[str]
|
||||
|
||||
|
||||
_TOOL_CALL_ARGS_ADAPTER: Final = TypeAdapter(_CodeExecutionToolArgs)
|
||||
_EXEC_RESULT_FILES_ADAPTER: Final = TypeAdapter(Sequence[_GeneratedFileEntry])
|
||||
|
||||
|
||||
def _first_choice(response: _ChatCompletion) -> _ChatChoice:
|
||||
"""The first choice of an OpenAI shaped completion response."""
|
||||
return response.choices[0]
|
||||
|
|
@ -620,7 +639,7 @@ class SkillsInjectionHook(CustomLogger):
|
|||
|
||||
# Collect generated files
|
||||
if exec_result.get("files"):
|
||||
files: Final[Sequence[Mapping[str, str]]] = exec_result["files"]
|
||||
files: Final = _EXEC_RESULT_FILES_ADAPTER.validate_python(exec_result["files"])
|
||||
for f in files:
|
||||
generated_files.append(
|
||||
{
|
||||
|
|
@ -833,7 +852,7 @@ print('No executable skill module found')
|
|||
) -> str:
|
||||
"""Execute a litellm_code_execution tool call and return result string."""
|
||||
try:
|
||||
args: Final[Mapping[str, str]] = json.loads(tool_call.function.arguments)
|
||||
args: Final = _TOOL_CALL_ARGS_ADAPTER.validate_python(json.loads(tool_call.function.arguments))
|
||||
code: Final[str] = args.get("code", "")
|
||||
|
||||
verbose_proxy_logger.debug("SkillsInjectionHook: Executing code (%s chars)", len(code))
|
||||
|
|
@ -849,7 +868,7 @@ print('No executable skill module found')
|
|||
# Collect generated files
|
||||
if exec_result.get("files"):
|
||||
tool_result += "\n\nGenerated files:"
|
||||
files: Final[Sequence[Mapping[str, str]]] = exec_result["files"]
|
||||
files: Final = _EXEC_RESULT_FILES_ADAPTER.validate_python(exec_result["files"])
|
||||
for f in files:
|
||||
file_content = base64.b64decode(f["content_base64"])
|
||||
generated_files.append(
|
||||
|
|
|
|||
|
|
@ -39,7 +39,6 @@ from typing import (
|
|||
Optional,
|
||||
Protocol,
|
||||
TypeAlias,
|
||||
TypedDict,
|
||||
Union,
|
||||
cast,
|
||||
get_args,
|
||||
|
|
@ -50,9 +49,9 @@ from typing import (
|
|||
import anyio
|
||||
import websockets
|
||||
import websockets.exceptions
|
||||
from pydantic import BaseModel, Json, JsonValue, TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, Json, JsonValue, TypeAdapter, ValidationError, with_config
|
||||
from pydantic.fields import FieldInfo, PydanticUndefined
|
||||
from typing_extensions import NotRequired, ReadOnly, assert_never
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict, assert_never
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import (
|
||||
|
|
@ -1191,6 +1190,34 @@ async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = N
|
|||
_AiohttpAddrInfo: TypeAlias = tuple[int | socket.AddressFamily, int | socket.SocketKind, int, str, tuple[object, ...]]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _LoginRequestBody(TypedDict, total=False):
|
||||
username: ReadOnly[object]
|
||||
password: ReadOnly[object]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _LoginExchangeRequestBody(TypedDict, total=False):
|
||||
code: ReadOnly[object]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _AnthropicBetaHeadersReloadConfig(TypedDict, total=False):
|
||||
interval_hours: ReadOnly[int | float | None]
|
||||
force_reload: ReadOnly[object]
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True))
|
||||
class _WebsearchInterceptionLitellmSettings(TypedDict, total=False):
|
||||
websearch_interception_params: ReadOnly[object]
|
||||
|
||||
|
||||
_LOGIN_REQUEST_BODY: Final = TypeAdapter(_LoginRequestBody)
|
||||
_LOGIN_EXCHANGE_REQUEST_BODY: Final = TypeAdapter(_LoginExchangeRequestBody)
|
||||
_ANTHROPIC_BETA_HEADERS_RELOAD_CONFIG: Final = TypeAdapter(_AnthropicBetaHeadersReloadConfig)
|
||||
_WEBSEARCH_INTERCEPTION_LITELLM_SETTINGS: Final = TypeAdapter(_WebsearchInterceptionLitellmSettings)
|
||||
|
||||
|
||||
class _AiohttpConnectorKwargs(TypedDict, total=False):
|
||||
keepalive_timeout: float
|
||||
ttl_dns_cache: int
|
||||
|
|
@ -8471,9 +8498,10 @@ class ProxyConfig:
|
|||
if config_record is None or config_record.param_value is None:
|
||||
return
|
||||
|
||||
litellm_settings = config_record.param_value
|
||||
if isinstance(litellm_settings, str):
|
||||
litellm_settings = json.loads(litellm_settings)
|
||||
raw_litellm_settings: Final = config_record.param_value
|
||||
litellm_settings: Final = _WEBSEARCH_INTERCEPTION_LITELLM_SETTINGS.validate_python(
|
||||
json.loads(raw_litellm_settings) if isinstance(raw_litellm_settings, str) else raw_litellm_settings
|
||||
)
|
||||
|
||||
websearch_config: Final = litellm_settings.get("websearch_interception_params", None)
|
||||
|
||||
|
|
@ -8726,7 +8754,7 @@ class ProxyConfig:
|
|||
if config_record is None or config_record.param_value is None:
|
||||
return # No configuration found, skip reload
|
||||
|
||||
config: Final = config_record.param_value
|
||||
config: Final = _ANTHROPIC_BETA_HEADERS_RELOAD_CONFIG.validate_python(config_record.param_value)
|
||||
interval_hours: Final = config.get("interval_hours")
|
||||
force_reload: Final = config.get("force_reload", False)
|
||||
|
||||
|
|
@ -17127,7 +17155,7 @@ async def login_v2(request: Request):
|
|||
from litellm.proxy.utils import get_custom_url
|
||||
|
||||
try:
|
||||
body: Final = await request.json()
|
||||
body: Final = _LOGIN_REQUEST_BODY.validate_python(await request.json())
|
||||
username: Final = str(body.get("username"))
|
||||
password: Final = str(body.get("password"))
|
||||
|
||||
|
|
@ -17199,7 +17227,7 @@ async def login_v3(request: Request):
|
|||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
body: Final = await request.json()
|
||||
body: Final = _LOGIN_REQUEST_BODY.validate_python(await request.json())
|
||||
username: Final = str(body.get("username"))
|
||||
password: Final = str(body.get("password"))
|
||||
|
||||
|
|
@ -17271,7 +17299,7 @@ async def login_v3_exchange(request: Request):
|
|||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
body: Final = await request.json()
|
||||
body: Final = _LOGIN_EXCHANGE_REQUEST_BODY.validate_python(await request.json())
|
||||
code: Final = body.get("code")
|
||||
if not code:
|
||||
raise ProxyException(
|
||||
|
|
|
|||
|
|
@ -9335,10 +9335,10 @@ class Router:
|
|||
if access_windows_error is not None:
|
||||
raise ValueError(access_windows_error)
|
||||
zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None
|
||||
merged_params: Final[Mapping[str, Any]] = (
|
||||
merged_params: Final[Mapping[str, object]] = (
|
||||
_litellm_params if zeroed_pricing is None else MappingProxyType({**_litellm_params, **zeroed_pricing})
|
||||
)
|
||||
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(**merged_params)
|
||||
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params.model_validate(dict(**merged_params))
|
||||
warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params)
|
||||
deployment = Deployment(
|
||||
**deployment_info,
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ from httpx import Proxy
|
|||
from httpx._utils import get_environment_proxies
|
||||
from openai.lib import _parsing, _pydantic # pyright: ignore[reportPrivateUsage] # OpenAI parser module is private
|
||||
from openai.types.chat.completion_create_params import ResponseFormat
|
||||
from pydantic import BaseModel, JsonValue
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, with_config
|
||||
|
||||
import litellm
|
||||
import litellm.litellm_core_utils
|
||||
|
|
@ -5506,7 +5506,7 @@ def get_max_tokens(model: str) -> int | None:
|
|||
response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx)
|
||||
|
||||
# Parse the JSON response
|
||||
config_json: Final[Mapping[str, int]] = response.json()
|
||||
config_json: Final = _HUGGINGFACE_MODEL_CONFIG.validate_python(response.json())
|
||||
# Extract and return the max_position_embeddings
|
||||
max_position_embeddings: Final = config_json.get("max_position_embeddings")
|
||||
if max_position_embeddings is not None:
|
||||
|
|
@ -5776,6 +5776,14 @@ def _check_provider_match(model_info: dict, custom_llm_provider: str | None) ->
|
|||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow", strict=True, hide_input_in_errors=True))
|
||||
class _HuggingFaceModelConfig(TypedDict, total=False):
|
||||
max_position_embeddings: ReadOnly[int | None]
|
||||
|
||||
|
||||
_HUGGINGFACE_MODEL_CONFIG: Final = TypeAdapter(_HuggingFaceModelConfig)
|
||||
|
||||
|
||||
class PotentialModelNamesAndCustomLLMProvider(TypedDict):
|
||||
split_model: str
|
||||
combined_model_name: str
|
||||
|
|
@ -5923,7 +5931,7 @@ def _get_max_position_embeddings(model_name: str) -> int | None:
|
|||
response.raise_for_status() # Raise an exception for bad responses (4xx or 5xx)
|
||||
|
||||
# Parse the JSON response
|
||||
config_json: Final[Mapping[str, int]] = response.json()
|
||||
config_json: Final = _HUGGINGFACE_MODEL_CONFIG.validate_python(response.json())
|
||||
|
||||
# Extract and return the max_position_embeddings
|
||||
max_position_embeddings: Final = config_json.get("max_position_embeddings")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue