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