fix(rust): repair retained callback CI gates

This commit is contained in:
Yujong Lee 2026-09-08 09:15:13 -07:00
parent 3be0ced670
commit 27bf5d9364
4 changed files with 99 additions and 24 deletions

View file

@ -20,9 +20,7 @@ TerminalAction = Literal[
"sync_failure",
"async_failure",
]
_OPTIONAL_ARGUMENTS_ADAPTER: Final[TypeAdapter[dict[str, object] | None]] = TypeAdapter(
dict[str, object] | None
)
_OPTIONAL_ARGUMENTS_ADAPTER: Final[TypeAdapter[dict[str, object] | None]] = TypeAdapter(dict[str, object] | None)
class NativeOutcome(IntEnum):

View file

@ -56,8 +56,8 @@ _UNSET: Final[_Unset] = _Unset()
@dataclass(slots=True)
class _RustMessagesState:
messages: RustMessages | None = None
amessages: RustAmessages | None = None
messages: RustMessages | None | _Unset = _UNSET
amessages: RustAmessages | None | _Unset = _UNSET
_STATE: Final = _RustMessagesState()
@ -87,13 +87,13 @@ def set_rust_messages(
def load_rust_messages() -> RustMessages | None:
if _STATE.messages is not None:
if not isinstance(_STATE.messages, _Unset):
return _STATE.messages
return _MESSAGES.load()
def load_rust_amessages() -> RustAmessages | None:
if _STATE.amessages is not None:
if not isinstance(_STATE.amessages, _Unset):
return _STATE.amessages
return _AMESSAGES.load()

View file

@ -257,7 +257,12 @@ def invoke_terminal(
return None
if action == "sync_success_if_needed":
if logging._should_run_sync_callbacks_for_async_calls(): # pyright: ignore[reportPrivateUsage] # preserves Logging's async callback policy
return invoke_terminal("sync_success", roots, logger, record, value, fallback_start_time, fallback_end_time)
def run() -> None:
_retained: Final = roots
logging.success_handler(value, start_time, end_time)
return utils.executor.submit(copy_context().run, run)
return None
exception: Final = cast(Exception, value) # cast-ok: Rust routes terminal failure values as Python exceptions
trace: Final = "".join(traceback.format_exception(type(exception), exception, exception.__traceback__))

View file

@ -229,6 +229,51 @@ def restore_ocr_context(logger: object) -> None:
pass
def drive_ocr_sync(arguments: dict[str, object], bindings: object) -> object:
logger: Final = initialize_ocr_logging(arguments, False)
try:
state: Final = bindings.prepare(arguments, logger, False)
except RuntimeError as error:
raise NotImplementedError(str(error)) from error
bindings.pre_call(state)
return bindings.send_sync(state)
async def drive_ocr_async(arguments: dict[str, object], bindings: object) -> object:
logger: Final = initialize_ocr_logging(arguments, True)
try:
state: Final = bindings.prepare(arguments, logger, True)
except RuntimeError as error:
raise NotImplementedError(str(error)) from error
bindings.pre_call(state)
return bindings.finish(await bindings.send(state))
class MessagesLogging:
def pre_call(self, **_kwargs: object) -> None:
pass
def drive_messages_sync(arguments: dict[str, object], bindings: object) -> object:
state: Final = bindings.prepare(arguments, MessagesLogging())
return bindings.send_sync(state)
async def drive_messages_async(arguments: dict[str, object], bindings: object) -> object:
state: Final = bindings.prepare(arguments, MessagesLogging())
return await bindings.send(state)
def drive_chat_sync(arguments: dict[str, object], bindings: object) -> object:
state: Final = bindings.prepare(arguments, MessagesLogging())
return bindings.send_sync(state)
async def drive_chat_async(arguments: dict[str, object], bindings: object) -> object:
state: Final = bindings.prepare(arguments, MessagesLogging())
return await bindings.send(state)
class WheelCallTypes:
aocr = "aocr"
@ -265,22 +310,28 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
},
}
if route == "messages":
return common | {
"model": "claude-sonnet-4-5",
"body": {
return {
"arguments": common
| {
"model": "claude-sonnet-4-5",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hello-from-messages"}],
},
"api_key": "sk-native",
"custom_llm_provider": "anthropic",
"body": {
"model": "claude-sonnet-4-5",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hello-from-messages"}],
},
"api_key": "sk-native",
"custom_llm_provider": "anthropic",
}
}
if route == "chat_completions":
return common | {
"model": "anthropic/claude-sonnet-4-5",
"messages": [{"role": "user", "content": "hello-from-chat"}],
"optional_params": {"max_tokens": 17},
"api_key": "sk-native",
return {
"arguments": common
| {
"model": "anthropic/claude-sonnet-4-5",
"messages": [{"role": "user", "content": "hello-from-chat"}],
"optional_params": {"max_tokens": 17},
"api_key": "sk-native",
}
}
raise AssertionError(f"unknown route: {route}")
@ -318,7 +369,7 @@ def assert_rate_limit(native: object, route: str, error: BaseException) -> None:
assert error.llm_provider == "mistral"
assert str(error) == "OCR provider request failed (HTTP 429)"
return
if route == "chat_completions":
if route in {"messages", "chat_completions"}:
upstream_error: Final = native.RustUpstreamError
if not isinstance(error, upstream_error) or error.args[0] != 429:
raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
@ -367,7 +418,8 @@ async def exercise_unsupported_ocr(native: object, api_base: str) -> None:
else:
native.ocr(arguments)
except NotImplementedError as error:
assert operation in str(error) # noqa: PT017 # the selected native entry point is parametrized at runtime
if operation not in str(error):
raise AssertionError(f"native OCR returned the wrong decline: {error}") from error
else:
raise AssertionError(f"native OCR accepted unsupported operation: {operation}")
assert arguments["litellm_logging_obj"].calls == ()
@ -396,6 +448,14 @@ def exercise_routes(native_path: Path, api_base: str) -> object:
ocr_bridge: Final = ModuleType("litellm.rust_bridge.ocr")
ocr_bridge.initialize_logging = initialize_ocr_logging
ocr_bridge.invoke_terminal = invoke_ocr_terminal
ocr_bridge._drive_sync = drive_ocr_sync
ocr_bridge._drive_async = drive_ocr_async
messages_bridge: Final = ModuleType("litellm.rust_bridge.messages")
messages_bridge._drive_sync = drive_messages_sync
messages_bridge._drive_async = drive_messages_async
chat_bridge: Final = ModuleType("litellm.rust_bridge.chat_completions")
chat_bridge._drive_sync = drive_chat_sync
chat_bridge._drive_async = drive_chat_async
utils: Final = ModuleType("litellm.utils")
utils.is_internal_call = ContextVar("wheel_internal_call", default=False)
utils.async_pre_call_deployment_hook = pre_ocr_deployment
@ -418,7 +478,19 @@ def exercise_routes(native_path: Path, api_base: str) -> object:
with patch.dict(
sys.modules,
packages
| {module.__name__: module for module in (transformation, exceptions, httpx, ocr_bridge, utils, types_utils)},
| {
module.__name__: module
for module in (
transformation,
exceptions,
httpx,
ocr_bridge,
messages_bridge,
chat_bridge,
utils,
types_utils,
)
},
):
exercise_sync(native, api_base)
asyncio.run(exercise_async(native, api_base))