diff --git a/litellm/harness/endpoint.py b/litellm/harness/endpoint.py index ce3586735c6..21e789ed1fc 100644 --- a/litellm/harness/endpoint.py +++ b/litellm/harness/endpoint.py @@ -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: diff --git a/litellm/harness/runtime.py b/litellm/harness/runtime.py index 4ec167a06b4..da90b28a751 100644 --- a/litellm/harness/runtime.py +++ b/litellm/harness/runtime.py @@ -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 diff --git a/litellm/harness/sync.py b/litellm/harness/sync.py index 1788543f9b5..1cb8fbcbca2 100644 --- a/litellm/harness/sync.py +++ b/litellm/harness/sync.py @@ -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: diff --git a/litellm/responses/main.py b/litellm/responses/main.py index d145f8cc6b8..14c12fc571d 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -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,