From 7294bc111187bb04204ca6e8e0177387915fcec1 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 22 Jun 2026 15:18:33 +0000 Subject: [PATCH] refactor(completion): move dispatch types into litellm/types/main.py Relocate the private _CompletionDispatchContext and _CompletionDispatchResult out of main.py into a dedicated litellm/types/main.py so the entrypoint file is not carrying the type classes inline, per review feedback. The definitions are moved verbatim, so the dispatch helpers and completion() are unchanged and the basedpyright/ruff budgets stay identical (the same annotations now live in the new module). Adds a mapped test pinning the frozen+slots invariant the dispatch shape relies on. --- litellm/main.py | 46 ++------------------ litellm/types/main.py | 54 ++++++++++++++++++++++++ tests/test_litellm/types/test_main.py | 61 +++++++++++++++++++++++++++ 3 files changed, 119 insertions(+), 42 deletions(-) create mode 100644 litellm/types/main.py create mode 100644 tests/test_litellm/types/test_main.py diff --git a/litellm/main.py b/litellm/main.py index 304b22e85e7..83d16643d85 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -22,7 +22,6 @@ import traceback from concurrent import futures from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait from copy import deepcopy -from dataclasses import dataclass from functools import partial from typing import ( TYPE_CHECKING, @@ -119,6 +118,10 @@ from litellm.llms.vertex_ai.common_utils import ( ) from litellm.realtime_api.main import _realtime_health_check from litellm.secret_managers.main import get_secret_bool, get_secret_str +from litellm.types.main import ( + _CompletionDispatchContext, + _CompletionDispatchResult, +) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ( CustomPricingLiteLLMParams, @@ -1085,47 +1088,6 @@ def _build_custom_pricing_entry( return entry -@dataclass(frozen=True, slots=True) -class _CompletionDispatchContext: - _azure_detection_model: str - acompletion: bool - api_base: Optional[str] - api_key: Optional[str] - api_version: Optional[str] - client: Any - custom_llm_provider: str - custom_prompt_dict: dict - extra_headers: Optional[dict] - headers: dict - hf_model_name: Optional[str] - kwargs: dict - litellm_params: dict - logger_fn: Optional[Callable] - logging: LiteLLMLoggingObj - max_retries: Optional[int] - max_tokens: Optional[int] - messages: list - metadata: Optional[dict] - model: str - model_response: ModelResponse - optional_params: dict - organization: Optional[str] - provider_config: Optional[BaseConfig] - shared_session: Optional["ClientSession"] - stream: Optional[bool] - temperature: Optional[float] - text_completion: bool - timeout: Optional[Union[float, str, httpx.Timeout]] - top_p: Optional[float] - - -_CompletionDispatchResult = Union[ - Coroutine[Any, Any, Union[ModelResponse, CustomStreamWrapper]], - ModelResponse, - CustomStreamWrapper, -] - - def _complete_azure(ctx: _CompletionDispatchContext) -> _CompletionDispatchResult: _azure_detection_model = ctx._azure_detection_model acompletion = ctx.acompletion diff --git a/litellm/types/main.py b/litellm/types/main.py new file mode 100644 index 00000000000..2636a1db543 --- /dev/null +++ b/litellm/types/main.py @@ -0,0 +1,54 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Callable, Coroutine, Optional, Union + +from litellm.utils import CustomStreamWrapper, ModelResponse + +if TYPE_CHECKING: + import httpx + from aiohttp import ClientSession + + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.base_llm import BaseConfig + + +@dataclass(frozen=True, slots=True) +class _CompletionDispatchContext: + _azure_detection_model: str + acompletion: bool + api_base: Optional[str] + api_key: Optional[str] + api_version: Optional[str] + client: Any + custom_llm_provider: str + custom_prompt_dict: dict + extra_headers: Optional[dict] + headers: dict + hf_model_name: Optional[str] + kwargs: dict + litellm_params: dict + logger_fn: Optional[Callable] + logging: LiteLLMLoggingObj + max_retries: Optional[int] + max_tokens: Optional[int] + messages: list + metadata: Optional[dict] + model: str + model_response: ModelResponse + optional_params: dict + organization: Optional[str] + provider_config: Optional[BaseConfig] + shared_session: Optional[ClientSession] + stream: Optional[bool] + temperature: Optional[float] + text_completion: bool + timeout: Optional[Union[float, str, httpx.Timeout]] + top_p: Optional[float] + + +_CompletionDispatchResult = Union[ + Coroutine[Any, Any, Union[ModelResponse, CustomStreamWrapper]], + ModelResponse, + CustomStreamWrapper, +] diff --git a/tests/test_litellm/types/test_main.py b/tests/test_litellm/types/test_main.py new file mode 100644 index 00000000000..a6c0cc9fa34 --- /dev/null +++ b/tests/test_litellm/types/test_main.py @@ -0,0 +1,61 @@ +import dataclasses +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.types.main import _CompletionDispatchContext + + +def _build_context() -> _CompletionDispatchContext: + return _CompletionDispatchContext( + _azure_detection_model="gpt-4o", + acompletion=False, + api_base=None, + api_key=None, + api_version=None, + client=None, + custom_llm_provider="openai", + custom_prompt_dict={}, + extra_headers=None, + headers={}, + hf_model_name=None, + kwargs={}, + litellm_params={}, + logger_fn=None, + logging=None, # type: ignore[arg-type] + max_retries=None, + max_tokens=None, + messages=[], + metadata=None, + model="gpt-4o", + model_response=None, # type: ignore[arg-type] + optional_params={}, + organization=None, + provider_config=None, + shared_session=None, + stream=None, + temperature=None, + text_completion=False, + timeout=None, + top_p=None, + ) + + +def test_dispatch_context_is_frozen(): + """A helper must not be able to re-route the call by rebinding a dispatch + input mid-flight; this pins the frozen invariant the dispatch shape relies on.""" + ctx = _build_context() + with pytest.raises(dataclasses.FrozenInstanceError): + ctx.model = "claude-haiku-4-5" # type: ignore[misc] + with pytest.raises(dataclasses.FrozenInstanceError): + ctx.custom_llm_provider = "anthropic" # type: ignore[misc] + + +def test_dispatch_context_uses_slots(): + """slots=True keeps the per-call context lightweight (no per-instance __dict__).""" + ctx = _build_context() + assert not hasattr(ctx, "__dict__") + assert hasattr(type(ctx), "__slots__")