diff --git a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs index d15c032f0bc..3cecbbbab19 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs @@ -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, diff --git a/litellm-rust/crates/python-bridge/src/callback_bindings.rs b/litellm-rust/crates/python-bridge/src/callback_bindings.rs index a8f05e26226..f7180d84823 100644 --- a/litellm-rust/crates/python-bridge/src/callback_bindings.rs +++ b/litellm-rust/crates/python-bridge/src/callback_bindings.rs @@ -71,6 +71,19 @@ impl PythonProviderObserver { } } +pub(crate) fn python_async_session( + adapter: Py, + py: Python<'_>, +) -> PyResult> { + let module = py.import("litellm.rust_bridge._native")?; + let runtime = module + .getattr("__python_callback_runtime__")? + .extract::>()? + .0 + .clone(); + PythonSession::new(adapter.bind(py), runtime.async_context(py)?) +} + impl ProviderAttemptObserver for PythonProviderObserver { type Error = PyErr; diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 35bf02ae227..6d31a7b937a 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -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>, ) -> PyResult> { - 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) -> 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"); diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index c47c48264dc..80cef256428 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -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, ) diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py b/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py index 4a831972154..d9a821ccb77 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py @@ -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) diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/transcription/case.py b/tests/rust-python-harness/strategies/trace_parity/sdk/transcription/case.py index a80eb8f07b5..7ac7b6d9590 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/transcription/case.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/transcription/case.py @@ -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"), diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index e72500dbc2f..35357e949d1 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -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(