diff --git a/litellm/types/llms/base.py b/litellm/types/llms/base.py index e63e1b040d0..43b4adc099c 100644 --- a/litellm/types/llms/base.py +++ b/litellm/types/llms/base.py @@ -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 diff --git a/tests/unit/types/llms/test_types_llms_base.py b/tests/unit/types/llms/test_types_llms_base.py index 132fb95cad1..964689cdeeb 100644 --- a/tests/unit/types/llms/test_types_llms_base.py +++ b/tests/unit/types/llms/test_types_llms_base.py @@ -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)