mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
test(rust): isolate callback registries per hypothesis example
Replace the module-level LATEST_EDITS list with per-example callback registry isolation, and import litellm names with from-imports in the legacy callback shim so the module uses one import style. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
parent
41873e2bc7
commit
ba6b22cf56
4 changed files with 88 additions and 71 deletions
|
|
@ -65,13 +65,17 @@ def setup(
|
|||
|
||||
|
||||
def check_limits(kwargs: Mapping[str, object]) -> None:
|
||||
import litellm
|
||||
from litellm import (
|
||||
BudgetExceededError,
|
||||
_current_cost, # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor
|
||||
max_budget,
|
||||
num_retries_per_request,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit
|
||||
|
||||
current_cost: Final = litellm._current_cost # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor
|
||||
if litellm.max_budget and current_cost > litellm.max_budget:
|
||||
raise litellm.BudgetExceededError(current_cost=current_cost, max_budget=litellm.max_budget)
|
||||
if max_retries_per_request_hit(kwargs, litellm.num_retries_per_request):
|
||||
if max_budget and _current_cost > max_budget:
|
||||
raise BudgetExceededError(current_cost=_current_cost, max_budget=max_budget)
|
||||
if max_retries_per_request_hit(kwargs, num_retries_per_request):
|
||||
raise RuntimeError("Max retries per request hit!")
|
||||
|
||||
|
||||
|
|
@ -281,9 +285,9 @@ def is_internal_call() -> bool:
|
|||
|
||||
|
||||
def credential_list() -> list[CredentialItem]:
|
||||
import litellm
|
||||
from litellm import credential_list as credentials
|
||||
|
||||
return litellm.credential_list
|
||||
return credentials
|
||||
|
||||
|
||||
def warn_unknown_credential(name: str, loaded: int) -> None:
|
||||
|
|
|
|||
|
|
@ -1,10 +1,9 @@
|
|||
import asyncio
|
||||
import os
|
||||
from collections.abc import AsyncIterator, Generator, Iterator
|
||||
from collections.abc import AsyncIterator, Generator
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from types import ModuleType
|
||||
from typing import Final, cast
|
||||
from contextlib import ExitStack
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
|
@ -18,62 +17,20 @@ from litellm.rust_bridge.configuration import ( # pyright: ignore[reportPrivate
|
|||
_parse_env_bool,
|
||||
)
|
||||
from tests.test_litellm_rust.support.callback_recorder import drain_logging
|
||||
from tests.test_litellm_rust.support.isolation import isolated_callback_registries, rebound
|
||||
from tests.test_litellm_rust.support.recording_server import RecordingServer, recording_service
|
||||
|
||||
CALLBACK_ATTRIBUTES: Final = (
|
||||
"callbacks",
|
||||
"input_callback",
|
||||
"success_callback",
|
||||
"failure_callback",
|
||||
"_async_input_callback",
|
||||
"_async_success_callback",
|
||||
"_async_failure_callback",
|
||||
)
|
||||
|
||||
|
||||
def _list_attribute(container: ModuleType, attribute: str) -> list[object]:
|
||||
value: Final = getattr(container, attribute)
|
||||
if not isinstance(value, list):
|
||||
raise AssertionError(f"{container.__name__}.{attribute} is not a list")
|
||||
return cast(list[object], value)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _isolated_list(container: ModuleType, attribute: str) -> Iterator[None]:
|
||||
source: Final = _list_attribute(container, attribute)
|
||||
original: Final = list(source)
|
||||
source.clear() # mutable-ok: test isolation mutates global registries by design
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
source.clear()
|
||||
source.extend(original)
|
||||
setattr(container, attribute, source)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _rebound(container: object, attribute: str, value: object) -> Iterator[None]:
|
||||
original: Final[object] = getattr(container, attribute)
|
||||
setattr(container, attribute, value)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
setattr(container, attribute, original)
|
||||
|
||||
|
||||
@pytest_asyncio.fixture(autouse=True, loop_scope="function")
|
||||
async def isolate_ocr_test_state() -> AsyncIterator[None]:
|
||||
with ExitStack() as stack:
|
||||
for attribute in CALLBACK_ATTRIBUTES:
|
||||
stack.enter_context(_isolated_list(litellm, attribute))
|
||||
stack.enter_context(_isolated_list(litellm_logging, "_in_memory_loggers")) # pyright: ignore[reportPrivateUsage] # no public callback-cache accessor
|
||||
stack.enter_context(_rebound(utils, "callback_list", [])) # rebind-ok: isolate legacy callback registry
|
||||
stack.enter_context(_rebound(litellm, "cache", None)) # test-quality-ok: isolate process-global cache
|
||||
stack.enter_context(_rebound(_CONFIGURATION, "override", None))
|
||||
stack.enter_context(isolated_callback_registries())
|
||||
stack.enter_context(rebound(litellm, "cache", None)) # test-quality-ok: isolate process-global cache
|
||||
stack.enter_context(rebound(_CONFIGURATION, "override", None))
|
||||
executor: Final = ThreadPoolExecutor(thread_name_prefix="rust-ocr-test-logging")
|
||||
stack.enter_context(_rebound(litellm_logging, "executor", executor))
|
||||
stack.enter_context(_rebound(utils, "executor", executor))
|
||||
stack.enter_context(_rebound(thread_pool_executor, "executor", executor))
|
||||
stack.enter_context(rebound(litellm_logging, "executor", executor))
|
||||
stack.enter_context(rebound(utils, "executor", executor))
|
||||
stack.enter_context(rebound(thread_pool_executor, "executor", executor))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ import copy
|
|||
import gc
|
||||
import queue
|
||||
import threading
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
|
@ -13,6 +15,7 @@ import litellm
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging
|
||||
from tests.test_litellm_rust.support.isolation import isolated_callback_registries
|
||||
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
|
||||
from tests.test_litellm_rust.support.requests import (
|
||||
OCR_DOCUMENT,
|
||||
|
|
@ -307,18 +310,13 @@ JSON_VALUES: Final = st.recursive(
|
|||
)
|
||||
|
||||
|
||||
LATEST_EDITS: Final[list[dict[str, object]]] = []
|
||||
|
||||
|
||||
class ApplyLatestEdits(CustomLogger):
|
||||
"""Registrations can outlive one hypothesis example, so every instance applies the current example's edits."""
|
||||
|
||||
def __init__(self, latest: list[dict[str, object]]) -> None:
|
||||
class ApplyEdits(CustomLogger):
|
||||
def __init__(self, edits: Mapping[str, object]) -> None:
|
||||
super().__init__()
|
||||
self.latest = latest
|
||||
self.edits: Final = edits
|
||||
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
request_body(kwargs).update(copy.deepcopy(self.latest[-1]))
|
||||
request_body(kwargs).update(copy.deepcopy(dict(self.edits)))
|
||||
|
||||
|
||||
@settings(max_examples=25, deadline=None, suppress_health_check=[HealthCheck.function_scoped_fixture])
|
||||
|
|
@ -327,9 +325,9 @@ def test_native_ocr_provider_receives_the_body_exactly_as_pre_call_callbacks_lef
|
|||
ocr_server: RecordingServer, edits: dict[str, object]
|
||||
) -> None:
|
||||
ocr_server.expected_requests = None
|
||||
LATEST_EDITS.append(edits)
|
||||
|
||||
call_native_ocr_with_callbacks(ocr_server, [ApplyLatestEdits(LATEST_EDITS)])
|
||||
with isolated_callback_registries():
|
||||
call_native_ocr_with_callbacks(ocr_server, [ApplyEdits(MappingProxyType(edits))])
|
||||
|
||||
assert ocr_server.requests[-1].body == {"model": "mistral-ocr-latest", "document": OCR_DOCUMENT, **edits}
|
||||
|
||||
|
|
|
|||
58
tests/test_litellm_rust/support/isolation.py
Normal file
58
tests/test_litellm_rust/support/isolation.py
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
from collections.abc import Generator
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from types import ModuleType
|
||||
from typing import Final, cast
|
||||
|
||||
import litellm
|
||||
from litellm import utils
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
|
||||
CALLBACK_ATTRIBUTES: Final = (
|
||||
"callbacks",
|
||||
"input_callback",
|
||||
"success_callback",
|
||||
"failure_callback",
|
||||
"_async_input_callback",
|
||||
"_async_success_callback",
|
||||
"_async_failure_callback",
|
||||
)
|
||||
|
||||
|
||||
def _list_attribute(container: ModuleType, attribute: str) -> list[object]:
|
||||
value: Final = getattr(container, attribute)
|
||||
if not isinstance(value, list):
|
||||
raise AssertionError(f"{container.__name__}.{attribute} is not a list")
|
||||
return cast(list[object], value)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _isolated_list(container: ModuleType, attribute: str) -> Generator[None]:
|
||||
source: Final = _list_attribute(container, attribute)
|
||||
original: Final = list(source)
|
||||
source.clear() # mutable-ok: test isolation mutates global registries by design
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
source.clear()
|
||||
source.extend(original)
|
||||
setattr(container, attribute, source)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def rebound(container: object, attribute: str, value: object) -> Generator[None]:
|
||||
original: Final[object] = getattr(container, attribute)
|
||||
setattr(container, attribute, value)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
setattr(container, attribute, original)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def isolated_callback_registries() -> Generator[None]:
|
||||
with ExitStack() as stack:
|
||||
for attribute in CALLBACK_ATTRIBUTES:
|
||||
stack.enter_context(_isolated_list(litellm, attribute))
|
||||
stack.enter_context(_isolated_list(litellm_logging, "_in_memory_loggers")) # pyright: ignore[reportPrivateUsage] # no public callback-cache accessor
|
||||
stack.enter_context(rebound(utils, "callback_list", [])) # rebind-ok: isolate legacy callback registry
|
||||
yield
|
||||
Loading…
Add table
Reference in a new issue