fix(rust): complete callback trace integration

This commit is contained in:
Yujong Lee 2026-09-05 16:47:30 -07:00
parent 304243f865
commit 01fbd47f37
7 changed files with 200 additions and 11 deletions

View file

@ -603,6 +603,12 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
transform_document_intelligence_response(model, response_json, false)
}
#[tracing::instrument(
name = "transform_ocr_response",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
fn transform_ocr_response_with_params(
&self,
model: &str,

View file

@ -71,6 +71,19 @@ impl PythonProviderObserver {
}
}
pub(crate) fn python_async_session(
adapter: Py<PyAny>,
py: Python<'_>,
) -> PyResult<PythonSession<AsyncContext>> {
let module = py.import("litellm.rust_bridge._native")?;
let runtime = module
.getattr("__python_callback_runtime__")?
.extract::<PyRef<'_, PythonCallbackRuntime>>()?
.0
.clone();
PythonSession::new(adapter.bind(py), runtime.async_context(py)?)
}
impl ProviderAttemptObserver for PythonProviderObserver {
type Error = PyErr;

View file

@ -1,12 +1,21 @@
mod callback_bindings;
#[cfg(test)]
#[path = "../tests/callbacks/mod.rs"]
mod callback_tests;
mod constants;
mod diagnostics;
mod errors;
mod execution;
#[cfg(feature = "trace-parity")]
mod function_trace;
mod marshal;
mod python_hook_bindings;
mod routes;
use std::sync::atomic::{AtomicU64, Ordering};
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use litellm_core::provider_callbacks::{CallbackDecision, SessionEvent, SessionObserver};
use litellm_core::responses::types::ResponsesWebSocketRequest;
use pyo3::prelude::*;
use pyo3::types::PyAny;
@ -14,6 +23,8 @@ use pyo3::types::PyAny;
use crate::errors::core_error_to_pyerr;
use crate::marshal::{NativeRequestContext, NativeRequestOptions};
static NEXT_WEBSOCKET_SESSION_ID: AtomicU64 = AtomicU64::new(1);
#[derive(FromPyObject)]
struct WebSocketConnectRequest {
url: String,
@ -36,7 +47,6 @@ impl ResponsesWebSocketConnection {
context: NativeRequestContext,
callback_adapter: Option<Py<PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let _ = callback_adapter;
let provider_supported = litellm_core::responses::websocket::native_websocket_supported(
options.provider("openai"),
);
@ -45,11 +55,52 @@ impl ResponsesWebSocketConnection {
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
}
let options: litellm_core::request_options::RequestOptions = options.into();
let call_id = context.litellm_call_id.clone().unwrap_or_default();
let session_id = format!(
"responses-websocket-{}",
NEXT_WEBSOCKET_SESSION_ID.fetch_add(1, Ordering::Relaxed)
);
let mut observer = callback_adapter
.map(|adapter| crate::callback_bindings::python_async_session(adapter, py))
.transpose()?;
let request = ResponsesWebSocketRequest { url: request.url };
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let inner = RustResponsesWebSocketConnection::connect(request, &options, &context)
.await
.map_err(core_error_to_pyerr)?;
if let Some(observer) = observer.as_mut() {
let decision = observer
.before_connect(&session_event(&session_id, &call_id, None))
.await?;
match decision {
CallbackDecision::Unchanged => {}
CallbackDecision::Replace { .. } => {
return Err(pyo3::exceptions::PyValueError::new_err(
"before_connect cannot replace WebSocket setup",
));
}
CallbackDecision::Reject { message, .. } => {
return Err(pyo3::exceptions::PyValueError::new_err(message));
}
}
}
let inner = match RustResponsesWebSocketConnection::connect(request, &options, &context).await {
Ok(inner) => inner,
Err(error) => {
if let Some(observer) = observer.as_mut() {
observer
.error(&session_event(
&session_id,
&call_id,
Some(error.to_string()),
))
.await?;
}
return Err(core_error_to_pyerr(error));
}
};
if let Some(observer) = observer.as_mut() {
observer
.connected(&session_event(&session_id, &call_id, None))
.await?;
}
Ok(ResponsesWebSocketConnection { inner })
})
}
@ -76,6 +127,18 @@ impl ResponsesWebSocketConnection {
}
}
fn session_event(session_id: &str, call_id: &str, message: Option<String>) -> SessionEvent {
SessionEvent {
session_id: session_id.to_string(),
call_id: call_id.to_string(),
trace_id: None,
event: None,
response_id: None,
sequence: None,
message,
}
}
#[pyfunction]
#[pyo3(signature = (_model, custom_llm_provider, *, context))]
fn responses_websocket_decline(
@ -231,6 +294,16 @@ mod tests {
let code = CString::new(
r#"
import asyncio
import sys
import types
litellm_module = types.ModuleType('litellm')
rust_bridge_module = types.ModuleType('litellm.rust_bridge')
litellm_module.rust_bridge = rust_bridge_module
rust_bridge_module._native = native
sys.modules['litellm'] = litellm_module
sys.modules['litellm.rust_bridge'] = rust_bridge_module
sys.modules['litellm.rust_bridge._native'] = native
async def exercise():
for request, request_options, request_context, field in (
@ -250,7 +323,56 @@ async def exercise():
else:
raise AssertionError('invalid WebSocket input reached execution')
connection = await native.ResponsesWebSocketConnection.connect(Request(url=url), options=options, context=context)
events = []
class Adapter:
async def before_connect(self, event):
events.append(('before_connect', event))
return {'action': 'unchanged'}
async def connected(self, event):
events.append(('connected', event))
async def before_send(self, event):
return {'action': 'unchanged'}
async def after_receive(self, event):
return {'action': 'unchanged'}
async def response_complete(self, event):
pass
async def response_error(self, event):
pass
async def error(self, event):
events.append(('error', event))
async def close(self, event):
pass
class RejectingAdapter(Adapter):
async def before_connect(self, event):
return {'action': 'reject', 'message': 'blocked', 'status_code': 400}
try:
await native.ResponsesWebSocketConnection.connect(
Request(url=url), options=options, context=context, callback_adapter=RejectingAdapter()
)
except ValueError as error:
assert str(error) == 'blocked'
else:
raise AssertionError('rejected WebSocket setup reached execution')
connection = await native.ResponsesWebSocketConnection.connect(
Request(url=url),
options=options,
context=replace(context, litellm_call_id='call-1'),
callback_adapter=Adapter(),
)
assert [name for name, _ in events] == ['before_connect', 'connected']
assert events[0][1]['call_id'] == 'call-1'
assert events[0][1]['session_id'] == events[1][1]['session_id']
assert type(connection) is native.ResponsesWebSocketConnection
await connection.send_text("from-python")
assert await connection.recv_text() == "from-server"
@ -258,6 +380,9 @@ async def exercise():
assert await connection.recv_text() is None
asyncio.run(asyncio.wait_for(exercise(), timeout=5))
sys.modules.pop('litellm.rust_bridge._native')
sys.modules.pop('litellm.rust_bridge')
sys.modules.pop('litellm')
"#,
)
.expect("Python source should not contain null bytes");

View file

@ -95,6 +95,7 @@ async def connect(
timeout: float | httpx.Timeout | None,
model: str = "responses websocket",
provider: str = "openai",
callback_adapter: SessionCallbackHandle | None = None,
fallback: Callable[[], Awaitable[Connection | None]] = async_none,
context: NativeRequestContext | None = None,
) -> Connection | None:
@ -116,6 +117,7 @@ async def connect(
requires_connection=True,
),
),
callback_adapter=callback_adapter,
),
call=lambda connection_type, request: call_native(connection_type.connect, request),
preflight=lambda: assess_route(_PREFLIGHT, model, provider),
@ -133,6 +135,7 @@ async def open_connection(
timeout: float | httpx.Timeout | None,
model: str,
provider: str,
callback_adapter: SessionCallbackHandle | None = None,
fallback: Callable[[], AbstractAsyncContextManager[Connection]],
context: NativeRequestContext | None = None,
) -> AsyncGenerator[Connection]:
@ -147,6 +150,7 @@ async def open_connection(
timeout=timeout,
model=model,
provider=provider,
callback_adapter=callback_adapter,
fallback=python_connection,
context=context,
)

View file

@ -61,8 +61,9 @@ def _collect(
from litellm.rust_bridge import get_native_bridge
from litellm.rust_bridge.transcription import (
RustAtranscription,
RustRouteDecline,
RustTranscription,
set_rust_transcription,
configure_rust_transcription,
)
native: Final = get_native_bridge()
@ -70,11 +71,17 @@ def _collect(
raise RuntimeError("native transcription is required for diagnostic trace parity")
# This required-native SDK route is injected only for diagnostic comparison.
# Normal callback readiness remains empty before and after this scope.
set_rust_transcription(
sync=cast(RustTranscription, native.transcription),
asynchronous=cast(RustAtranscription, native.atranscription),
configure_rust_transcription(
transcription=cast(RustTranscription, native.transcription),
atranscription=cast(RustAtranscription, native.atranscription),
decline=cast(RustRouteDecline, native.transcription_decline),
)
stack.callback(
configure_rust_transcription,
transcription=None,
atranscription=None,
decline=None,
)
stack.callback(set_rust_transcription, sync=None, asynchronous=None)
with profile_python(Path(litellm.__file__).parent, threads=True) as profiler:
_invoke(function, kwargs, asynchronous=asynchronous)
return tuple(profiler.events)

View file

@ -23,7 +23,7 @@ MAPPINGS: Final = (
mapping(rust_span="transform_transcription_request"),
mapping(
rust_span="execute_audio_transcription_provider_call",
python_frame=r"BedrockAudioTranscriptionRustDispatch\.(?:async_)?audio_transcriptions$",
python_frame=r"rust_bridge/request\.py:\d+ call_native$",
),
mapping(rust_span="transform_transcription_response"),
mapping(rust_span="http_request"),

View file

@ -1,8 +1,11 @@
from __future__ import annotations
from typing import cast
import pytest
from litellm.rust_bridge import configuration, responses_websocket
from litellm.rust_bridge.callbacks import SessionCallbackHandle
from litellm.rust_bridge.request import NativeRequestContext, NativeResponsesWebSocketRequest
@ -98,6 +101,37 @@ async def test_enabled_bridge_connects_and_adapts_socket(
await connection.close()
@pytest.mark.asyncio
async def test_connection_forwards_session_callback_adapter() -> None:
configuration.rust(True)
received: list[object] = []
callback_adapter = cast(SessionCallbackHandle, object())
class Native:
@classmethod
async def connect(
cls,
request: NativeResponsesWebSocketRequest,
*,
context: NativeRequestContext,
callback_adapter: object | None = None,
) -> _FakeNativeConnection:
received.append(callback_adapter)
return _FakeNativeConnection()
responses_websocket.set_rust_responses_websocket(connection=Native)
connection = await responses_websocket.connect(
url="wss://example.test/responses",
headers={},
timeout=None,
callback_adapter=callback_adapter,
)
assert connection is not None
assert received == [callback_adapter]
class _FailingNativeBridge:
@classmethod
async def connect(