mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
fix(rust): complete callback trace integration
This commit is contained in:
parent
304243f865
commit
01fbd47f37
7 changed files with 200 additions and 11 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue