mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(types): serialize deferred pydantic schema builds across threads (#45034)
A deferred LiteLLM model whose first use happens on two threads at once could lose its freshly built validator to the second thread's rebuild, so GenericLiteLLMParams.model_validate handed back a CredentialLiteLLMParams and the request failed 400 on use_litellm_proxy. LiteLLMBaseModel.model_rebuild now runs under one process-wide re-entrant lock, so a thread arriving mid-build waits for the finished validator instead of rebuilding over it Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
a546a1720f
commit
d7c6c4b80f
2 changed files with 63 additions and 7 deletions
|
|
@ -1,3 +1,4 @@
|
|||
import threading
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
|
|
@ -6,6 +7,8 @@ from pydantic import BaseModel, ConfigDict
|
|||
|
||||
from litellm.constants import DEFER_PYDANTIC_BUILD
|
||||
|
||||
_SCHEMA_BUILD_LOCK: Final = threading.RLock()
|
||||
|
||||
|
||||
class LiteLLMBaseModel(BaseModel):
|
||||
model_config = ConfigDict(defer_build=DEFER_PYDANTIC_BUILD)
|
||||
|
|
@ -25,12 +28,13 @@ class LiteLLMBaseModel(BaseModel):
|
|||
) -> bool | None:
|
||||
# Resolve names from the model's own module, never a caller frame: a deferred first-use build
|
||||
# reads f_locals 5 frames up, and on Python < 3.13 that rewrites the dict the caller's locals() returned
|
||||
return super().model_rebuild(
|
||||
force=force,
|
||||
raise_errors=raise_errors,
|
||||
_parent_namespace_depth=0,
|
||||
_types_namespace=_types_namespace,
|
||||
)
|
||||
with _SCHEMA_BUILD_LOCK:
|
||||
return super().model_rebuild(
|
||||
force=force,
|
||||
raise_errors=raise_errors,
|
||||
_parent_namespace_depth=0,
|
||||
_types_namespace=_types_namespace,
|
||||
)
|
||||
|
||||
def model_post_init(self, context: object, /) -> None:
|
||||
# Instances built by a parent's validator or by model_construct skip this class's own
|
||||
|
|
|
|||
|
|
@ -1,10 +1,13 @@
|
|||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from itertools import chain
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from pydantic import ConfigDict
|
||||
from pydantic import ConfigDict, create_model
|
||||
|
||||
from litellm.types.llms.base import LiteLLMBaseModel
|
||||
|
||||
|
|
@ -90,3 +93,52 @@ def test_deferred_first_use_build_leaves_caller_locals_snapshot_untouched() -> N
|
|||
assert not DeferredProbe.__pydantic_complete__
|
||||
assert build(DeferredProbe) == ["model"]
|
||||
assert DeferredProbe.__pydantic_complete__
|
||||
|
||||
|
||||
_RACE_ROUNDS: Final = 100
|
||||
_RACE_THREADS: Final = 16
|
||||
_RACE_VALIDATIONS_PER_THREAD: Final = 5
|
||||
_RACE_FIELD_COUNT: Final = 20
|
||||
|
||||
|
||||
class _Deferred(LiteLLMBaseModel):
|
||||
model_config = ConfigDict(defer_build=True)
|
||||
|
||||
|
||||
def _fresh_deferred_subclass(round_id: int) -> tuple[type[LiteLLMBaseModel], type[LiteLLMBaseModel]]:
|
||||
fields: Final = {f"field_{index}": (str | int | None, None) for index in range(_RACE_FIELD_COUNT)}
|
||||
parent: Final = create_model(f"Parent{round_id}", __base__=_Deferred, **fields)
|
||||
child: Final = create_model(f"Child{round_id}", __base__=parent, extra_flag=(bool | None, False))
|
||||
return parent, child
|
||||
|
||||
|
||||
def _first_use_outcome(child: type[LiteLLMBaseModel]) -> str:
|
||||
try:
|
||||
return type(child.model_validate({"field_0": "x"})).__name__
|
||||
except AttributeError as error:
|
||||
return f"{type(error).__name__}: {error}"
|
||||
|
||||
|
||||
def _validate_after_barrier(child: type[LiteLLMBaseModel], gate: threading.Barrier) -> tuple[str, ...]:
|
||||
gate.wait()
|
||||
return tuple(_first_use_outcome(child) for _ in range(_RACE_VALIDATIONS_PER_THREAD))
|
||||
|
||||
|
||||
def _concurrent_first_use_outcomes(child: type[LiteLLMBaseModel]) -> frozenset[str]:
|
||||
gate: Final = threading.Barrier(_RACE_THREADS)
|
||||
with ThreadPoolExecutor(max_workers=_RACE_THREADS) as executor:
|
||||
per_thread: Final = tuple(executor.map(lambda _: _validate_after_barrier(child, gate), range(_RACE_THREADS)))
|
||||
return frozenset(chain.from_iterable(per_thread))
|
||||
|
||||
|
||||
def test_concurrent_first_use_of_a_deferred_subclass_always_builds_that_subclass() -> None:
|
||||
previous_switch_interval: Final = sys.getswitchinterval()
|
||||
sys.setswitchinterval(1e-6)
|
||||
try:
|
||||
for round_id in range(_RACE_ROUNDS):
|
||||
parent, child = _fresh_deferred_subclass(round_id)
|
||||
parent.model_validate({})
|
||||
assert not child.__pydantic_complete__
|
||||
assert _concurrent_first_use_outcomes(child) == {child.__name__}, f"round {round_id}"
|
||||
finally:
|
||||
sys.setswitchinterval(previous_switch_interval)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue