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:
Yujong Lee 2026-09-18 15:57:50 -07:00
parent 41873e2bc7
commit ba6b22cf56
4 changed files with 88 additions and 71 deletions

View file

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

View file

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

View file

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

View 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