refactor(rust): encapsulate trace parity in python bridge

This commit is contained in:
Yujong Lee 2026-09-08 14:57:57 -07:00
parent b72c704c14
commit 81bb960d8c
17 changed files with 98 additions and 141 deletions

View file

@ -1466,7 +1466,6 @@ dependencies = [
"tokio",
"tokio-tungstenite 0.24.0",
"tracing",
"tracing-subscriber",
]
[[package]]
@ -1484,6 +1483,7 @@ dependencies = [
name = "litellm-python-bridge"
version = "0.1.0"
dependencies = [
"axum",
"criterion",
"litellm-ai-gateway",
"litellm-core",
@ -1493,7 +1493,9 @@ dependencies = [
"serde",
"serde_json",
"tokio",
"tower",
"tracing",
"tracing-subscriber",
]
[[package]]

View file

@ -24,7 +24,6 @@ futures-util.workspace = true
serde_json.workspace = true
axum = { workspace = true, features = ["ws"], optional = true }
serde.workspace = true
tower = { version = "0.5.3", features = ["util"], optional = true }
[features]
default = []
@ -32,7 +31,6 @@ server = ["dep:axum", "dep:litellm-gateway-auth"]
# Build the gateway's config from the proxy YAML via an embedded Python
# interpreter (links libpython; requires `litellm` importable at runtime).
python-config = ["litellm-config/python"]
trace-parity = ["server", "dep:tower", "litellm-core/observability"]
[dev-dependencies]
futures-channel = "0.3"

View file

@ -18,7 +18,8 @@ deployment, and adapts frames while core dials OpenAI and splices the session.
Dependency direction is acyclic: config depends on core, the gateway depends on
config and core, and the Python bridge depends on core and Python interop. Its
optional `trace-parity` diagnostics also depend on the gateway.
private `trace_parity` module owns trace collection and gateway fixtures, enabled
by the bridge’s `trace-parity` feature and the gateway’s `server` feature
- **Client endpoint:** `wss://<host>/v1/realtime?model=<model>` (WebSocket)
- **Auth:** `Authorization: Bearer $LITELLM_MASTER_KEY` (fails closed if unset)

View file

@ -14,7 +14,5 @@ pub mod io;
pub mod routes;
#[cfg(feature = "server")]
pub mod state;
#[cfg(feature = "trace-parity")]
pub mod trace_parity;
mod constants;

View file

@ -19,7 +19,6 @@ thiserror.workspace = true
tracing.workspace = true
tokio = { workspace = true, features = ["rt", "sync", "time"] }
tokio-tungstenite.workspace = true
tracing-subscriber = { workspace = true, optional = true }
litellm-auth-aws = { workspace = true, optional = true }
[features]
@ -27,10 +26,8 @@ default = []
bedrock-auth = [
"dep:litellm-auth-aws",
]
observability = ["dep:tracing-subscriber"]
[dev-dependencies]
futures-channel = "0.3"
rstest.workspace = true
tokio = { workspace = true, features = ["io-util", "macros", "net", "rt-multi-thread"] }
tracing-subscriber.workspace = true

View file

@ -7,8 +7,6 @@ pub mod http_utils;
pub mod integrations;
pub mod lifecycle;
pub mod messages;
#[cfg(any(feature = "observability", test))]
pub mod observability;
pub mod ocr;
pub mod providers;
pub mod realtime;

View file

@ -1,59 +0,0 @@
use tracing::span::Id;
use tracing::{Level, Metadata, Subscriber};
use tracing_subscriber::filter::{FilterFn, LevelFilter, filter_fn};
use tracing_subscriber::layer::Context;
use tracing_subscriber::registry::LookupSpan;
use crate::constants::FUNCTION_TRACE_TARGET;
pub mod function_trace;
pub use function_trace::{FunctionTrace, FunctionTraceEvent};
pub fn function_trace_filter() -> FilterFn<impl Fn(&Metadata<'_>) -> bool> {
filter_fn(|metadata| {
metadata.is_span()
&& metadata.target() == FUNCTION_TRACE_TARGET
&& *metadata.level() == Level::TRACE
})
.with_max_level_hint(LevelFilter::TRACE)
}
pub fn span_depth<S>(context: &Context<'_, S>, id: &Id) -> usize
where
S: Subscriber + for<'lookup> LookupSpan<'lookup>,
{
context
.span(id)
.map(|span| span.scope().skip(1).count())
.unwrap_or_default()
}
#[cfg(test)]
mod tests {
use tracing::instrument::WithSubscriber;
use super::*;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn instrumented_with_literal_target() {}
#[tokio::test]
async fn literal_instrument_target_matches_filter_constant() {
assert_eq!(FUNCTION_TRACE_TARGET, "litellm::function_trace");
let trace = FunctionTrace::default();
instrumented_with_literal_target()
.with_subscriber(trace.dispatcher())
.await;
let events = trace.events();
assert_eq!(events.len(), 1);
assert_eq!(events[0].id, 0);
assert_eq!(events[0].parent_id, None);
assert_eq!(events[0].function, "instrumented_with_literal_target");
assert_eq!(events[0].module_path, Some(module_path!()));
assert_eq!(events[0].file, Some(file!()));
assert!(events[0].line.is_some());
}
}

View file

@ -17,11 +17,16 @@ panic-test = []
trace-parity = [
"dep:tracing",
"dep:litellm-ai-gateway",
"litellm-core/observability",
"litellm-ai-gateway/trace-parity",
"dep:tracing-subscriber",
"dep:axum",
"dep:tower",
"litellm-ai-gateway/server",
]
[dependencies]
tracing-subscriber = { workspace = true, optional = true }
axum = { workspace = true, optional = true }
tower = { version = "0.5.3", features = ["util"], optional = true }
tracing = { workspace = true, optional = true }
litellm-core = { workspace = true, features = ["bedrock-auth"] }
litellm-ai-gateway = { workspace = true, default-features = false, optional = true }

View file

@ -3,11 +3,11 @@
mod diagnostics;
mod driver;
mod errors;
#[cfg(feature = "trace-parity")]
mod function_trace;
mod marshal;
mod retained;
mod routes;
#[cfg(feature = "trace-parity")]
mod trace_parity;
use pyo3::prelude::*;
@ -88,6 +88,7 @@ mod tests {
"amessages",
"chat_completions",
"achat_completions",
"chat_completions_decline",
"gateway_messages",
]
);

View file

@ -150,7 +150,7 @@ mod trace {
})?;
litellm_python_interop::run_sync(
py,
crate::function_trace::capture(future),
crate::trace_parity::capture(future),
core_error_to_pyerr,
)
}
@ -181,7 +181,7 @@ mod trace {
})?;
litellm_python_interop::run_async(
py,
crate::function_trace::capture(future),
crate::trace_parity::capture(future),
core_error_to_pyerr,
)
}

View file

@ -79,7 +79,7 @@ mod tests {
let future = prepare_echo(EchoInputs { value })?;
litellm_python_interop::run_sync(
py,
crate::function_trace::capture(future),
crate::trace_parity::capture(future),
map_error,
)
}
@ -89,7 +89,7 @@ mod tests {
let future = prepare_echo(EchoInputs { value })?;
litellm_python_interop::run_async(
py,
crate::function_trace::capture(future),
crate::trace_parity::capture(future),
map_error,
)
}

View file

@ -1,30 +0,0 @@
use pyo3::prelude::*;
use serde_json::Value;
use crate::errors::core_error_to_pyerr;
use litellm_python_interop::run_async;
#[pyfunction]
fn gateway_messages<'py>(
py: Python<'py>,
model_alias: String,
provider_model: String,
api_base: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)] body: Value,
) -> PyResult<Bound<'py, PyAny>> {
let future = litellm_ai_gateway::trace_parity::messages_request(
model_alias,
provider_model,
api_base,
body,
);
run_async(
py,
crate::function_trace::capture(future),
core_error_to_pyerr,
)
}
pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> {
super::definition::add_function(module, wrap_pyfunction!(gateway_messages, module)?)
}

View file

@ -2,9 +2,6 @@ use pyo3::prelude::*;
mod definition;
#[cfg(feature = "trace-parity")]
mod gateway_messages;
mod audio_transcription;
mod chat_completions;
mod messages;
@ -24,7 +21,7 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
audio_transcription::register_trace(&trace)?;
messages::register_trace(&trace)?;
chat_completions::register_trace(&trace)?;
gateway_messages::register_trace(&trace)?;
crate::trace_parity::register_gateway(&trace)?;
module.add_submodule(&trace)?;
}
Ok(())

View file

@ -3,32 +3,42 @@ use std::sync::{Arc, Mutex};
use serde::Serialize;
use tracing::span::{Attributes, Id};
use tracing::{Dispatch, Subscriber};
use tracing::{Dispatch, Level, Metadata, Subscriber};
use tracing_subscriber::filter::{FilterFn, LevelFilter, filter_fn};
use tracing_subscriber::layer::Context;
use tracing_subscriber::prelude::*;
use tracing_subscriber::registry::LookupSpan;
use tracing_subscriber::{Layer, Registry};
use super::function_trace_filter;
use litellm_core::constants::FUNCTION_TRACE_TARGET;
fn function_trace_filter() -> FilterFn<impl Fn(&Metadata<'_>) -> bool> {
filter_fn(|metadata| {
metadata.is_span()
&& metadata.target() == FUNCTION_TRACE_TARGET
&& *metadata.level() == Level::TRACE
})
.with_max_level_hint(LevelFilter::TRACE)
}
#[derive(Clone, Debug, PartialEq, Serialize)]
pub struct FunctionTraceEvent {
pub id: usize,
pub parent_id: Option<usize>,
pub function: &'static str,
pub module_path: Option<&'static str>,
pub file: Option<&'static str>,
pub line: Option<u32>,
pub(super) struct FunctionTraceEvent {
id: usize,
parent_id: Option<usize>,
function: &'static str,
module_path: Option<&'static str>,
file: Option<&'static str>,
line: Option<u32>,
}
#[derive(Clone, Default)]
pub struct FunctionTrace {
pub(super) struct FunctionTrace {
events: Arc<Mutex<Vec<FunctionTraceEvent>>>,
span_events: Arc<Mutex<HashMap<Id, usize>>>,
}
impl FunctionTrace {
pub fn dispatcher(&self) -> Dispatch {
pub(super) fn dispatcher(&self) -> Dispatch {
Dispatch::new(
Registry::default().with(
FunctionTraceLayer {
@ -39,7 +49,7 @@ impl FunctionTrace {
)
}
pub fn events(&self) -> Vec<FunctionTraceEvent> {
pub(super) fn events(&self) -> Vec<FunctionTraceEvent> {
self.events
.lock()
.unwrap_or_else(|error| error.into_inner())
@ -90,9 +100,8 @@ where
#[cfg(test)]
mod tests {
use crate::constants::FUNCTION_TRACE_TARGET;
use super::*;
use tracing::instrument::WithSubscriber;
fn event(
id: usize,
@ -129,8 +138,6 @@ mod tests {
#[tokio::test]
async fn concurrent_futures_keep_separate_traces_across_yields() {
use tracing::instrument::WithSubscriber;
let first = FunctionTrace::default();
let second = FunctionTrace::default();
let outside = FunctionTrace::default();
@ -161,8 +168,6 @@ mod tests {
#[tokio::test]
async fn concurrent_siblings_keep_the_same_parent() {
use tracing::instrument::WithSubscriber;
let trace = FunctionTrace::default();
concurrent_parent()
.with_subscriber(trace.dispatcher())
@ -212,4 +217,26 @@ mod tests {
vec![event(0, None, "outer"), event(1, Some(0), "inner")]
);
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn instrumented_with_literal_target() {}
#[tokio::test]
async fn literal_instrument_target_matches_filter_constant() {
assert_eq!(FUNCTION_TRACE_TARGET, "litellm::function_trace");
let trace = FunctionTrace::default();
instrumented_with_literal_target()
.with_subscriber(trace.dispatcher())
.await;
let events = trace.events();
assert_eq!(events.len(), 1);
assert_eq!(events[0].id, 0);
assert_eq!(events[0].parent_id, None);
assert_eq!(events[0].function, "instrumented_with_literal_target");
assert_eq!(events[0].module_path, Some(module_path!()));
assert_eq!(events[0].file, Some(file!()));
assert!(events[0].line.is_some());
}
}

View file

@ -1,27 +1,28 @@
//! Harness-only in-process adapters. Never mounted as production routes.
use std::sync::Arc;
use axum::body::{Body, to_bytes};
use axum::http::header::{AUTHORIZATION, CONTENT_TYPE};
use axum::http::{Request, StatusCode};
use litellm_ai_gateway::io::realtime_pool::RealtimePool;
use litellm_ai_gateway::routes;
use litellm_ai_gateway::state::AppState;
use litellm_core::Error;
use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter};
use litellm_python_interop::run_async;
use pyo3::prelude::*;
use serde::Serialize;
use serde_json::Value;
use tower::ServiceExt;
use crate::io::realtime_pool::RealtimePool;
use crate::routes;
use crate::state::AppState;
use crate::errors::core_error_to_pyerr;
#[derive(Debug, Serialize)]
pub struct GatewayResponse {
pub status: u16,
pub body: Value,
struct GatewayResponse {
status: u16,
body: Value,
}
pub async fn messages_request(
async fn messages_request(
model_alias: String,
provider_model: String,
api_base: String,
@ -63,3 +64,19 @@ pub async fn messages_request(
body,
})
}
#[pyfunction]
fn gateway_messages<'py>(
py: Python<'py>,
model_alias: String,
provider_model: String,
api_base: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)] body: Value,
) -> PyResult<Bound<'py, PyAny>> {
let future = messages_request(model_alias, provider_model, api_base, body);
run_async(py, super::capture(future), core_error_to_pyerr)
}
pub(crate) fn register_gateway(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_function(wrap_pyfunction!(gateway_messages, module)?)
}

View file

@ -1,10 +1,15 @@
mod collector;
mod gateway;
use std::fmt::Display;
use std::future::Future;
use litellm_core::observability::{FunctionTrace, FunctionTraceEvent};
use serde::Serialize;
use tracing::instrument::WithSubscriber;
use self::collector::{FunctionTrace, FunctionTraceEvent};
pub(crate) use gateway::register_gateway;
#[derive(Serialize)]
pub(crate) struct TracedResponse<T> {
#[serde(skip_serializing_if = "Option::is_none")]

View file

@ -1 +1 @@
Maps Python profiler frames onto feature-gated Rust span names via an explicit per-case mapping (Rust span name is the identity) and compares steps, order, and nesting of both live traces against a replayed provider response.
Maps Python profiler frames onto Rust span names captured by the bridge’s `trace-parity` feature via an explicit per-case mapping (Rust span name is the identity) and compares steps, order, and nesting of both live traces against a replayed provider response.