mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
325e54390f
commit
7294bc1111
3 changed files with 119 additions and 42 deletions
|
|
@ -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
54
litellm/types/main.py
Normal 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,
|
||||
]
|
||||
61
tests/test_litellm/types/test_main.py
Normal file
61
tests/test_litellm/types/test_main.py
Normal 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__")
|
||||
Loading…
Add table
Reference in a new issue