mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(types): replace Any with proven types in 4 files (#44370)
* refactor(types): replace Any with proven types in 8 files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): revert unproven email logger protocol and iterator annotation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): revert unproven deepagents and http handler annotations --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e340e546e2
commit
e768ad55ce
4 changed files with 76 additions and 51 deletions
|
|
@ -18,7 +18,7 @@ import secrets
|
|||
from collections.abc import AsyncIterable, AsyncIterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType, ModuleType
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
|
||||
import httpx
|
||||
import openai
|
||||
|
|
@ -41,7 +41,8 @@ from litellm.types.llms.custom_http import httpxSpecialProvider
|
|||
if TYPE_CHECKING:
|
||||
from starlette.applications import Starlette
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
from starlette.responses import JSONResponse, Response, StreamingResponse
|
||||
from starlette.routing import Route
|
||||
from uvicorn import Server
|
||||
|
||||
verbose_logger: Final = logging.getLogger("LiteLLM")
|
||||
|
|
@ -85,12 +86,31 @@ COST_HEADER = "x-litellm-response-cost"
|
|||
SSE_MEDIA_TYPE = "text/event-stream"
|
||||
|
||||
|
||||
class _ApplicationsModule(Protocol):
|
||||
"""The starlette.applications attributes the endpoint uses."""
|
||||
|
||||
Starlette: type[Starlette]
|
||||
|
||||
|
||||
class _RoutingModule(Protocol):
|
||||
"""The starlette.routing attributes the endpoint uses."""
|
||||
|
||||
Route: type[Route]
|
||||
|
||||
|
||||
class _ResponsesModule(Protocol):
|
||||
"""The starlette.responses attributes the endpoint uses."""
|
||||
|
||||
JSONResponse: type[JSONResponse]
|
||||
StreamingResponse: type[StreamingResponse]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ServerDeps:
|
||||
uvicorn: ModuleType
|
||||
applications: ModuleType
|
||||
routing: ModuleType
|
||||
responses: ModuleType
|
||||
applications: _ApplicationsModule
|
||||
routing: _RoutingModule
|
||||
responses: _ResponsesModule
|
||||
|
||||
|
||||
def _load_server_deps() -> _ServerDeps:
|
||||
|
|
@ -193,7 +213,7 @@ class SSEUsageParser:
|
|||
if isinstance(event, Mapping):
|
||||
self.absorb(event)
|
||||
|
||||
def absorb(self, event: Mapping[str, Any]) -> None:
|
||||
def absorb(self, event: Mapping[str, object]) -> None:
|
||||
event_type = event.get("type")
|
||||
if event_type == "message_start":
|
||||
self._absorb_message_start(event)
|
||||
|
|
@ -204,16 +224,16 @@ class SSEUsageParser:
|
|||
elif isinstance(event.get("usage"), Mapping):
|
||||
self._set(*usage_from_mapping(event["usage"]))
|
||||
|
||||
def _absorb_message_start(self, event: Mapping[str, Any]) -> None:
|
||||
def _absorb_message_start(self, event: Mapping[str, object]) -> None:
|
||||
message = event.get("message")
|
||||
if isinstance(message, Mapping):
|
||||
self._set(*usage_from_mapping(message.get("usage")))
|
||||
|
||||
def _absorb_message_delta(self, event: Mapping[str, Any]) -> None:
|
||||
def _absorb_message_delta(self, event: Mapping[str, object]) -> None:
|
||||
# message_delta output_tokens is cumulative for the whole message.
|
||||
self._set(*usage_from_mapping(event.get("usage")))
|
||||
|
||||
def _absorb_response_completed(self, event: Mapping[str, Any]) -> None:
|
||||
def _absorb_response_completed(self, event: Mapping[str, object]) -> None:
|
||||
response = event.get("response")
|
||||
if isinstance(response, Mapping):
|
||||
self._set(*usage_from_mapping(response.get("usage")))
|
||||
|
|
@ -271,7 +291,7 @@ def gateway_headers(
|
|||
incoming: Mapping[str, str],
|
||||
gateway: GatewayTarget,
|
||||
harness: Harness,
|
||||
metadata: Mapping[str, Any] | None,
|
||||
metadata: Mapping[str, object] | None,
|
||||
) -> Mapping[str, str]:
|
||||
"""Incoming headers minus hop-by-hop/auth/x-litellm-*, plus gateway auth, tags, metadata."""
|
||||
kept = (
|
||||
|
|
@ -309,7 +329,7 @@ def error_status(exc: BaseException) -> int:
|
|||
return 500
|
||||
|
||||
|
||||
def error_body(exc: BaseException, message: str) -> dict[str, Any]: # mutable-ok: JSONResponse body
|
||||
def error_body(exc: BaseException, message: str) -> dict[str, dict[str, str]]: # mutable-ok: JSONResponse body
|
||||
return {"error": {"type": type(exc).__name__, "message": message}} # mutable-ok: JSONResponse body
|
||||
|
||||
|
||||
|
|
@ -373,7 +393,7 @@ class ModelEndpoint:
|
|||
gateway: GatewayTarget | None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
*,
|
||||
client: httpx.AsyncClient | None = None,
|
||||
) -> None:
|
||||
|
|
@ -382,14 +402,14 @@ class ModelEndpoint:
|
|||
self.gateway = gateway
|
||||
self.api_key = api_key
|
||||
self.api_base = api_base
|
||||
self.metadata: Mapping[str, Any] = MappingProxyType(dict(metadata or ()))
|
||||
self.metadata: Mapping[str, object] = MappingProxyType(dict(metadata or ()))
|
||||
self.token = secrets.token_urlsafe(HARNESS_SESSION_TOKEN_BYTES)
|
||||
self.usage = UsageTracker()
|
||||
self.port = 0
|
||||
self._injected_client = client
|
||||
self._deps: _ServerDeps | None = None
|
||||
self._client: httpx.AsyncClient | None = None
|
||||
self._server: Any = None
|
||||
self._server: Server | None = None
|
||||
self._task: asyncio.Task[None] | None = None
|
||||
|
||||
@property
|
||||
|
|
@ -410,7 +430,7 @@ class ModelEndpoint:
|
|||
self._server = self._build_server(self._deps)
|
||||
self._task = asyncio.create_task(self._server.serve())
|
||||
try:
|
||||
await asyncio.wait_for(self._wait_started(), HARNESS_ENDPOINT_STARTUP_TIMEOUT_SECONDS)
|
||||
await asyncio.wait_for(self._wait_started(self._server), HARNESS_ENDPOINT_STARTUP_TIMEOUT_SECONDS)
|
||||
except BaseException:
|
||||
await self.stop()
|
||||
raise
|
||||
|
|
@ -439,8 +459,8 @@ class ModelEndpoint:
|
|||
)
|
||||
return handler.client
|
||||
|
||||
async def _wait_started(self) -> None:
|
||||
while not self._server.started:
|
||||
async def _wait_started(self, server: Server) -> None:
|
||||
while not server.started:
|
||||
if self._task is not None and self._task.done():
|
||||
raise HarnessError("harness model endpoint failed to start")
|
||||
await asyncio.sleep(DEFAULT_POLLING_INTERVAL)
|
||||
|
|
@ -483,7 +503,7 @@ class ModelEndpoint:
|
|||
)
|
||||
|
||||
@property
|
||||
def _responses(self) -> ModuleType:
|
||||
def _responses(self) -> _ResponsesModule:
|
||||
if self._deps is None:
|
||||
raise HarnessError("harness model endpoint is not started")
|
||||
return self._deps.responses
|
||||
|
|
@ -530,7 +550,7 @@ class ModelEndpoint:
|
|||
return await self._forward(request, route, body)
|
||||
return await self._call_sdk(route, body)
|
||||
|
||||
def _cost_model(self, body: Mapping[str, Any]) -> str | None:
|
||||
def _cost_model(self, body: Mapping[str, object]) -> str | None:
|
||||
model = self.model or body.get("model")
|
||||
return model if isinstance(model, str) else None
|
||||
|
||||
|
|
@ -545,7 +565,7 @@ class ModelEndpoint:
|
|||
cost = compute_cost(model, input_tokens, output_tokens)
|
||||
self.usage.add(input_tokens, output_tokens, cost)
|
||||
|
||||
async def _forward(self, request: Request, route: str, body: Mapping[str, Any]) -> Response:
|
||||
async def _forward(self, request: Request, route: str, body: Mapping[str, object]) -> Response:
|
||||
if self._client is None or self.gateway is None:
|
||||
raise HarnessError("gateway client is not started")
|
||||
if self.model:
|
||||
|
|
@ -601,7 +621,7 @@ class ModelEndpoint:
|
|||
self._record(model, tokens[0], tokens[1], header_cost(upstream.headers))
|
||||
|
||||
def _sdk_kwargs(
|
||||
self, body: Mapping[str, Any]
|
||||
self, body: Mapping[str, object]
|
||||
) -> dict[str, Any]: # mutable-ok: SDK call kwargs, mutated by _invoke_sdk then splatted
|
||||
kwargs: dict[str, Any] = {**body} # mutable-ok: SDK call kwargs built from the JSON body, then overridden
|
||||
if self.model:
|
||||
|
|
@ -629,7 +649,7 @@ class ModelEndpoint:
|
|||
return await litellm.acompletion(**kwargs)
|
||||
return await litellm.aresponses(**kwargs)
|
||||
|
||||
async def _call_sdk(self, route: str, body: Mapping[str, Any]) -> Response:
|
||||
async def _call_sdk(self, route: str, body: Mapping[str, object]) -> Response:
|
||||
kwargs = self._sdk_kwargs(body)
|
||||
model = self._cost_model(kwargs)
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -79,7 +79,7 @@ class SessionConfig:
|
|||
api_key: str | None = None
|
||||
api_base: str | None = None
|
||||
instructions: str | None = None
|
||||
tools: Sequence[Callable[..., Any]] = ()
|
||||
tools: Sequence[Callable[..., object]] = ()
|
||||
skills: Sequence[str] = ()
|
||||
disable_tools: Sequence[str] = ()
|
||||
permissions: PermissionMode = "full"
|
||||
|
|
@ -87,7 +87,7 @@ class SessionConfig:
|
|||
output: type[BaseModel] | None = None
|
||||
max_turns: int | None = None
|
||||
timeout: float | None = None
|
||||
metadata: Mapping[str, Any] = field(default_factory=dict)
|
||||
metadata: Mapping[str, object] = field(default_factory=dict)
|
||||
options: HarnessOptions | None = None
|
||||
install: bool = False
|
||||
|
||||
|
|
@ -184,7 +184,7 @@ def build_config(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -192,7 +192,7 @@ def build_config(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> SessionConfig:
|
||||
|
|
@ -335,7 +335,7 @@ async def call_approval_handler(handler: ApprovalHandler, approval: Approval) ->
|
|||
"""Run on_approval (sync in a worker thread, or async) and resolve approval."""
|
||||
try:
|
||||
if inspect.iscoroutinefunction(handler):
|
||||
decision: Any = await handler(approval)
|
||||
decision: object = await handler(approval)
|
||||
else:
|
||||
decision = await asyncio.to_thread(handler, approval)
|
||||
if inspect.isawaitable(decision):
|
||||
|
|
@ -794,7 +794,7 @@ def aagent_session(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -802,7 +802,7 @@ def aagent_session(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> AsyncSession:
|
||||
|
|
@ -838,7 +838,7 @@ async def arun_agent(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -846,7 +846,7 @@ async def arun_agent(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Result:
|
||||
|
|
@ -883,7 +883,7 @@ def astream_agent(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -891,7 +891,7 @@ def astream_agent(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> AsyncEventStream:
|
||||
|
|
@ -935,7 +935,7 @@ def aagent_resume(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -943,7 +943,7 @@ def aagent_resume(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> AsyncSession:
|
||||
|
|
@ -988,7 +988,7 @@ def aagent(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -996,10 +996,10 @@ def aagent(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Coroutine[Any, Any, Result] | AsyncEventStream:
|
||||
) -> Coroutine[object, object, Result] | AsyncEventStream:
|
||||
"""Run an agent harness on one prompt.
|
||||
|
||||
`await litellm.aagent(...)` returns a Result. With stream=True it returns an async
|
||||
|
|
|
|||
|
|
@ -62,7 +62,7 @@ class _LoopThread:
|
|||
self._thread.start()
|
||||
return self._loop
|
||||
|
||||
def submit(self, coro: Coroutine[Any, Any, T]) -> Future[T]:
|
||||
def submit(self, coro: Coroutine[object, None, T]) -> Future[T]:
|
||||
return asyncio.run_coroutine_threadsafe(coro, self.loop())
|
||||
|
||||
|
||||
|
|
@ -77,7 +77,7 @@ def _ensure_sync_context(name: str) -> None:
|
|||
raise RuntimeError(IN_LOOP_MESSAGE.format(name=name))
|
||||
|
||||
|
||||
def run_sync(coro: Coroutine[Any, Any, T], name: str) -> T:
|
||||
def run_sync(coro: Coroutine[object, None, T], name: str) -> T:
|
||||
"""Run coro on the harness loop thread and block for its result."""
|
||||
try:
|
||||
_ensure_sync_context(name)
|
||||
|
|
@ -215,7 +215,7 @@ def _run(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -223,7 +223,7 @@ def _run(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Result:
|
||||
|
|
@ -262,7 +262,7 @@ def _stream(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -270,7 +270,7 @@ def _stream(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> EventStream:
|
||||
|
|
@ -307,7 +307,7 @@ def agent_session(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -315,7 +315,7 @@ def agent_session(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Session:
|
||||
|
|
@ -352,7 +352,7 @@ def agent_resume(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -360,7 +360,7 @@ def agent_resume(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Session:
|
||||
|
|
@ -399,7 +399,7 @@ def agent(
|
|||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: Sequence[Callable[..., Any]] = (),
|
||||
tools: Sequence[Callable[..., object]] = (),
|
||||
skills: Sequence[str | os.PathLike[str]] = (),
|
||||
disable_tools: Sequence[str] = (),
|
||||
permissions: PermissionMode = "full",
|
||||
|
|
@ -407,7 +407,7 @@ def agent(
|
|||
output: type[BaseModel] | None = None,
|
||||
max_turns: int | None = None,
|
||||
timeout: float | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: HarnessOptions | None = None,
|
||||
install: bool = False,
|
||||
) -> Result | EventStream:
|
||||
|
|
|
|||
|
|
@ -994,7 +994,12 @@ def _responses_try_dispatch_mcp_gateway(
|
|||
kwargs: dict[str, object],
|
||||
_is_async: bool,
|
||||
skip_mcp_handler: bool,
|
||||
) -> Any | None:
|
||||
) -> (
|
||||
ResponsesAPIResponse
|
||||
| BaseResponsesAPIStreamingIterator
|
||||
| Coroutine[object, object, ResponsesAPIResponse | BaseResponsesAPIStreamingIterator]
|
||||
| None
|
||||
):
|
||||
"""Return a response when MCP gateway handles the call; otherwise None."""
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue