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.
This commit is contained in:
mateo-berri 2026-06-22 15:18:33 +00:00
parent 325e54390f
commit 7294bc1111
No known key found for this signature in database
3 changed files with 119 additions and 42 deletions

View file

@ -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

54
litellm/types/main.py Normal file
View file

@ -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,
]

View file

@ -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__")