refactor(rust): extract inference-chat crate (#44827)

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:26 -07:00 • committed by GitHub
parent d7f5b40ab4
commit 6a55e0a8aa
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
24 changed files with 726 additions and 373 deletions

View file

@ -3836,6 +3836,7 @@ dependencies = [
"litellm-host-http",
"litellm-http",
"litellm-inference",
"litellm-inference-chat",
"litellm-inference-messages",
"litellm-inference-responses",
"litellm-inference-transcription",
@ -4042,6 +4043,30 @@ dependencies = [
"wiremock",
]
[[package]]
name = "litellm-inference-chat"
version = "0.1.0"
dependencies = [
"futures-util",
"litellm-auth",
"litellm-cache-memory",
"litellm-cache-response",
"litellm-core-utils",
"litellm-host",
"litellm-host-native",
"litellm-http",
"litellm-inference",
"litellm-llms",
"litellm-llms-types",
"litellm-secrets",
"litellm-tracing",
"rstest",
"serde_json",
"tokio",
"tracing",
"wiremock",
]
[[package]]
name = "litellm-inference-messages"
version = "0.1.0"
@ -4199,6 +4224,7 @@ dependencies = [
"litellm-host-python",
"litellm-http",
"litellm-inference",
"litellm-inference-chat",
"litellm-inference-messages",
"litellm-inference-responses",
"litellm-inference-transcription",

View file

@ -20,6 +20,7 @@ litellm-inference = { path = "crates/inference" }
litellm-inference-transcription = { path = "crates/inference-transcription" }
litellm-inference-responses = { path = "crates/inference-responses" }
litellm-inference-messages = { path = "crates/inference-messages" }
litellm-inference-chat = { path = "crates/inference-chat" }
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
litellm-gateway = { path = "crates/gateway" }
litellm-gateway-inference = { path = "crates/gateway-inference" }

View file

@ -16,6 +16,7 @@ litellm-inference.workspace = true
litellm-inference-transcription.workspace = true
litellm-inference-responses.workspace = true
litellm-inference-messages.workspace = true
litellm-inference-chat.workspace = true
litellm-host-http.workspace = true
litellm-host.workspace = true
litellm-http.workspace = true

View file

@ -6,7 +6,7 @@ use axum::{
extract::{Path, State},
response::{IntoResponse, Response},
};
use litellm_inference::chat_completions::types::ChatCompletionsCall;
use litellm_inference_chat::types::ChatCompletionsCall;
use serde_json::{Map, Value};
use crate::{Error, Gateway, JsonObject, request};

View file

@ -16,9 +16,8 @@ use std::sync::Arc;
use axum::{Router, routing::post};
use litellm_http::{ClientVariant, HttpClientConfig, media::UrlPolicy};
use litellm_inference::{
chat_completions::ChatCompletionsRoute, ocr::OcrRoute, resources::CoreResources,
};
use litellm_inference::{ocr::OcrRoute, resources::CoreResources};
use litellm_inference_chat::ChatCompletionsRoute;
use litellm_inference_messages::MessagesRoute;
use litellm_inference_responses::ResponsesRoute;
use litellm_inference_transcription::AudioTranscriptionRoute;

View file

@ -9,10 +9,8 @@ use axum::{
use litellm_http::{
ClientVariant, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver,
};
use litellm_inference::{
chat_completions::{ChatCompletionsRoute, types::ChatCompletionsRequest},
resources::CoreResources,
};
use litellm_inference::resources::CoreResources;
use litellm_inference_chat::{ChatCompletionsRoute, types::ChatCompletionsRequest};
use rstest::rstest;
use serde_json::{Value, json};
use tower::ServiceExt;

View file

@ -0,0 +1,29 @@
[package]
name = "litellm-inference-chat"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
futures-util.workspace = true
litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] }
litellm-cache-response.workspace = true
litellm-core-utils.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
serde_json = { workspace = true, features = ["preserve_order"] }
tracing.workspace = true
[dev-dependencies]
litellm-cache-memory.workspace = true
litellm-host-native.workspace = true
litellm-inference = { workspace = true, features = ["test-support"] }
litellm-tracing.workspace = true
rstest.workspace = true
tokio.workspace = true
wiremock.workspace = true

View file

@ -0,0 +1,6 @@
/// Full-request timeout ceiling for chat completions provider calls, in
/// seconds. Mirrors the Python chat completions default.
pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600;
/// `object` field every non-streaming chat completion response carries.
pub const CHAT_COMPLETION_OBJECT: &str = "chat.completion";

View file

@ -12,10 +12,7 @@ use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
use serde_json::Value;
use super::Error;
use crate::{
chat_completions::types::ProviderChatCompletionsRequest,
constants::CHAT_COMPLETIONS_TIMEOUT_SECS,
};
use crate::{constants::CHAT_COMPLETIONS_TIMEOUT_SECS, types::ProviderChatCompletionsRequest};
pub(super) async fn execute(
http: &Client,
@ -61,9 +58,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_unary::<super::route::ChatCompletions, _, _>(
let cache_request = litellm_inference::caching::CacheRequest::from_wire(
identity,
cache.as_ref().map(|_| &wire),
);
litellm_inference::caching::execute_unary::<super::route::ChatCompletions, _, _>(
cache_request,
cache.as_ref().map(|cache| cache.service.clone()),
cache.as_ref().map(|cache| cache.options(cache_options)),
@ -80,16 +79,18 @@ pub(super) async fn execute(
timeout,
)?;
let response = crate::outbound::send(outbound, http).await.map_err(|err| {
// Failing to establish the connection means the request never went out,
// so the host can still serve it. Everything else here, a timeout
// above all, may have reached the provider and been answered.
if err.is_connect() || err.is_builder() {
Error::Transport(litellm_http::transport::Error::Connect(err.to_string()))
} else {
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
}
})?;
let response = litellm_inference::outbound::send(outbound, http)
.await
.map_err(|err| {
// Failing to establish the connection means the request never went out,
// so the host can still serve it. Everything else here, a timeout
// above all, may have reached the provider and been answered.
if err.is_connect() || err.is_builder() {
Error::Transport(litellm_http::transport::Error::Connect(err.to_string()))
} else {
Error::Transport(litellm_http::transport::Error::Network(err.to_string()))
}
})?;
let status = response.status();
let text = response.text().await.map_err(|err| {
@ -151,7 +152,7 @@ pub(super) fn outbound_request(
body: &Value,
timeout: Option<Duration>,
) -> Result<OutboundRequest, Error> {
crate::outbound::outbound_request(
litellm_inference::outbound::outbound_request(
authenticated,
url,
body,
@ -176,7 +177,7 @@ mod tests {
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
use super::*;
use crate::chat_completions::{
use crate::{
prepare::{prepare_provider_request, resolve_request},
types::ChatCompletionsRequest,
};

View file

@ -1,17 +1,18 @@
use litellm_host::observation::ObservationSender;
pub mod route;
pub mod types;
pub use crate::error::RouteError as Error;
pub use litellm_inference::RouteError as Error;
mod common_utils;
pub mod constants;
pub(crate) mod handler;
mod prepare;
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
use prepare::{prepare_provider_request, resolve_request};
use crate::chat_completions::types::ChatCompletionsRequest;
use litellm_auth::AuthServices;
use litellm_secrets::source::SecretSource;
use std::sync::Arc;
use types::ChatCompletionsRequest;
#[derive(Clone)]
pub struct ChatCompletionsRoute {
@ -46,9 +47,9 @@ impl ChatCompletionsRoute {
&self,
request: ChatCompletionsRequest<'_>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
options: impl Into<crate::CallOptions>,
options: impl Into<litellm_inference::CallOptions>,
) -> Result<ChatCompletionsResponse, Error> {
let crate::CallOptions {
let litellm_inference::CallOptions {
cache: cache_options,
observers,
} = options.into();
@ -77,7 +78,7 @@ impl ChatCompletionsRoute {
.resolve(&resolved.config.secret_names())
.await?;
let prepared = prepare_provider_request(resolved, snapshot)?;
crate::diagnostic::provider(&prepared.model, &prepared.custom_llm_provider);
litellm_inference::diagnostic::provider(&prepared.model, &prepared.custom_llm_provider);
let execute: futures_util::future::BoxFuture<'_, Result<ChatCompletionsResponse, Error>> =
Box::pin(handler::execute(
&self.http,

View file

@ -10,10 +10,10 @@ use super::{
Error,
common_utils::{chat_completions_provider, string_headers},
};
use crate::chat_completions::types::{
use crate::types::{
ChatCompletionsRequest, ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest,
};
use crate::provider::resolve_llm_provider;
use litellm_inference::provider::resolve_llm_provider;
pub(super) struct ResolvedProvider {
pub(super) model: String,
@ -136,7 +136,7 @@ mod tests {
use serde_json::{Map, Value, json};
use super::{prepare_provider_request, resolve_request};
use crate::chat_completions::{
use crate::{
Error,
types::{ChatCompletionsRequest, ProviderChatCompletionsRequest},
};
@ -473,7 +473,7 @@ mod tests {
)
.await
.expect("resolves");
let signed = crate::chat_completions::handler::outbound_request(
let signed = crate::handler::outbound_request(
authenticated,
prepared.url,
&prepared.body,
@ -532,7 +532,7 @@ mod tests {
)
.await
.expect("resolves");
let error = crate::chat_completions::handler::outbound_request(
let error = crate::handler::outbound_request(
authenticated,
prepared.url,
&prepared.body,

View file

@ -27,9 +27,9 @@ impl ChatCompletionsRoute {
pub fn machine(
self,
call: ChatCompletionsCall,
options: impl Into<crate::CallOptions>,
options: impl Into<litellm_inference::CallOptions>,
) -> HostedMachine<ChatCompletions> {
let crate::CallOptions {
let litellm_inference::CallOptions {
cache: cache_options,
observers,
} = options.into();
@ -59,7 +59,7 @@ impl ChatCompletionsRoute {
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {
crate::diagnostic::unary(async {
litellm_inference::diagnostic::unary(async {
let request = ChatCompletionsRequest {
model: &call.model,
messages: call.messages,
@ -77,6 +77,6 @@ impl ChatCompletionsRoute {
}
}
impl crate::caching::Cachable for ChatCompletions {
impl litellm_inference::caching::Cachable for ChatCompletions {
const SURFACE: &'static str = "chat_completions";
}

View file

@ -0,0 +1,337 @@
mod support;
use std::{
num::NonZeroUsize,
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
time::Duration,
};
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{
CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService,
};
use litellm_host::{
interceptors::{
ExecutionFacts, Interceptors, ProviderIdentity, RawResponse, RequestContext, ResultSource,
WireRequest,
},
lifecycle::{CallEvent, ExecutionEvent},
observation::observation_channel,
};
use litellm_inference::RouteError;
use rstest::{fixture, rstest};
use serde_json::{Value, json};
use support::traces;
#[fixture]
fn cache() -> 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: 4096,
}),
)
}
#[rstest]
#[case::without_cache(false)]
#[case::with_cache(true)]
#[tokio::test]
async fn the_same_route_entrypoint_reports_facts_with_or_without_caching(
cache: Arc<dyn ResponseCacheService>,
#[case] caching: bool,
traces: support::TraceCapture,
) {
use litellm_cache_response::ScopedCache;
use litellm_inference_chat::types::ChatCompletionsRequest;
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let upstream = MockServer::start().await;
let body = json!({"id":"msg-test","type":"message","role":"assistant","model":"cache-test-model",
"content":[{"type":"text","text":"cached answer"}],"stop_reason":"end_turn",
"stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}});
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(body))
.expect(if caching { 1 } else { 2 })
.mount(&upstream)
.await;
let route = support::chat_completions_route();
let route = if caching {
route.with_cache(ScopedCache::new(cache, CacheScope::Shared))
} else {
route
};
let (observer, mut events) = observation_channel(NonZeroUsize::new(16).unwrap());
let base = upstream.uri();
for _ in 0..2 {
let response = traces
.logger()
.instrument(route.execute(
ChatCompletionsRequest {
model: "anthropic/cache-test-model",
messages: json!([{"role":"user","content":"hello"}]),
optional_params: [("max_tokens".into(), json!(16))].into_iter().collect(),
api_key: Some("test-key"),
api_base: Some(&base),
custom_llm_provider: None,
extra_headers: None,
timeout: None,
},
&(),
Some(observer.clone()),
))
.await
.unwrap();
assert_eq!(
serde_json::to_value(response).unwrap()["usage"]["total_tokens"],
15
);
}
let facts: Vec<_> = std::iter::from_fn(|| events.try_recv().ok())
.filter_map(|event| match event {
CallEvent::Execution(ExecutionEvent::ResultReady { facts }) => Some(facts),
_ => None,
})
.collect();
assert_eq!(facts.len(), 2);
assert_eq!(
facts[0].provider,
ProviderIdentity {
model: "cache-test-model".into(),
provider: "anthropic".into()
}
);
assert_eq!(facts[1].provider, facts[0].provider);
assert_eq!(facts[0].source, ResultSource::Provider);
match &facts[1].source {
ResultSource::Provider => assert!(!caching),
ResultSource::Cache { key } => {
assert!(caching);
assert!(!key.is_empty());
}
}
let summaries = traces.summaries("litellm.route");
assert_eq!(summaries.len(), 2);
for summary in summaries {
assert_eq!(summary["provider"], "anthropic");
assert_eq!(summary["resolved_model"], "cache-test-model");
assert_eq!(summary["outcome"], "success");
}
upstream.verify().await;
}
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::credentials("credentials")]
#[case::endpoint("endpoint")]
#[case::callback("callback")]
#[tokio::test]
async fn chat_cache_identity_follows_resolved_configuration_and_request_callbacks(
cache: Arc<dyn ResponseCacheService>,
#[case] change: &str,
) {
use litellm_cache_response::ScopedCache;
use litellm_inference_chat::{ChatCompletionsRoute, types::ChatCompletionsRequest};
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let first = MockServer::start().await;
let second = MockServer::start().await;
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 })
.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);
ChatCompletionsRoute::new(
litellm_http::Client::plain_for_test(),
Arc::new(Default::default()),
secrets.clone(),
)
.with_cache(cache)
.execute(
ChatCompletionsRequest {
model: "anthropic/cache-test-model",
messages: json!([{"role":"user","content":"hello"}]),
optional_params: [("max_tokens".into(), json!(32))].into_iter().collect(),
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["x-api-key"],
requests[1].headers["x-api-key"]
);
}
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;
}
#[rstest]
#[tokio::test]
async fn signed_requests_bypass_response_caching(cache: Arc<dyn ResponseCacheService>) {
use litellm_cache_response::ScopedCache;
use litellm_inference_chat::types::ChatCompletionsRequest;
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let upstream = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"output":{"message":{"role":"assistant","content":[{"text":"answer"}]}},
"stopReason":"end_turn", "usage":{"inputTokens":3,"outputTokens":2,"totalTokens":5}
})))
.expect(2)
.mount(&upstream)
.await;
let route =
support::chat_completions_route().with_cache(ScopedCache::new(cache, CacheScope::Shared));
let hooks = ChangingHooks::default();
for _ in 0..2 {
let response = route.execute(ChatCompletionsRequest {
model:"bedrock/anthropic.cache-test-model",
messages:json!([{"role":"user","content":"hello"}]),
optional_params:json!({"aws_access_key_id":"test-access","aws_secret_access_key":"test-secret","aws_region_name":"eu-west-1"}).as_object().unwrap().clone(),
api_key:None,api_base:Some(&upstream.uri()),custom_llm_provider:None,extra_headers:None,timeout:None,
}, &hooks, None).await.unwrap();
assert_eq!(
serde_json::to_value(response).unwrap()["usage"]["total_tokens"],
5
);
}
assert!(
hooks
.facts
.lock()
.unwrap()
.iter()
.all(|facts| facts.source == ResultSource::Provider)
);
upstream.verify().await;
}

View file

@ -6,7 +6,7 @@ use litellm_host::{
use std::time::Duration;
use litellm_http::transport::Error as TransportError;
use litellm_inference::chat_completions::{Error, types::ChatCompletionsRequest};
use litellm_inference_chat::{Error, types::ChatCompletionsRequest};
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
use rstest::{fixture, rstest};
use serde_json::{Map, Value, json};
@ -258,7 +258,7 @@ async fn direct_and_hosted_calls_share_hooks_and_lifecycle(
#[case] hosted: bool,
) {
use litellm_host::{call::HostedCompletion, lifecycle::CallEvent};
use litellm_inference::chat_completions::route::ChatCompletions;
use litellm_inference_chat::route::ChatCompletions;
let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await;
let base = upstream.uri();

View file

@ -0,0 +1,281 @@
//! 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, no_secrets, provider_http, resources};
use serde_json::Value;
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
/// A port nothing listens on, for calls that must fail before any request is sent.
pub const UNREACHABLE_BASE: &str = "http://127.0.0.1:1";
pub fn chat_completions_route() -> litellm_inference_chat::ChatCompletionsRoute {
let resources = resources();
litellm_inference_chat::ChatCompletionsRoute::new(
provider_http(&resources, &http_config()),
resources.auth,
no_secrets(),
)
}
/// Starts an upstream that answers its n-th request with the n-th response and 404s after.
pub async fn upstream(responses: impl IntoIterator<Item = ResponseTemplate>) -> MockServer {
let server = MockServer::start().await;
respond_in_order(&server, responses).await;
server
}
/// Scripts responses on a started server, for responses that need its address.
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

@ -1,8 +1 @@
pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com";
/// Full-request timeout ceiling for chat completions provider calls, in
/// seconds. Mirrors the Python chat completions default.
pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600;
/// `object` field every non-streaming chat completion response carries.
pub const CHAT_COMPLETION_OBJECT: &str = "chat.completion";

View file

@ -2,7 +2,6 @@ pub mod context;
pub mod diagnostic;
pub mod caching;
pub mod chat_completions;
pub mod constants;
pub mod error;
pub mod ocr;

View file

@ -729,309 +729,3 @@ async fn namespaces_and_surfaces_isolate_entries_on_shared_storage() {
);
assert_eq!(calls.load(Ordering::SeqCst), 3);
}
mod support;
use support::traces;
#[rstest]
#[case::without_cache(false)]
#[case::with_cache(true)]
#[tokio::test]
async fn the_same_route_entrypoint_reports_facts_with_or_without_caching(
cache: Arc<dyn ResponseCacheService>,
#[case] caching: bool,
traces: support::TraceCapture,
) {
use litellm_cache_response::ScopedCache;
use litellm_inference::chat_completions::types::ChatCompletionsRequest;
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let upstream = MockServer::start().await;
let body = json!({"id":"msg-test","type":"message","role":"assistant","model":"cache-test-model",
"content":[{"type":"text","text":"cached answer"}],"stop_reason":"end_turn",
"stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}});
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(body))
.expect(if caching { 1 } else { 2 })
.mount(&upstream)
.await;
let route = support::chat_completions_route();
let route = if caching {
route.with_cache(ScopedCache::new(cache, CacheScope::Shared))
} else {
route
};
let (observer, mut events) = observation_channel(NonZeroUsize::new(16).unwrap());
let base = upstream.uri();
for _ in 0..2 {
let response = traces
.logger()
.instrument(route.execute(
ChatCompletionsRequest {
model: "anthropic/cache-test-model",
messages: json!([{"role":"user","content":"hello"}]),
optional_params: [("max_tokens".into(), json!(16))].into_iter().collect(),
api_key: Some("test-key"),
api_base: Some(&base),
custom_llm_provider: None,
extra_headers: None,
timeout: None,
},
&(),
Some(observer.clone()),
))
.await
.unwrap();
assert_eq!(
serde_json::to_value(response).unwrap()["usage"]["total_tokens"],
15
);
}
let facts: Vec<_> = std::iter::from_fn(|| events.try_recv().ok())
.filter_map(|event| match event {
CallEvent::Execution(ExecutionEvent::ResultReady { facts }) => Some(facts),
_ => None,
})
.collect();
assert_eq!(facts.len(), 2);
assert_eq!(
facts[0].provider,
ProviderIdentity {
model: "cache-test-model".into(),
provider: "anthropic".into()
}
);
assert_eq!(facts[1].provider, facts[0].provider);
assert_eq!(facts[0].source, ResultSource::Provider);
match &facts[1].source {
ResultSource::Provider => assert!(!caching),
ResultSource::Cache { key } => {
assert!(caching);
assert!(!key.is_empty());
}
}
let summaries = traces.summaries("litellm.route");
assert_eq!(summaries.len(), 2);
for summary in summaries {
assert_eq!(summary["provider"], "anthropic");
assert_eq!(summary["resolved_model"], "cache-test-model");
assert_eq!(summary["outcome"], "success");
}
upstream.verify().await;
}
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_credentials("chat", "credentials")]
#[case::chat_endpoint("chat", "endpoint")]
#[case::chat_callback("chat", "callback")]
#[tokio::test]
async fn cache_identity_follows_resolved_configuration_and_request_callbacks(
cache: Arc<dyn ResponseCacheService>,
#[case] surface: &str,
#[case] change: &str,
) {
use litellm_cache_response::ScopedCache;
use litellm_inference::chat_completions::{
ChatCompletionsRoute, types::ChatCompletionsRequest,
};
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let first = MockServer::start().await;
let second = MockServer::start().await;
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 })
.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);
match surface {
"chat" => {
ChatCompletionsRoute::new(
litellm_http::Client::plain_for_test(),
Arc::new(Default::default()),
secrets.clone(),
)
.with_cache(cache)
.execute(
ChatCompletionsRequest {
model: "anthropic/cache-test-model",
messages: json!([{"role":"user","content":"hello"}]),
optional_params: [("max_tokens".into(), json!(32))].into_iter().collect(),
api_key: None,
api_base: None,
custom_llm_provider: None,
extra_headers: None,
timeout: None,
},
&hooks,
None,
)
.await
.unwrap();
}
_ => unreachable!(),
}
}
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["x-api-key"],
requests[1].headers["x-api-key"]
);
}
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;
}
#[rstest]
#[tokio::test]
async fn signed_requests_bypass_response_caching(cache: Arc<dyn ResponseCacheService>) {
use litellm_cache_response::ScopedCache;
use litellm_inference::chat_completions::types::ChatCompletionsRequest;
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let upstream = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(json!({
"output":{"message":{"role":"assistant","content":[{"text":"answer"}]}},
"stopReason":"end_turn", "usage":{"inputTokens":3,"outputTokens":2,"totalTokens":5}
})))
.expect(2)
.mount(&upstream)
.await;
let route =
support::chat_completions_route().with_cache(ScopedCache::new(cache, CacheScope::Shared));
let hooks = ChangingHooks::default();
for _ in 0..2 {
let response = route.execute(ChatCompletionsRequest {
model:"bedrock/anthropic.cache-test-model",
messages:json!([{"role":"user","content":"hello"}]),
optional_params:json!({"aws_access_key_id":"test-access","aws_secret_access_key":"test-secret","aws_region_name":"eu-west-1"}).as_object().unwrap().clone(),
api_key:None,api_base:Some(&upstream.uri()),custom_llm_provider:None,extra_headers:None,timeout:None,
}, &hooks, None).await.unwrap();
assert_eq!(
serde_json::to_value(response).unwrap()["usage"]["total_tokens"],
5
);
}
assert!(
hooks
.facts
.lock()
.unwrap()
.iter()
.all(|facts| facts.source == ResultSource::Provider)
);
upstream.verify().await;
}

View file

@ -9,7 +9,6 @@ use std::{
};
use litellm_http::HttpClientConfig;
use litellm_inference::test_support::{http_config, no_secrets, provider_http, resources};
use litellm_secrets::source::SecretSource;
use serde_json::Value;
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
@ -17,15 +16,6 @@ use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
/// A port nothing listens on, for calls that must fail before any request is sent.
pub const UNREACHABLE_BASE: &str = "http://127.0.0.1:1";
pub fn chat_completions_route() -> litellm_inference::chat_completions::ChatCompletionsRoute {
let resources = resources();
litellm_inference::chat_completions::ChatCompletionsRoute::new(
provider_http(&resources, &http_config()),
resources.auth,
no_secrets(),
)
}
pub fn build_ocr_route(
resources: &litellm_inference::resources::CoreResources,
config: &HttpClientConfig,

View file

@ -40,6 +40,7 @@ litellm-inference.workspace = true
litellm-inference-transcription.workspace = true
litellm-inference-responses.workspace = true
litellm-inference-messages.workspace = true
litellm-inference-chat.workspace = true
litellm-core-utils.workspace = true
litellm-http.workspace = true
litellm-llms.workspace = true

View file

@ -3,9 +3,7 @@ mod host;
use pyo3::types::{PyDict, PyTuple};
use crate::execution::{run_async, run_sync};
use litellm_inference::chat_completions::{
ChatCompletionsRoute, Error, types::ChatCompletionsRequest,
};
use litellm_inference_chat::{ChatCompletionsRoute, Error, types::ChatCompletionsRequest};
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
use pyo3::prelude::*;
use serde_json::{Map, Value};

View file

@ -2,9 +2,7 @@ use std::convert::Infallible;
use super::super::inference::InferenceHost;
use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned};
use litellm_inference::chat_completions::{
Error, route::ChatCompletions, types::ChatCompletionsCall,
};
use litellm_inference_chat::{Error, route::ChatCompletions, types::ChatCompletionsCall};
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,