refactor(rust): extract inference-responses crate (#44811)

Co-authored-by: Yujong Lee <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-06 07:17:25 -07:00 • committed by GitHub
parent 63babf23e6
commit 785cb12f37
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
22 changed files with 718 additions and 161 deletions

View file

@ -3836,6 +3836,7 @@ dependencies = [
"litellm-host-http",
"litellm-http",
"litellm-inference",
"litellm-inference-responses",
"litellm-inference-transcription",
"litellm-llms",
"litellm-llms-types",
@ -4040,6 +4041,33 @@ dependencies = [
"wiremock",
]
[[package]]
name = "litellm-inference-responses"
version = "0.1.0"
dependencies = [
"bytes",
"futures-util",
"litellm-auth",
"litellm-cache",
"litellm-cache-memory",
"litellm-cache-response",
"litellm-host",
"litellm-host-native",
"litellm-http",
"litellm-inference",
"litellm-llms",
"litellm-llms-types",
"litellm-secrets",
"litellm-tracing",
"reqwest 0.12.28",
"rstest",
"serde_json",
"tokio",
"tokio-tungstenite",
"tracing",
"wiremock",
]
[[package]]
name = "litellm-inference-transcription"
version = "0.1.0"
@ -4142,6 +4170,7 @@ dependencies = [
"litellm-host-python",
"litellm-http",
"litellm-inference",
"litellm-inference-responses",
"litellm-inference-transcription",
"litellm-llms",
"litellm-llms-types",

View file

@ -18,6 +18,7 @@ litellm-traces-clickhouse = { path = "crates/traces-clickhouse" }
litellm-storage-clickhouse = { path = "crates/storage-clickhouse" }
litellm-inference = { path = "crates/inference" }
litellm-inference-transcription = { path = "crates/inference-transcription" }
litellm-inference-responses = { path = "crates/inference-responses" }
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
litellm-gateway = { path = "crates/gateway" }
litellm-gateway-inference = { path = "crates/gateway-inference" }

View file

@ -14,6 +14,7 @@ litellm-auth.workspace = true
litellm-gateway-auth.workspace = true
litellm-inference.workspace = true
litellm-inference-transcription.workspace = true
litellm-inference-responses.workspace = true
litellm-host-http.workspace = true
litellm-host.workspace = true
litellm-http.workspace = true

View file

@ -18,8 +18,9 @@ use axum::{Router, routing::post};
use litellm_http::{ClientVariant, HttpClientConfig, media::UrlPolicy};
use litellm_inference::{
chat_completions::ChatCompletionsRoute, messages::MessagesRoute, ocr::OcrRoute,
resources::CoreResources, responses::ResponsesRoute,
resources::CoreResources,
};
use litellm_inference_responses::ResponsesRoute;
use litellm_inference_transcription::AudioTranscriptionRoute;
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
use litellm_secrets::source::SecretSource;

View file

@ -3,7 +3,7 @@ use std::sync::Arc;
use axum::{Json, body::Bytes, extract::State, response::Response};
use litellm_gateway_auth::AuthenticatedRequest;
use litellm_host_http::Sse;
use litellm_inference::responses::{route::Responses, types::ResponsesCall};
use litellm_inference_responses::{route::Responses, types::ResponsesCall};
use serde_json::json;
use crate::{Error, Gateway, JsonObject, request};

View file

@ -0,0 +1,36 @@
[package]
name = "litellm-inference-responses"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
bytes.workspace = true
futures-util.workspace = true
litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] }
litellm-cache-response.workspace = true
litellm-host.workspace = true
litellm-http.workspace = true
litellm-inference.workspace = true
litellm-llms.workspace = true
litellm-llms-types.workspace = true
litellm-secrets.workspace = true
reqwest.workspace = true
serde_json = { workspace = true, features = ["preserve_order"] }
tokio = { workspace = true, features = ["sync"] }
tokio-tungstenite.workspace = true
tracing.workspace = true
[dev-dependencies]
futures-util.workspace = true
litellm-cache.workspace = true
litellm-cache-memory.workspace = true
litellm-cache-response.workspace = true
litellm-host-native.workspace = true
litellm-http.workspace = true
litellm-inference = { workspace = true, features = ["test-support"] }
litellm-tracing.workspace = true
rstest.workspace = true
tokio.workspace = true
wiremock.workspace = true

View file

@ -36,9 +36,11 @@ pub(super) async fn execute(
)
.await?;
let cache = cache.filter(|_| authenticated.signer.is_none());
let cache_request =
crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire));
crate::caching::execute_streaming::<super::route::Responses, _, _>(
let cache_request = litellm_inference::caching::CacheRequest::from_wire(
identity,
cache.as_ref().map(|_| &wire),
);
litellm_inference::caching::execute_streaming::<super::route::Responses, _, _>(
cache_request,
cache.as_ref().map(|cache| cache.service.clone()),
cache.as_ref().map(|cache| cache.options(cache_options)),
@ -50,7 +52,7 @@ pub(super) async fn execute(
Some(serde_json::Value::Bool(value)) => *value,
Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())),
};
let outbound = crate::outbound::outbound_request(
let outbound = litellm_inference::outbound::outbound_request(
Authenticated {
headers: wire.headers,
signer: authenticated.signer,
@ -59,7 +61,7 @@ pub(super) async fn execute(
&wire.body,
Some(request.timeout.unwrap_or(Duration::from_secs(600))),
)?;
let response = crate::outbound::send(outbound, http)
let response = litellm_inference::outbound::send(outbound, http)
.await
.map_err(network)?;
let status = response.status().as_u16();

View file

@ -1,11 +1,12 @@
pub use crate::error::RouteError as Error;
use litellm_host::observation::ObservationSender;
pub use litellm_inference::RouteError as Error;
pub mod route;
pub mod types;
pub mod websocket;
mod handler;
mod prepare;
pub mod route;
pub mod types;
use std::sync::Arc;
@ -47,9 +48,9 @@ impl ResponsesRoute {
&self,
call: ResponsesCall,
interceptors: &impl Interceptors<Error>,
options: impl Into<crate::CallOptions>,
options: impl Into<litellm_inference::CallOptions>,
) -> Result<ResponsesOutput, Error> {
let crate::CallOptions {
let litellm_inference::CallOptions {
cache: cache_options,
observers,
} = options.into();
@ -75,7 +76,7 @@ impl ResponsesRoute {
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {
crate::diagnostic::call(async {
litellm_inference::diagnostic::call(async {
self.run_provider(call, cache_options, interceptors, observers)
.await
})
@ -90,7 +91,10 @@ impl ResponsesRoute {
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
crate::diagnostic::provider(&request.context.model, &request.context.custom_llm_provider);
litellm_inference::diagnostic::provider(
&request.context.model,
&request.context.custom_llm_provider,
);
let execute: futures_util::future::BoxFuture<'_, Result<ResponsesOutput, Error>> =
Box::pin(handler::execute(
&self.http,

View file

@ -27,9 +27,9 @@ impl ResponsesRoute {
pub fn machine(
self,
call: ResponsesCall,
options: impl Into<crate::CallOptions>,
options: impl Into<litellm_inference::CallOptions>,
) -> HostedMachine<Responses> {
let crate::CallOptions {
let litellm_inference::CallOptions {
cache: cache_options,
observers,
} = options.into();
@ -44,7 +44,7 @@ impl ResponsesRoute {
}
}
impl crate::caching::Cachable for Responses {
impl litellm_inference::caching::Cachable for Responses {
const SURFACE: &'static str = "responses";
fn reusable(response: &Self::Response) -> bool {
@ -56,7 +56,7 @@ impl crate::caching::Cachable for Responses {
}
}
impl crate::caching::StreamCachable for Responses {
impl litellm_inference::caching::StreamCachable for Responses {
const TERMINAL_EVENT: &'static str = "response.completed";
fn replay(data: bytes::Bytes) -> Option<litellm_host::call::OutputOf<Self>> {

View file

@ -40,7 +40,7 @@ impl ResponsesWebSocketConnection {
headers: &HashMap<String, String>,
timeout: Option<Duration>,
) -> Result<Self, Error> {
crate::diagnostic::operation("litellm.websocket.connect_url", async {
litellm_inference::diagnostic::operation("litellm.websocket.connect_url", async {
let mut request = url.into_client_request().map_err(|error| {
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
})?;
@ -86,7 +86,7 @@ impl ResponsesWebSocketConnection {
fields(outcome)
)]
pub async fn send_text(&self, text: String) -> Result<(), Error> {
crate::diagnostic::operation("litellm.websocket.send_text", async {
litellm_inference::diagnostic::operation("litellm.websocket.send_text", async {
let mut socket = self.socket.lock().await;
let Some(socket) = socket.as_mut() else {
return Err(Error::Transport(litellm_http::transport::Error::Network(
@ -107,7 +107,7 @@ impl ResponsesWebSocketConnection {
fields(outcome)
)]
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
crate::diagnostic::operation("litellm.websocket.recv_text", async {
litellm_inference::diagnostic::operation("litellm.websocket.recv_text", async {
let mut socket = self.socket.lock().await;
let Some(socket) = socket.as_mut() else {
return Ok(None);
@ -134,7 +134,7 @@ impl ResponsesWebSocketConnection {
fields(outcome)
)]
pub async fn close(&self) -> Result<(), Error> {
crate::diagnostic::operation("litellm.websocket.close", async {
litellm_inference::diagnostic::operation("litellm.websocket.close", async {
let mut socket = self.socket.lock().await;
if let Some(socket) = socket.as_mut() {
socket.close(None).await.map_err(|error| {

View file

@ -0,0 +1,327 @@
mod support;
use std::{
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Duration,
};
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{
CacheOptions, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService,
ResponseEnvelope, ScopedCache,
};
use litellm_host::interceptors::{
ExecutionFacts, Interceptors, ProviderIdentity, RawResponse, RequestContext, ResultSource,
WireRequest,
};
use litellm_inference::{
RouteError,
caching::{CacheRequest, execute_unary},
};
use rstest::{fixture, rstest};
use serde_json::{Value, json};
fn cache_request(input: Value) -> CacheRequest {
CacheRequest {
identity: ProviderIdentity {
model: "test-model".into(),
provider: "test-provider".into(),
},
input,
}
}
#[fixture]
fn cache() -> Arc<dyn ResponseCacheService> {
cache_with_limit(4096)
}
fn cache_with_limit(max_entry_bytes: usize) -> Arc<dyn ResponseCacheService> {
Arc::new(
ResponseCache::new(Arc::new(InMemoryCache::new(
Some(100),
Some(Duration::from_secs(60)),
)))
.with_config(ResponseCacheConfig {
namespace: "test".into(),
max_entry_bytes,
}),
)
}
struct InvalidEntryCache(
ResponseCache<InMemoryCache<litellm_cache_response::CacheEntry>>,
Value,
);
impl ResponseCacheService for InvalidEntryCache {
fn config(&self) -> &ResponseCacheConfig {
self.0.config()
}
fn lookup<'a>(
&'a self,
request: &'a litellm_cache_response::ResponseCacheRequest,
now: Duration,
) -> futures_util::future::BoxFuture<'a, Result<Option<Value>, litellm_cache::Error>> {
Box::pin(async move {
Ok(self
.0
.async_lookup(request, now)
.await?
.or_else(|| Some(self.1.clone())))
})
}
fn store<'a>(
&'a self,
request: &'a litellm_cache_response::ResponseCacheRequest,
response: Value,
now: Duration,
) -> futures_util::future::BoxFuture<'a, Result<(), litellm_cache::Error>> {
Box::pin(self.0.async_store(request, response, now))
}
}
struct ChangingSecrets {
revision: AtomicUsize,
endpoints: [String; 2],
change_credentials: bool,
}
impl litellm_secrets::source::SecretSource for ChangingSecrets {
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> futures_util::future::BoxFuture<
'a,
Result<Option<litellm_secrets::SecretValue>, litellm_secrets::Error>,
> {
Box::pin(async move {
let revision = self.revision.load(Ordering::SeqCst);
let value = if name.ends_with("_API_KEY") {
Some(format!(
"key-{}",
if self.change_credentials { revision } else { 0 }
))
} else if name.ends_with("_API_BASE") {
Some(self.endpoints[revision].clone())
} else {
None
};
Ok(value.map(litellm_secrets::SecretValue::new))
})
}
}
#[derive(Default)]
struct ChangingHooks {
calls: AtomicUsize,
rewrite: bool,
facts: std::sync::Mutex<Vec<ExecutionFacts>>,
}
impl Interceptors<RouteError> for ChangingHooks {
async fn before_provider_request(
&self,
mut wire: WireRequest,
_: RequestContext,
) -> Result<WireRequest, RouteError> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
if self.rewrite {
wire.body["temperature"] = json!(if call < 2 { 0.1 } else { 0.8 });
}
Ok(wire)
}
async fn after_provider_response(&self, _: RawResponse) -> Result<(), RouteError> {
Ok(())
}
async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), RouteError> {
self.facts.lock().unwrap().push(facts);
Ok(())
}
}
#[rstest]
#[case::chat_completion(json!({"kind":"Response","value":{"id":"chat-1","model":"test","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2}}}))]
#[case::wrong_envelope(json!({"kind":"Stream","value":"data: [DONE]\n\n"}))]
#[tokio::test]
async fn responses_refetches_instead_of_deserializing_another_api_response(
#[case] poisoned: Value,
) {
use litellm_inference_responses::route::Responses;
use litellm_llms_types::formats::responses::ResponsesApiResponse;
let cache: Arc<dyn ResponseCacheService> = Arc::new(InvalidEntryCache(
ResponseCache::new(Arc::new(InMemoryCache::default())),
serde_json::to_value(ResponseEnvelope::new("responses", poisoned)).unwrap(),
));
let calls = AtomicUsize::new(0);
for _ in 0..2 {
let response = execute_unary::<Responses, _, _>(
cache_request(json!({"input":"hello"})),
Some(cache.clone()),
Some(CacheOptions::new(CacheScope::Shared)),
&(),
None,
|| async {
calls.fetch_add(1, Ordering::SeqCst);
Ok(ResponsesApiResponse {
id: "fresh-response".into(),
model: "test".into(),
output: vec![
json!({"type":"message","content":[{"type":"output_text","text":"fresh"}]}),
],
extra: [("status".into(), json!("completed"))]
.into_iter()
.collect(),
})
},
)
.await
.unwrap();
assert_eq!(response.id, "fresh-response");
assert_eq!(response.output[0]["content"][0]["text"], "fresh");
}
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[rstest]
#[case::completed("completed", 1)]
#[case::incomplete("incomplete", 2)]
#[tokio::test]
async fn responses_cache_only_reuses_completed_responses(
cache: Arc<dyn ResponseCacheService>,
#[case] status: &str,
#[case] expected_calls: usize,
) {
use litellm_inference_responses::route::Responses;
use litellm_llms_types::formats::responses::ResponsesApiResponse;
let calls = AtomicUsize::new(0);
for _ in 0..2 {
let response = execute_unary::<Responses, _, _>(
cache_request(json!({"input":"hello"})),
Some(cache.clone()),
Some(CacheOptions::new(CacheScope::Shared)),
&(),
None,
|| async {
let call = calls.fetch_add(1, Ordering::SeqCst);
Ok(ResponsesApiResponse {
id: call.to_string(),
model: "test".into(),
output: Vec::new(),
extra: [("status".into(), json!(status))].into_iter().collect(),
})
},
)
.await
.unwrap();
assert_eq!(response.extra.get("status"), Some(&json!(status)));
}
assert_eq!(calls.load(Ordering::SeqCst), expected_calls);
}
#[rstest]
#[case::credentials("credentials")]
#[case::endpoint("endpoint")]
#[case::callback("callback")]
#[tokio::test]
async fn responses_cache_identity_follows_resolved_configuration_and_request_callbacks(
cache: Arc<dyn ResponseCacheService>,
#[case] change: &str,
) {
use litellm_inference_responses::types::ResponsesCall;
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let first = MockServer::start().await;
let second = MockServer::start().await;
let response = json!({"id":"response-test", "model":"test", "output":[], "status":"completed"});
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(response.clone()))
.expect(if change == "endpoint" { 1 } else { 2 })
.mount(&first)
.await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(response))
.expect(if change == "endpoint" { 1 } else { 0 })
.mount(&second)
.await;
let secrets = Arc::new(ChangingSecrets {
revision: AtomicUsize::new(0),
endpoints: [
first.uri(),
if change == "endpoint" {
second.uri()
} else {
first.uri()
},
],
change_credentials: change == "credentials",
});
let hooks = ChangingHooks {
rewrite: change == "callback",
..Default::default()
};
for call in 0..4 {
secrets
.revision
.store(usize::from(call >= 2), Ordering::SeqCst);
let cache = ScopedCache::new(cache.clone(), CacheScope::Shared);
support::responses_route(secrets.clone())
.with_cache(cache)
.execute(
ResponsesCall {
model: "test".into(),
input: json!("hello"),
optional_params: Default::default(),
api_key: None,
api_base: None,
custom_llm_provider: None,
extra_headers: None,
timeout: None,
},
&hooks,
None,
)
.await
.unwrap();
}
assert_eq!(hooks.calls.load(Ordering::SeqCst), 4);
{
let facts = hooks.facts.lock().unwrap();
assert_eq!(facts[0].source, ResultSource::Provider);
assert_eq!(facts[2].source, ResultSource::Provider);
let (ResultSource::Cache { key: first_key }, ResultSource::Cache { key: second_key }) =
(&facts[1].source, &facts[3].source)
else {
panic!("unchanged effective requests must hit the cache");
};
assert_ne!(first_key, second_key);
}
let requests = first.received_requests().await.unwrap();
if change == "credentials" {
assert_ne!(
requests[0].headers["authorization"],
requests[1].headers["authorization"]
);
}
if change == "callback" {
assert_eq!(
serde_json::from_slice::<Value>(&requests[0].body).unwrap()["temperature"],
0.1
);
assert_eq!(
serde_json::from_slice::<Value>(&requests[1].body).unwrap()["temperature"],
0.8
);
}
first.verify().await;
second.verify().await;
}

View file

@ -6,11 +6,11 @@ use std::sync::Arc;
use futures_util::TryStreamExt;
use litellm_host::{call::HostedCompletion, lifecycle::CallEvent};
use litellm_inference::responses::{
use litellm_inference::test_support::{RecordingSecrets, no_secrets};
use litellm_inference_responses::{
route::Responses,
types::{ResponsesCall, ResponsesOutput},
};
use litellm_inference::test_support::{RecordingSecrets, no_secrets};
use rstest::{fixture, rstest};
use serde_json::json;
use wiremock::ResponseTemplate;
@ -418,7 +418,7 @@ async fn websocket_operations_trace_outcomes_without_capturing_frames_or_credent
traces: TraceCapture,
) {
use futures_util::{SinkExt, StreamExt};
use litellm_inference::responses::websocket::ResponsesWebSocketConnection;
use litellm_inference_responses::websocket::ResponsesWebSocketConnection;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();

View file

@ -0,0 +1,279 @@
//! Shared fixtures for route integration tests: a scripted upstream and a recording
//! secret source.
#![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset
use std::{
ops::ControlFlow,
sync::{Arc, Mutex},
};
use litellm_inference::test_support::{http_config, provider_http, resources};
use litellm_secrets::source::SecretSource;
use serde_json::Value;
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
pub fn responses_route(
secrets: Arc<dyn SecretSource>,
) -> litellm_inference_responses::ResponsesRoute {
let resources = resources();
litellm_inference_responses::ResponsesRoute::new(
provider_http(&resources, &http_config()),
resources.auth,
secrets,
)
}
pub async fn upstream(responses: impl IntoIterator<Item = ResponseTemplate>) -> MockServer {
let server = MockServer::start().await;
respond_in_order(&server, responses).await;
server
}
pub async fn respond_in_order(
server: &MockServer,
responses: impl IntoIterator<Item = ResponseTemplate>,
) {
for response in responses {
Mock::given(any())
.respond_with(response)
.up_to_n_times(1)
.mount(server)
.await;
}
}
pub async fn received(server: &MockServer) -> Vec<Request> {
server
.received_requests()
.await
.expect("request recording is on")
}
pub async fn only_request(server: &MockServer) -> Request {
let [request] = <[Request; 1]>::try_from(received(server).await)
.unwrap_or_else(|requests| panic!("expected one request, got {}", requests.len()));
request
}
pub fn json_response(body: Value) -> ResponseTemplate {
ResponseTemplate::new(200).set_body_json(body)
}
pub fn status_response(status: u16, body: Value) -> ResponseTemplate {
ResponseTemplate::new(status).set_body_json(body)
}
pub trait ReceivedRequest {
fn header(&self, name: &str) -> Option<&str>;
fn header_values(&self, name: &str) -> Vec<&str>;
fn json(&self) -> Value;
fn body_text(&self) -> String;
/// The path and query, as the request line carried them.
fn target(&self) -> String;
fn query(&self, name: &str) -> Option<String>;
}
impl ReceivedRequest for Request {
fn header(&self, name: &str) -> Option<&str> {
self.headers.get(name).and_then(|value| value.to_str().ok())
}
fn header_values(&self, name: &str) -> Vec<&str> {
self.headers
.get_all(name)
.iter()
.filter_map(|value| value.to_str().ok())
.collect()
}
fn json(&self) -> Value {
serde_json::from_slice(&self.body).expect("request body is json")
}
fn body_text(&self) -> String {
String::from_utf8_lossy(&self.body).into_owned()
}
fn target(&self) -> String {
match self.url.query() {
Some(query) => format!("{}?{query}", self.url.path()),
None => self.url.path().to_string(),
}
}
fn query(&self, name: &str) -> Option<String> {
self.url
.query_pairs()
.find_map(|(key, value)| (key == name).then(|| value.into_owned()))
}
}
pub struct RecordingCall<P: litellm_host::protocol::Protocol> {
pub request: Mutex<Option<P::Request>>,
pub events: Arc<CallEvents>,
pub chunks: Mutex<Vec<P::Chunk>>,
pub head: Mutex<Option<P::StreamHead>>,
}
#[derive(Default)]
pub struct CallEvents(pub Observations);
pub struct Observations {
pub sender: litellm_host::observation::ObservationSender,
receiver: Mutex<tokio::sync::mpsc::Receiver<litellm_host::lifecycle::CallEvent>>,
recorded: Mutex<Vec<litellm_host::lifecycle::CallEvent>>,
}
impl Default for Observations {
fn default() -> Self {
let (sender, receiver) = litellm_host::observation::observation_channel(
std::num::NonZeroUsize::new(128).unwrap(),
);
Self {
sender,
receiver: Mutex::new(receiver),
recorded: Mutex::new(Vec::new()),
}
}
}
impl Observations {
pub fn lock(
&self,
) -> std::sync::LockResult<std::sync::MutexGuard<'_, Vec<litellm_host::lifecycle::CallEvent>>>
{
let mut events = self.recorded.lock()?;
let mut receiver = self.receiver.lock().unwrap();
while let Ok(event) = receiver.try_recv() {
events.push(event);
}
Ok(events)
}
}
impl litellm_host::lifecycle::CallObserver for CallEvents {
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
self.0.sender.emit(event);
}
}
impl<P: litellm_host::protocol::Protocol> RecordingCall<P> {
pub fn new(request: P::Request) -> Self {
Self {
request: Mutex::new(Some(request)),
events: Arc::new(CallEvents::default()),
chunks: Mutex::new(Vec::new()),
head: Mutex::new(None),
}
}
}
impl<P: litellm_host::protocol::Protocol> litellm_host::interceptors::Interceptors<P::Error>
for RecordingCall<P>
{
async fn before_provider_request(
&self,
wire: litellm_host::interceptors::WireRequest,
_: litellm_host::interceptors::RequestContext,
) -> Result<litellm_host::interceptors::WireRequest, P::Error> {
Ok(litellm_host::interceptors::WireRequest {
headers: wire
.headers
.into_iter()
.chain([("x-hook".into(), "called".into())])
.collect(),
..wire
})
}
async fn after_provider_response(
&self,
_: litellm_host::interceptors::RawResponse,
) -> Result<(), P::Error> {
Ok(())
}
}
impl<P> RecordingCall<P>
where
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
P::Error: From<litellm_host::machine::MachineFault>,
{
pub fn request(&self) -> Result<P::Request, P::Error> {
self.request
.lock()
.unwrap()
.take()
.ok_or_else(|| litellm_host::machine::MachineFault::Abandoned.into())
}
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, Self> {
litellm_host_native::in_process::Host {
services: &(),
interceptors: self,
stream: self,
observers: Some(&self.events.0.sender),
}
}
}
impl<P> litellm_host_native::in_process::StreamConsumer<P> for RecordingCall<P>
where
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
P::Error: From<litellm_host::machine::MachineFault>,
{
async fn open_stream(&self, head: P::StreamHead) -> Result<ControlFlow<()>, P::Error> {
*self.head.lock().unwrap() = Some(head);
Ok(ControlFlow::Continue(()))
}
async fn send_chunk(&self, chunk: P::Chunk) -> Result<ControlFlow<()>, P::Error> {
self.chunks.lock().unwrap().push(chunk);
Ok(ControlFlow::Continue(()))
}
}
impl<P> litellm_host::lifecycle::CallObserver for RecordingCall<P>
where
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
P::Error: From<litellm_host::machine::MachineFault>,
{
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
self.events.0.sender.emit(event);
}
}
#[derive(Clone, Default)]
pub struct TraceCapture(Arc<Mutex<Vec<Value>>>);
impl TraceCapture {
pub fn logger(&self) -> litellm_tracing::Logger {
litellm_tracing::Logger::new(self.clone())
}
pub fn records(&self) -> Vec<Value> {
self.0.lock().unwrap().clone()
}
pub fn summaries(&self, name: &str) -> Vec<Value> {
self.records()
.into_iter()
.filter(|record| record["span_name"] == name)
.collect()
}
}
impl litellm_tracing::Sink for TraceCapture {
fn enabled(&self, metadata: &litellm_tracing::Metadata<'_>) -> bool {
metadata.target().starts_with("litellm_inference")
}
fn emit(&self, record: &litellm_tracing::Record) {
self.0
.lock()
.unwrap()
.push(Value::Object(record.fields.clone()));
}
}
#[rstest::fixture]
pub fn traces() -> TraceCapture {
TraceCapture::default()
}

View file

@ -52,7 +52,7 @@ impl From<litellm_host::machine::MachineFault> for RouteError {
}
impl RouteError {
pub(crate) fn post_call(error: Self) -> Self {
pub fn post_call(error: Self) -> Self {
Self::PostCallHook(Arc::new(error))
}

View file

@ -10,7 +10,6 @@ pub mod ocr;
pub mod outbound;
pub mod provider;
pub mod resources;
pub mod responses;
#[cfg(feature = "test-support")]
pub mod test_support;

View file

@ -12,8 +12,7 @@ use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt, stream};
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{
CacheOptions, CachePolicy, CacheScope, ResponseCache, ResponseCacheConfig,
ResponseCacheService, ResponseEnvelope,
CacheOptions, CachePolicy, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService,
};
use litellm_host::{
call::{CallOutput, OutputOf},
@ -403,50 +402,6 @@ async fn an_invalid_cached_envelope_is_replaced_by_a_provider_result(#[case] poi
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[rstest]
#[case::chat_completion(json!({"kind":"Response","value":{"id":"chat-1","model":"test","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2}}}))]
#[case::wrong_envelope(json!({"kind":"Stream","value":"data: [DONE]\n\n"}))]
#[tokio::test]
async fn responses_refetches_instead_of_deserializing_another_api_response(
#[case] poisoned: Value,
) {
use litellm_inference::responses::route::Responses;
use litellm_llms_types::formats::responses::ResponsesApiResponse;
let cache: Arc<dyn ResponseCacheService> = Arc::new(InvalidEntryCache(
ResponseCache::new(Arc::new(InMemoryCache::default())),
serde_json::to_value(ResponseEnvelope::new("responses", poisoned)).unwrap(),
));
let calls = AtomicUsize::new(0);
for _ in 0..2 {
let response = execute_unary::<Responses, _, _>(
cache_request(json!({"input":"hello"})),
Some(cache.clone()),
Some(CacheOptions::new(CacheScope::Shared)),
&(),
None,
|| async {
calls.fetch_add(1, Ordering::SeqCst);
Ok(ResponsesApiResponse {
id: "fresh-response".into(),
model: "test".into(),
output: vec![
json!({"type":"message","content":[{"type":"output_text","text":"fresh"}]}),
],
extra: [("status".into(), json!("completed"))]
.into_iter()
.collect(),
})
},
)
.await
.unwrap();
assert_eq!(response.id, "fresh-response");
assert_eq!(response.output[0]["content"][0]["text"], "fresh");
}
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[rstest]
#[case::system("system", json!("answer ALPHA"), json!("answer BETA"))]
#[case::stop_sequences("stop_sequences", json!(["STOP"]), json!(["END"]))]
@ -833,43 +788,6 @@ async fn namespaces_and_surfaces_isolate_entries_on_shared_storage() {
assert_eq!(calls.load(Ordering::SeqCst), 3);
}
#[rstest]
#[case::completed("completed", 1)]
#[case::incomplete("incomplete", 2)]
#[tokio::test]
async fn responses_cache_only_reuses_completed_responses(
cache: Arc<dyn ResponseCacheService>,
#[case] status: &str,
#[case] expected_calls: usize,
) {
use litellm_inference::responses::route::Responses;
use litellm_llms_types::formats::responses::ResponsesApiResponse;
let calls = AtomicUsize::new(0);
for _ in 0..2 {
let response = execute_unary::<Responses, _, _>(
cache_request(json!({"input":"hello"})),
Some(cache.clone()),
Some(CacheOptions::new(CacheScope::Shared)),
&(),
None,
|| async {
let call = calls.fetch_add(1, Ordering::SeqCst);
Ok(ResponsesApiResponse {
id: call.to_string(),
model: "test".into(),
output: Vec::new(),
extra: [("status".into(), json!(status))].into_iter().collect(),
})
},
)
.await
.unwrap();
assert_eq!(response.extra.get("status"), Some(&json!(status)));
}
assert_eq!(calls.load(Ordering::SeqCst), expected_calls);
}
mod support;
use support::traces;
@ -1028,9 +946,6 @@ impl Interceptors<RouteError> for ChangingHooks {
#[case::messages_credentials("messages", "credentials")]
#[case::messages_endpoint("messages", "endpoint")]
#[case::messages_callback("messages", "callback")]
#[case::responses_credentials("responses", "credentials")]
#[case::responses_endpoint("responses", "endpoint")]
#[case::responses_callback("responses", "callback")]
#[tokio::test]
async fn cache_identity_follows_resolved_configuration_and_request_callbacks(
cache: Arc<dyn ResponseCacheService>,
@ -1041,19 +956,14 @@ async fn cache_identity_follows_resolved_configuration_and_request_callbacks(
use litellm_inference::{
chat_completions::{ChatCompletionsRoute, types::ChatCompletionsRequest},
messages::MessagesCall,
responses::types::ResponsesCall,
};
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let first = MockServer::start().await;
let second = MockServer::start().await;
let response = if surface == "responses" {
json!({"id":"response-test", "model":"test", "output":[], "status":"completed"})
} else {
json!({"id":"message-test", "type":"message", "role":"assistant", "model":"test",
"content":[{"type":"text", "text":"answer"}], "stop_reason":"end_turn", "stop_sequence":null,
"usage":{"input_tokens":3,"output_tokens":2}})
};
let response = json!({"id":"message-test", "type":"message", "role":"assistant", "model":"test",
"content":[{"type":"text", "text":"answer"}], "stop_reason":"end_turn", "stop_sequence":null,
"usage":{"input_tokens":3,"output_tokens":2}});
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(response.clone()))
.expect(if change == "endpoint" { 1 } else { 2 })
@ -1116,26 +1026,6 @@ async fn cache_identity_follows_resolved_configuration_and_request_callbacks(
api_key:None,api_base:None,custom_llm_provider:None,extra_headers:None,provider_specific_header:None,timeout:None,shaping:Default::default(),
}, &hooks, None).await.unwrap();
}
"responses" => {
support::responses_route(secrets.clone())
.with_cache(cache)
.execute(
ResponsesCall {
model: "test".into(),
input: json!("hello"),
optional_params: Default::default(),
api_key: None,
api_base: None,
custom_llm_provider: None,
extra_headers: None,
timeout: None,
},
&hooks,
None,
)
.await
.unwrap();
}
_ => unreachable!(),
}
}
@ -1153,12 +1043,10 @@ async fn cache_identity_follows_resolved_configuration_and_request_callbacks(
}
let requests = first.received_requests().await.unwrap();
if change == "credentials" {
let header = if surface == "responses" {
"authorization"
} else {
"x-api-key"
};
assert_ne!(requests[0].headers[header], requests[1].headers[header]);
assert_ne!(
requests[0].headers["x-api-key"],
requests[1].headers["x-api-key"]
);
}
if change == "callback" {
assert_eq!(

View file

@ -37,17 +37,6 @@ pub fn chat_completions_route() -> litellm_inference::chat_completions::ChatComp
)
}
pub fn responses_route(
secrets: Arc<dyn SecretSource>,
) -> litellm_inference::responses::ResponsesRoute {
let resources = resources();
litellm_inference::responses::ResponsesRoute::new(
provider_http(&resources, &http_config()),
resources.auth,
secrets,
)
}
pub fn build_ocr_route(
resources: &litellm_inference::resources::CoreResources,
config: &HttpClientConfig,

View file

@ -38,6 +38,7 @@ litellm-auth-aws.workspace = true
litellm-callbacks-legacy-python.workspace = true
litellm-inference.workspace = true
litellm-inference-transcription.workspace = true
litellm-inference-responses.workspace = true
litellm-core-utils.workspace = true
litellm-http.workspace = true
litellm-llms.workspace = true

View file

@ -1,6 +1,6 @@
mod host;
use litellm_inference::responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use litellm_inference_responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
@ -81,7 +81,7 @@ fn run_public(
py,
arguments,
move |py, arguments, request| {
let route = litellm_inference::responses::ResponsesRoute::new(
let route = litellm_inference_responses::ResponsesRoute::new(
crate::http::provider_client(py, arguments, asynchronous)?
.map_err(crate::http::client_error)?,
crate::http::resources().auth.clone(),

View file

@ -2,7 +2,7 @@ use std::convert::Infallible;
use super::super::inference::InferenceHost;
use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned};
use litellm_inference::responses::{Error, route::Responses, types::ResponsesCall};
use litellm_inference_responses::{Error, route::Responses, types::ResponsesCall};
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,