refactor(rust): orchestrate Messages route execution (#43719)

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-09-29 19:54:12 +00:00 • committed by GitHub
parent fb74957ddd
commit 273489824a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
30 changed files with 711 additions and 379 deletions

View file

@ -24,6 +24,6 @@ Keep unary caching independent of stream-only methods. Store streams only after
Test each contract in its owner: storage capabilities in backend tests, envelopes and freshness here, reuse and replay in core, Python callback and fallback behavior at the bridge, and HTTP behavior at the gateway. Run backend contract checks and Python response-codec fixtures before exposing a new backend
`ScopedCache` requires an explicit shared or isolated scope at construction. `CacheOptions` has no default sharing policy. Callers may override policy per invocation without replacing the attached service. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec
`ScopedCache` requires an explicit shared or isolated scope at construction. Per-call `CachePolicy` controls reads, writes, expiry, and freshness without replacing the attached scope or service. `CacheOptions` binds that policy to an explicit scope for storage requests and has no default sharing policy. Versioned native envelopes reject incompatible API surfaces and versions as misses; this envelope is distinct from the legacy Python response codec
Response storage is not the source of budget or rate-limit coordination dependencies. Keep counters, reservations, and atomic admission operations out of `ResponseCacheService`, including when both services happen to use Redis

View file

@ -17,6 +17,6 @@ pub use exact::{ConnectionProbe, ExactResponseCache};
pub use response::{ResponseCache, ResponseCacheRequest};
pub use service::{
CacheOptions, CacheScope, ResponseCacheConfig, ResponseCacheService, ResponseEnvelope,
ScopedCache,
CacheOptions, CachePolicy, CacheScope, ResponseCacheConfig, ResponseCacheService,
ResponseEnvelope, ScopedCache,
};

View file

@ -73,32 +73,35 @@ pub enum CacheScope {
Isolated(String),
}
#[derive(Clone)]
pub struct CacheOptions {
#[derive(Clone, Copy, Default)]
pub struct CachePolicy {
pub caching: Option<bool>,
pub no_cache: bool,
pub no_store: bool,
pub ttl: Option<Duration>,
pub max_age: Option<Duration>,
}
impl CachePolicy {
pub fn enabled(&self) -> bool {
self.caching != Some(false) && !(self.no_cache && self.no_store)
}
}
#[derive(Clone)]
pub struct CacheOptions {
pub policy: CachePolicy,
pub scope: CacheScope,
}
impl CacheOptions {
pub fn new(scope: CacheScope) -> Self {
Self {
caching: None,
no_cache: false,
no_store: false,
ttl: None,
max_age: None,
policy: CachePolicy::default(),
scope,
}
}
pub fn enabled(&self) -> bool {
self.caching != Some(false) && !(self.no_cache && self.no_store)
}
pub fn request(self, namespace: &str, surface: &str, mut input: Value) -> ResponseCacheRequest {
input.sort_all_objects();
let scope = match self.scope {
@ -128,13 +131,15 @@ impl CacheOptions {
supported_call_type: true,
native_backend: true,
default_on: true,
caching: self.caching,
no_cache: self.no_cache,
no_store: self.no_store,
caching: self.policy.caching,
no_cache: self.policy.no_cache,
no_store: self.policy.no_store,
..Default::default()
},
context: ExactCacheContext { ttl: self.ttl },
max_age: self.max_age,
context: ExactCacheContext {
ttl: self.policy.ttl,
},
max_age: self.policy.max_age,
}
}
}
@ -171,7 +176,10 @@ impl ScopedCache {
Self { service, scope }
}
pub fn options(&self, overrides: Option<CacheOptions>) -> CacheOptions {
overrides.unwrap_or_else(|| CacheOptions::new(self.scope.clone()))
pub fn options(&self, policy: Option<CachePolicy>) -> CacheOptions {
CacheOptions {
policy: policy.unwrap_or_default(),
scope: self.scope.clone(),
}
}
}

View file

@ -129,11 +129,20 @@ async fn isolated_policy_controls_actual_entry_reuse(
#[case] first: &str,
#[case] second: &str,
#[case] hit: bool,
#[values(false, true)] override_policy: bool,
) {
use litellm_cache_response::{CacheOptions, CacheScope};
let service = ResponseCache::new(Arc::new(InMemoryCache::<CacheEntry>::default()));
let request =
|scope| CacheOptions::new(scope).request("test", "messages", json!({"prompt":"hello"}));
use litellm_cache_response::{CachePolicy, CacheScope, ScopedCache};
let service = Arc::new(ResponseCache::new(Arc::new(
InMemoryCache::<CacheEntry>::default(),
)));
let request = |scope| {
ScopedCache::new(service.clone(), scope)
.options(override_policy.then_some(CachePolicy {
ttl: Some(Duration::from_secs(30)),
..CachePolicy::default()
}))
.request("test", "messages", json!({"prompt":"hello"}))
};
service
.async_store(
&request(CacheScope::Isolated(first.into())),

View file

@ -36,7 +36,9 @@ Not here: serving HTTP (axum routes, extractors), config file reading, rollout s
## Response caching and accounting boundary
Attach a `litellm_cache_response::ScopedCache` with `route.with_cache(cache)`. Cached and uncached routes use the same `execute` and `machine` methods. `CallOptions` carries per-call cache overrides and observation; attaching a service does not change the execution contract
Attach a `litellm_cache_response::ScopedCache` with `route.with_cache(cache)`. Cached and uncached routes use the same `execute` and `machine` methods. `CallOptions` carries a scope-free `CachePolicy` and observation; per-call policy never replaces the attached scope or service
Messages groups per-call dependencies in `CallContext` and explicitly sequences cache lookup, provider execution, result acceptance, and cache storage. Provider transport does not own cache orchestration. Stream capture remains in the shared cache implementation
Core owns request identity, typed response reconstruction and stream capture/replay. `cache-response` owns cache policy, namespacing, scope encoding, versioned envelopes and freshness. The SDK explicitly chooses shared scope. The gateway derives isolated scope from authenticated identity before attaching its service

View file

@ -1,5 +1,6 @@
use std::{
future::Future,
marker::PhantomData,
sync::Arc,
time::{Duration, SystemTime, UNIX_EPOCH},
};
@ -7,7 +8,8 @@ use std::{
use bytes::{Bytes, BytesMut};
use futures_util::{StreamExt, TryStreamExt, stream};
use litellm_cache_response::{
CacheOptions, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope, cache_key,
CacheOptions, CachePolicy, ResponseCacheRequest, ResponseCacheService, ResponseEnvelope,
ScopedCache, cache_key,
};
use litellm_host::{
call::{CallOutput, OutputOf},
@ -77,7 +79,7 @@ impl CacheSession {
options: Option<CacheOptions>,
request: &CacheRequest,
) -> Option<Self> {
let options = options.filter(CacheOptions::enabled)?;
let options = options.filter(|options| options.policy.enabled())?;
let service = service?;
let input = request.input.clone();
let request = options.request(&service.config().namespace, P::SURFACE, input);
@ -198,81 +200,137 @@ where
let identity = request.identity.clone();
crate::diagnostic::provider(&identity.model, &identity.provider);
let session = CacheSession::prepare::<P>(cache, options, &request);
let hit = match &session {
Some(session) => session.lookup::<P>().await.and_then(|entry| {
let output = match entry {
CachedOutput::Response(response) => Some(CallOutput::Complete(response)),
CachedOutput::Stream(data) => P::replay(Bytes::from(data)),
};
output.map(|output| (output, cache_key(&session.request.key)))
}),
None => None,
let cache = CallCache::<P> {
session,
protocol: PhantomData,
};
let hit = cache.lookup().await;
let (output, source) = match hit {
Some((output, key)) => (output, ResultSource::Cache { key }),
Some(hit) => hit,
None => (provider().await?, ResultSource::Provider),
};
let from_provider = source == ResultSource::Provider;
publish(
ExecutionFacts {
provider: identity,
source,
source: source.clone(),
},
interceptors,
observers,
)
.await?;
let Some(session) =
session.filter(|session| from_provider && session.request.controls.writes())
else {
return Ok(output);
};
match output {
CallOutput::Complete(response) => {
session.store_response::<P>(&response).await;
Ok(CallOutput::Complete(response))
}
CallOutput::Stream { head, chunks } => {
let captured = stream::try_unfold(
(chunks, Some(Vec::<u8>::new()), session),
|(mut chunks, captured, session)| async move {
match chunks.try_next().await? {
Some(chunk) => {
let captured = captured.and_then(|mut data| {
let bytes = P::bytes(&chunk);
if data.len().saturating_add(bytes.len())
> session.service.config().max_entry_bytes
{
return None;
}
data.extend_from_slice(bytes);
Some(data)
});
Ok(Some((chunk, (chunks, captured, session))))
}
None => {
if let Some(data) = captured
&& let Ok(text) = String::from_utf8(data)
&& successful_stream(&text, P::TERMINAL_EVENT)
&& let Ok(entry) = serde_json::to_value(ResponseEnvelope::new(
P::SURFACE,
CachedOutput::<Value>::Stream(text),
))
{
session.store(entry).await;
}
Ok::<_, RouteError>(None)
}
}
},
)
.boxed();
Ok(CallOutput::Stream {
head,
chunks: captured,
Ok(cache.finish(output, &source).await)
}
pub(crate) struct CallCache<P> {
session: Option<CacheSession>,
protocol: PhantomData<P>,
}
impl<P: StreamCachable> CallCache<P> {
pub(crate) fn from_wire(
cache: Option<&ScopedCache>,
policy: CachePolicy,
identity: &ProviderIdentity,
wire: &WireRequest,
) -> Self {
let session = cache.and_then(|cache| {
if !policy.enabled() {
return None;
}
let options = cache.options(Some(policy));
let request = CacheRequest::from_wire(identity.clone(), Some(wire));
Some(CacheSession {
request: options.request(
&cache.service.config().namespace,
P::SURFACE,
request.input,
),
service: cache.service.clone(),
})
});
Self {
session,
protocol: PhantomData,
}
}
pub(crate) async fn lookup(&self) -> Option<(OutputOf<P>, ResultSource)>
where
P::Response: DeserializeOwned,
{
let session = self.session.as_ref()?;
let output = match session.lookup::<P>().await? {
CachedOutput::Response(response) => CallOutput::Complete(response),
CachedOutput::Stream(data) => P::replay(Bytes::from(data))?,
};
Some((
output,
ResultSource::Cache {
key: cache_key(&session.request.key),
},
))
}
pub(crate) async fn finish(self, output: OutputOf<P>, source: &ResultSource) -> OutputOf<P>
where
P::Response: Serialize,
{
let Some(session) = self.session.filter(|session| {
*source == ResultSource::Provider && session.request.controls.writes()
}) else {
return output;
};
match output {
CallOutput::Complete(response) => {
session.store_response::<P>(&response).await;
CallOutput::Complete(response)
}
CallOutput::Stream { head, chunks } => CallOutput::Stream {
head,
chunks: capture_stream::<P>(chunks, session),
},
}
}
}
fn capture_stream<P: StreamCachable>(
chunks: futures_util::stream::BoxStream<'static, Result<P::Chunk, RouteError>>,
session: CacheSession,
) -> futures_util::stream::BoxStream<'static, Result<P::Chunk, RouteError>> {
stream::try_unfold(
(chunks, Some(Vec::<u8>::new()), session),
|(mut chunks, captured, session)| async move {
match chunks.try_next().await? {
Some(chunk) => {
let captured = captured.and_then(|mut data| {
let bytes = P::bytes(&chunk);
if data.len().saturating_add(bytes.len())
> session.service.config().max_entry_bytes
{
return None;
}
data.extend_from_slice(bytes);
Some(data)
});
Ok(Some((chunk, (chunks, captured, session))))
}
None => {
if let Some(data) = captured
&& let Ok(text) = String::from_utf8(data)
&& successful_stream(&text, P::TERMINAL_EVENT)
&& let Ok(entry) = serde_json::to_value(ResponseEnvelope::new(
P::SURFACE,
CachedOutput::<Value>::Stream(text),
))
{
session.store(entry).await;
}
Ok::<_, RouteError>(None)
}
}
},
)
.boxed()
}
fn now() -> Duration {

View file

@ -22,7 +22,7 @@ pub(super) async fn execute(
auth: &AuthServices,
request: ProviderChatCompletionsRequest,
cache: Option<litellm_cache_response::ScopedCache>,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {

View file

@ -67,7 +67,7 @@ impl ChatCompletionsRoute {
async fn run(
&self,
request: ChatCompletionsRequest<'_>,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {

View file

@ -55,7 +55,7 @@ impl ChatCompletionsRoute {
pub(super) async fn run_call(
&self,
call: ChatCompletionsCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ChatCompletionsResponse, Error> {

View file

@ -0,0 +1,48 @@
use litellm_cache_response::CachePolicy;
use litellm_host::{
interceptors::{ExecutionFacts, Interceptors, RawResponse},
lifecycle::{CallEvent, ExecutionEvent},
observation::ObservationSender,
};
use crate::{CallOptions, RouteError};
pub(crate) struct CallContext<'a, I> {
pub interceptors: &'a I,
pub observers: Option<ObservationSender>,
pub cache: CachePolicy,
}
impl<'a, I: Interceptors<RouteError>> CallContext<'a, I> {
pub fn new(interceptors: &'a I, options: CallOptions) -> Self {
Self {
interceptors,
observers: options.observers,
cache: options.cache.unwrap_or_default(),
}
}
pub async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), RouteError> {
if let Some(observers) = &self.observers {
observers.emit(CallEvent::Execution(ExecutionEvent::ResultReady {
facts: facts.clone(),
}));
}
self.interceptors.result_ready(facts).await
}
pub async fn response_received(&self, body: &str) -> Result<(), RouteError> {
let raw = RawResponse {
body: body.to_owned(),
};
if let Some(observers) = &self.observers {
observers.emit(CallEvent::Execution(
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
));
}
self.interceptors
.after_provider_response(raw)
.await
.map_err(RouteError::post_call)
}
}

View file

@ -1,3 +1,4 @@
mod context;
mod diagnostic;
pub mod audio_transcription;
@ -16,7 +17,7 @@ pub use error::RouteError;
#[derive(Clone, Default)]
pub struct CallOptions {
pub cache: Option<litellm_cache_response::CacheOptions>,
pub cache: Option<litellm_cache_response::CachePolicy>,
pub observers: Option<litellm_host::observation::ObservationSender>,
}
@ -29,8 +30,8 @@ impl From<Option<litellm_host::observation::ObservationSender>> for CallOptions
}
}
impl From<litellm_cache_response::CacheOptions> for CallOptions {
fn from(cache: litellm_cache_response::CacheOptions) -> Self {
impl From<litellm_cache_response::CachePolicy> for CallOptions {
fn from(cache: litellm_cache_response::CachePolicy) -> Self {
Self {
cache: Some(cache),
observers: None,

View file

@ -1,10 +1,8 @@
use litellm_host::{lifecycle::ExecutionEvent, observation::ObservationSender};
use std::time::Duration;
use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt, stream::BoxStream};
use litellm_auth::AuthServices;
use litellm_host::interceptors::{Interceptors, RawResponse, RequestContext, WireRequest};
use litellm_host::interceptors::{Interceptors, ProviderIdentity, RequestContext, WireRequest};
use litellm_http::transport::Error as TransportError;
use litellm_llms::base_llm::{
auth::{Authenticated, resolve_auth},
@ -18,102 +16,130 @@ use litellm_tracing::ByteChunk;
use serde_json::Value;
use super::{
Error, MessagesCallResponse, common_utils::truncate_error_body,
Error, MessagesCallResponse, MessagesRoute, common_utils::truncate_error_body,
prepare::ProviderMessagesRequest,
};
use crate::{constants::MESSAGES_TIMEOUT_SECS, outbound::outbound_request};
use crate::{constants::MESSAGES_TIMEOUT_SECS, context::CallContext, outbound::outbound_request};
pub(super) async fn execute(
http: &litellm_http::Client,
auth: &AuthServices,
request: ProviderMessagesRequest,
cache: Option<litellm_cache_response::ScopedCache>,
cache_options: Option<litellm_cache_response::CacheOptions>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<MessagesCallResponse, Error> {
let ProviderMessagesRequest {
provider,
url,
body,
environment,
timeout,
api_key,
} = request;
let stream = body.params.stream == Some(true);
let context = RequestContext {
model: body.model.clone(),
custom_llm_provider: provider.as_str().to_string(),
optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?,
secret_fields: Vec::new(),
api_key,
};
let authenticated = resolve_auth(auth, environment, &|key| std::env::var(key).ok()).await?;
let identity = litellm_host::interceptors::ProviderIdentity {
model: context.model.clone(),
provider: context.custom_llm_provider.clone(),
};
let wire = interceptors
.before_provider_request(
WireRequest {
url,
headers: authenticated.headers,
body: serde_json::to_value(&body).map_err(serialize_failure)?,
},
context,
)
.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::Messages, _, _>(
cache_request,
cache.as_ref().map(|cache| cache.service.clone()),
cache.as_ref().map(|cache| cache.options(cache_options)),
interceptors,
observers,
|| async move {
let provider_name = provider.as_str();
log_request_body(provider_name, stream, &wire.body);
let response = send(
http,
Authenticated {
headers: wire.headers,
signer: authenticated.signer,
pub(super) struct ProviderCall {
pub identity: ProviderIdentity,
pub wire: WireRequest,
provider: super::common_utils::MessagesProvider,
signer: Option<litellm_auth_aws::SigV4Signer>,
timeout: Option<Duration>,
stream: bool,
}
impl ProviderCall {
pub fn cacheable(&self) -> bool {
self.signer.is_none()
}
}
impl MessagesRoute {
pub(super) async fn prepare_outbound(
&self,
request: ProviderMessagesRequest,
context: &CallContext<'_, impl Interceptors<Error>>,
) -> Result<ProviderCall, Error> {
let ProviderMessagesRequest {
provider,
url,
body,
environment,
timeout,
api_key,
} = request;
let request_context = RequestContext {
model: body.model.clone(),
custom_llm_provider: provider.as_str().to_string(),
optional_params: serde_json::to_value(&body.params).map_err(serialize_failure)?,
secret_fields: Vec::new(),
api_key,
};
let authenticated =
resolve_auth(&self.auth, environment, &|key| std::env::var(key).ok()).await?;
let identity = ProviderIdentity {
model: request_context.model.clone(),
provider: request_context.custom_llm_provider.clone(),
};
let wire = context
.interceptors
.before_provider_request(
WireRequest {
url,
headers: authenticated.headers,
body: serde_json::to_value(&body).map_err(serialize_failure)?,
},
&wire.url,
&wire.body,
timeout,
request_context,
)
.await?;
if !response.status().is_success() {
return Err(provider_error(response).await);
}
let config = provider.config();
if stream {
return Ok(streaming_response(
response,
config.stream_decoder(),
provider_name,
let stream = match wire.body.get("stream") {
None | Some(Value::Null) => false,
Some(Value::Bool(stream)) => *stream,
Some(value) => {
return Err(Error::InvalidRequest(
litellm_llms::ErrorDetail::InvalidValue {
field: "stream",
expected: "a boolean",
actual: value.clone(),
},
));
}
let text = response.text().await.map_err(network)?;
log_response_body(&text);
let raw = RawResponse { body: text.clone() };
if let Some(observers) = observers {
observers.emit(litellm_host::lifecycle::CallEvent::Execution(
ExecutionEvent::ProviderResponseReceived { raw: raw.clone() },
));
}
interceptors
.after_provider_response(raw)
.await
.map_err(Error::post_call)?;
decode_response(config, &body.model, &text)
.map(|message| MessagesCallResponse::Complete(Box::new(message)))
},
)
.await
};
Ok(ProviderCall {
identity,
wire,
provider,
signer: authenticated.signer,
timeout,
stream,
})
}
pub(super) async fn call_provider(
&self,
request: ProviderCall,
context: &CallContext<'_, impl Interceptors<Error>>,
) -> Result<MessagesCallResponse, Error> {
let ProviderCall {
identity,
wire,
provider,
signer,
timeout,
stream,
} = request;
let provider_name = provider.as_str();
log_request_body(provider_name, stream, &wire.body);
let response = send(
&self.http,
Authenticated {
headers: wire.headers,
signer,
},
&wire.url,
&wire.body,
timeout,
)
.await?;
if !response.status().is_success() {
return Err(provider_error(response).await);
}
let config = provider.config();
if stream {
return Ok(streaming_response(
response,
config.stream_decoder(),
provider_name,
));
}
let text = response.text().await.map_err(network)?;
log_response_body(&text);
context.response_received(&text).await?;
decode_response(config, &identity.model, &text)
.map(|message| MessagesCallResponse::Complete(Box::new(message)))
}
}
fn serialize_failure(err: serde_json::Error) -> Error {

View file

@ -1,11 +1,14 @@
use litellm_host::observation::ObservationSender;
mod common_utils;
mod handler;
mod prepare;
pub mod route;
mod types;
use futures_util::FutureExt;
use litellm_auth::AuthServices;
use litellm_host::interceptors::{ExecutionFacts, Interceptors, ResultSource};
use crate::{caching::CallCache, context::CallContext};
use litellm_secrets::source::SecretSource;
use std::sync::Arc;
@ -20,74 +23,18 @@ pub struct MessagesRoute {
cache: Option<litellm_cache_response::ScopedCache>,
}
#[must_use]
#[derive(Clone, Default)]
pub struct MessagesRouteBuilder<Http = (), Auth = (), Secrets = ()> {
http: Http,
auth: Auth,
secrets: Secrets,
cache: Option<litellm_cache_response::ScopedCache>,
}
impl<Http, Auth, Secrets> MessagesRouteBuilder<Http, Auth, Secrets> {
pub fn with_http(
self,
http: litellm_http::Client,
) -> MessagesRouteBuilder<litellm_http::Client, Auth, Secrets> {
MessagesRouteBuilder {
http,
auth: self.auth,
secrets: self.secrets,
cache: self.cache,
}
}
pub fn with_auth(
self,
auth: Arc<AuthServices>,
) -> MessagesRouteBuilder<Http, Arc<AuthServices>, Secrets> {
MessagesRouteBuilder {
http: self.http,
auth,
secrets: self.secrets,
cache: self.cache,
}
}
pub fn with_secrets(
self,
secrets: Arc<dyn SecretSource>,
) -> MessagesRouteBuilder<Http, Auth, Arc<dyn SecretSource>> {
MessagesRouteBuilder {
http: self.http,
auth: self.auth,
secrets,
cache: self.cache,
}
}
pub fn with_cache(self, cache: litellm_cache_response::ScopedCache) -> Self {
Self {
cache: Some(cache),
..self
}
}
}
impl MessagesRouteBuilder<litellm_http::Client, Arc<AuthServices>, Arc<dyn SecretSource>> {
pub fn build(self) -> MessagesRoute {
MessagesRoute {
http: self.http,
auth: self.auth,
secrets: self.secrets,
cache: self.cache,
}
}
}
impl MessagesRoute {
pub fn builder() -> MessagesRouteBuilder {
MessagesRouteBuilder::default()
pub fn new(
http: litellm_http::Client,
auth: Arc<AuthServices>,
secrets: Arc<dyn SecretSource>,
) -> Self {
Self {
http,
auth,
secrets,
cache: None,
}
}
#[must_use]
@ -104,15 +51,9 @@ impl MessagesRoute {
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
options: impl Into<crate::CallOptions>,
) -> Result<MessagesCallResponse, Error> {
let crate::CallOptions {
cache: cache_options,
observers,
} = options.into();
litellm_host::lifecycle::observe_call(
observers.clone(),
self.run(call, cache_options, interceptors, observers.as_ref()),
)
.await
let context = CallContext::new(interceptors, options.into());
litellm_host::lifecycle::observe_call(context.observers.clone(), self.run(call, context))
.await
}
#[tracing::instrument(name = "litellm.route", skip_all, fields(
@ -126,36 +67,34 @@ impl MessagesRoute {
async fn run(
&self,
call: MessagesCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
context: CallContext<'_, impl Interceptors<Error>>,
) -> Result<MessagesCallResponse, Error> {
crate::diagnostic::call(async {
self.run_provider(call, cache_options, interceptors, observers)
.await
let prepared = prepare::prepare(call, self.secrets.as_ref()).await?;
crate::diagnostic::provider(&prepared.body.model, prepared.provider.as_str());
let request = self.prepare_outbound(prepared, &context).boxed().await?;
let cache = CallCache::<route::Messages>::from_wire(
self.cache.as_ref().filter(|_| request.cacheable()),
context.cache,
&request.identity,
&request.wire,
);
let identity = request.identity.clone();
let (output, source) = match cache.lookup().await {
Some(hit) => hit,
None => (
self.call_provider(request, &context).await?,
ResultSource::Provider,
),
};
context
.result_ready(ExecutionFacts {
provider: identity,
source: source.clone(),
})
.await?;
Ok(cache.finish(output, &source).await)
})
.await
}
async fn run_provider(
&self,
call: MessagesCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<MessagesCallResponse, Error> {
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
crate::diagnostic::provider(&request.body.model, request.provider.as_str());
let execute: futures_util::future::BoxFuture<'_, Result<MessagesCallResponse, Error>> =
Box::pin(handler::execute(
&self.http,
&self.auth,
request,
self.cache.clone(),
cache_options,
interceptors,
observers,
));
execute.await
}
}

View file

@ -43,8 +43,14 @@ impl super::MessagesRoute {
request,
observers,
move |call, _, interceptors, observers| async move {
self.run(call, cache_options, &interceptors, observers.as_ref())
.await
let context = crate::context::CallContext::new(
&interceptors,
crate::CallOptions {
cache: cache_options,
observers,
},
);
self.run(call, context).await
},
)
}

View file

@ -16,7 +16,7 @@ pub(super) async fn execute(
auth: &litellm_auth::AuthServices,
request: ProviderResponsesRequest,
cache: Option<litellm_cache_response::ScopedCache>,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {

View file

@ -71,7 +71,7 @@ impl ResponsesRoute {
async fn run(
&self,
call: ResponsesCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {
@ -85,7 +85,7 @@ impl ResponsesRoute {
async fn run_provider(
&self,
call: ResponsesCall,
cache_options: Option<litellm_cache_response::CacheOptions>,
cache_options: Option<litellm_cache_response::CachePolicy>,
interceptors: &impl Interceptors<Error>,
observers: Option<&ObservationSender>,
) -> Result<ResponsesOutput, Error> {

View file

@ -12,8 +12,8 @@ use bytes::Bytes;
use futures_util::{StreamExt, TryStreamExt, stream};
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{
CacheOptions, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService,
ResponseEnvelope,
CacheOptions, CachePolicy, CacheScope, ResponseCache, ResponseCacheConfig,
ResponseCacheService, ResponseEnvelope,
};
use litellm_core::{
RouteError,
@ -117,9 +117,9 @@ async fn call(
#[rstest]
#[case::normal(CacheOptions::new(CacheScope::Shared), true, true)]
#[case::no_cache(CacheOptions { no_cache: true, ..CacheOptions::new(CacheScope::Shared) }, false, true)]
#[case::no_store(CacheOptions { no_store: true, ..CacheOptions::new(CacheScope::Shared) }, true, false)]
#[case::disabled(CacheOptions { caching: Some(false), ..CacheOptions::new(CacheScope::Shared) }, false, false)]
#[case::no_cache(CacheOptions { policy: CachePolicy { no_cache: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, true)]
#[case::no_store(CacheOptions { policy: CachePolicy { no_store: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, true, false)]
#[case::disabled(CacheOptions { policy: CachePolicy { caching: Some(false), ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, false)]
#[tokio::test]
async fn cache_controls_apply_to_both_reads_and_writes(
cache: Arc<dyn ResponseCacheService>,
@ -723,9 +723,9 @@ async fn unary_call(
#[rstest]
#[case::normal(CacheOptions::new(CacheScope::Shared), true, true)]
#[case::no_cache(CacheOptions { no_cache: true, ..CacheOptions::new(CacheScope::Shared) }, false, true)]
#[case::no_store(CacheOptions { no_store: true, ..CacheOptions::new(CacheScope::Shared) }, true, false)]
#[case::disabled(CacheOptions { caching: Some(false), ..CacheOptions::new(CacheScope::Shared) }, false, false)]
#[case::no_cache(CacheOptions { policy: CachePolicy { no_cache: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, true)]
#[case::no_store(CacheOptions { policy: CachePolicy { no_store: true, ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, true, false)]
#[case::disabled(CacheOptions { policy: CachePolicy { caching: Some(false), ..CachePolicy::default() }, ..CacheOptions::new(CacheScope::Shared) }, false, false)]
#[tokio::test]
async fn unary_cache_controls_do_not_change_the_shared_service(
cache: Arc<dyn ResponseCacheService>,

View file

@ -1,9 +1,9 @@
use litellm_host::lifecycle::ExecutionEvent;
use std::sync::Mutex;
use litellm_core::messages::route::Messages;
use litellm_core::messages::{MessagesCallResponse, route::Messages};
use litellm_host::{
interceptors::{RequestContext, WireRequest},
interceptors::{ExecutionFacts, RequestContext, ResultSource, WireRequest},
lifecycle::CallEvent,
};
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities as AnthropicModelCapabilities;
@ -20,6 +20,8 @@ struct RecordingHost {
rewrite: Rewrite,
events: super::support::Observations,
optional_params: Mutex<Vec<Value>>,
facts: Mutex<Vec<ExecutionFacts>>,
reject_result: bool,
}
impl RecordingHost {
@ -29,6 +31,8 @@ impl RecordingHost {
rewrite,
events: super::support::Observations::default(),
optional_params: Mutex::new(Vec::new()),
facts: Mutex::new(Vec::new()),
reject_result: false,
}
}
@ -73,6 +77,14 @@ impl litellm_host::lifecycle::CallObserver for RecordingHost {
impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protocol::Protocol>::Error>
for RecordingHost
{
async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), Error> {
self.facts.lock().unwrap().push(facts);
if self.reject_result {
return Err(Error::Unsupported("result rejected"));
}
Ok(())
}
async fn before_provider_request(
&self,
wire: WireRequest,
@ -98,6 +110,91 @@ impl litellm_host::interceptors::Interceptors<<Messages as litellm_host::protoco
}
}
#[rstest]
#[case::native_unary(false, false)]
#[case::native_stream(true, false)]
#[case::hosted_unary(false, true)]
#[case::hosted_stream(true, true)]
#[tokio::test]
async fn rejected_results_are_not_delivered_or_cached(
call: MessagesCall,
#[case] streaming: bool,
#[case] hosted: bool,
) {
use futures_util::TryStreamExt;
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{CacheScope, ResponseCache, ScopedCache};
let response = if streaming {
ResponseTemplate::new(200).set_body_raw(
"event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n",
"text/event-stream",
)
} else {
message_response()
};
let upstream = upstream([response.clone(), response]).await;
let route = messages_route(no_secrets()).with_cache(ScopedCache::new(
Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
Some(100),
Some(Duration::from_secs(60)),
)))),
CacheScope::Shared,
));
for (reject, expected_requests, cached) in [
(true, 1, false),
(false, 2, false),
(true, 2, true),
(false, 2, true),
] {
let request = authenticated(
with_fields(
MessagesCall {
body: call.body.clone(),
..super::call()
},
json!({"stream": streaming}),
),
upstream.uri(),
);
let host = RecordingHost {
reject_result: reject,
..RecordingHost::passthrough(request)
};
let result = if hosted {
litellm_host_native::in_process::run_hosted(
route.clone().machine(host.request().unwrap(), None),
host.runtime(),
)
.await
.map(|_| ())
} else {
match route.execute(host.request().unwrap(), &host, None).await {
Ok(MessagesCallResponse::Complete(_)) => Ok(()),
Ok(MessagesCallResponse::Stream { chunks, .. }) => {
chunks.try_collect::<Vec<_>>().await.map(|_| ())
}
Err(error) => Err(error),
}
};
assert_eq!(
result,
if reject {
Err(Error::Unsupported("result rejected"))
} else {
Ok(())
}
);
assert_eq!(received(&upstream).await.len(), expected_requests);
let facts = host.facts.lock().unwrap();
assert_eq!(facts.len(), 1);
assert_eq!(
matches!(facts[0].source, ResultSource::Cache { .. }),
cached
);
}
}
async fn run_through(host: &RecordingHost) -> Result<MessagesOutput, Error> {
litellm_host_native::in_process::run_hosted(
machine(Arc::new(RecordingSecrets::empty()))(host.request()?),
@ -143,6 +240,81 @@ async fn what_before_send_returns_is_what_the_provider_receives(call: MessagesCa
assert_eq!(request.header("x-api-key"), Some("sk-ant"));
}
#[rstest]
#[case::enable(false, json!(true), Some(true))]
#[case::disable(true, json!(false), Some(false))]
#[case::null(true, Value::Null, Some(false))]
#[case::invalid(false, json!("true"), None)]
#[tokio::test]
async fn response_mode_follows_the_intercepted_request(
call: MessagesCall,
traces: TraceCapture,
#[case] original_stream: bool,
#[case] rewritten_stream: Value,
#[case] expected_stream: Option<bool>,
) {
use futures_util::TryStreamExt;
let sse = "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n";
let response = if expected_stream == Some(true) {
ResponseTemplate::new(200).set_body_raw(sse, "text/event-stream")
} else {
message_response()
};
let upstream = upstream([response]).await;
let rewrite = rewritten_stream.clone();
let host = RecordingHost::new(
authenticated(
with_fields(call, json!({"stream": original_stream})),
upstream.uri(),
),
Box::new(move |wire| {
let mut body = wire.body;
body["stream"] = rewrite.clone();
Ok(WireRequest { body, ..wire })
}),
);
let result = traces
.logger()
.instrument(async {
let output = messages_route(no_secrets())
.execute(host.request()?, &host, None)
.await?;
match output {
MessagesCallResponse::Stream { chunks, .. } => {
assert_eq!(expected_stream, Some(true));
assert_eq!(
chunks.try_collect::<Vec<_>>().await?.concat(),
sse.as_bytes()
);
}
MessagesCallResponse::Complete(message) => {
assert_eq!(expected_stream, Some(false));
assert_eq!(*message, serde_json::from_value(message_body()).unwrap());
}
}
Ok::<_, Error>(())
})
.await;
let summaries = traces.summaries("litellm.route");
assert_eq!(summaries.len(), 1);
let Some(expected_stream) = expected_stream else {
assert!(matches!(result, Err(Error::InvalidRequest(_))));
assert!(received(&upstream).await.is_empty());
assert_eq!(summaries[0]["outcome"], "failure");
return;
};
result.unwrap();
assert_eq!(
only_request(&upstream).await.json()["stream"],
rewritten_stream
);
assert_eq!(host.raw_responses().len(), usize::from(!expected_stream));
assert_eq!(summaries[0]["stream"], expected_stream);
assert_eq!(summaries[0]["outcome"], "success");
}
#[rstest]
#[tokio::test]
async fn a_before_send_failure_never_sends(call: MessagesCall) {

View file

@ -269,25 +269,22 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes
};
let resources = support::resources();
let response = litellm_core::messages::MessagesRoute::builder()
.with_http(provider_http(
&resources,
&Resolution::from(&settings).config,
))
.with_auth(resources.auth)
.with_secrets(no_secrets())
.build()
.execute(
MessagesCall {
api_key: Some("sk-ant".into()),
api_base: Some(base),
..call
},
&(),
None,
)
.await
.expect("messages request succeeds");
let response = litellm_core::messages::MessagesRoute::new(
provider_http(&resources, &Resolution::from(&settings).config),
resources.auth,
no_secrets(),
)
.execute(
MessagesCall {
api_key: Some("sk-ant".into()),
api_base: Some(base),
..call
},
&(),
None,
)
.await
.expect("messages request succeeds");
let MessagesCallResponse::Complete(message) = response else {
panic!("a non-streaming request returns a message");
@ -345,7 +342,7 @@ async fn message_route_summary_excludes_payload_diagnostics(
#[case::uncached(false, 2)]
#[case::cached(true, 1)]
#[tokio::test]
async fn builder_preserves_dependencies_and_optional_cache(
async fn route_uses_injected_dependencies_and_optional_cache(
#[case] caching: bool,
#[case] expected_requests: usize,
) {
@ -355,9 +352,13 @@ async fn builder_preserves_dependencies_and_optional_cache(
let upstream = upstream([message_response(), message_response()]).await;
let resources = resources();
let builder = MessagesRoute::builder();
let builder = if caching {
builder.with_cache(ScopedCache::new(
let route = MessagesRoute::new(
provider_http(&resources, &http_config()),
resources.auth.clone(),
Arc::new(RecordingSecrets::new([("ANTHROPIC_API_KEY", "route-key")])),
);
let route = if caching {
route.with_cache(ScopedCache::new(
Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
Some(100),
Some(Duration::from_secs(60)),
@ -365,16 +366,8 @@ async fn builder_preserves_dependencies_and_optional_cache(
CacheScope::Shared,
))
} else {
builder
route
};
let route = builder
.with_secrets(Arc::new(RecordingSecrets::new([(
"ANTHROPIC_API_KEY",
"builder-key",
)])))
.with_auth(resources.auth.clone())
.with_http(provider_http(&resources, &http_config()))
.build();
for _ in 0..2 {
let request = MessagesCall {
api_base: Some(upstream.uri()),
@ -392,5 +385,72 @@ async fn builder_preserves_dependencies_and_optional_cache(
}
let requests = received(&upstream).await;
assert_eq!(requests.len(), expected_requests);
assert_eq!(requests[0].header("x-api-key"), Some("builder-key"));
assert_eq!(requests[0].header("x-api-key"), Some("route-key"));
}
#[rstest]
#[tokio::test]
async fn cache_overrides_preserve_the_routes_isolated_scope(call: MessagesCall) {
use litellm_cache_memory::InMemoryCache;
use litellm_cache_response::{CachePolicy, CacheScope, ResponseCache, ScopedCache};
let first_body = message_body();
let second_body = Value::Object(
first_body
.as_object()
.unwrap()
.iter()
.map(|(key, value)| {
(
key.clone(),
if key == "id" {
json!("msg_second")
} else {
value.clone()
},
)
})
.collect(),
);
let upstream = upstream([
json_response(first_body.clone()),
json_response(second_body.clone()),
])
.await;
let service = Arc::new(ResponseCache::new(Arc::new(InMemoryCache::new(
Some(100),
Some(Duration::from_secs(60)),
))));
let first = messages_route(no_secrets()).with_cache(ScopedCache::new(
service.clone(),
CacheScope::Isolated("first".into()),
));
let second = messages_route(no_secrets()).with_cache(ScopedCache::new(
service,
CacheScope::Isolated("second".into()),
));
for (route, expected) in [
(&first, &first_body),
(&second, &second_body),
(&first, &first_body),
(&second, &second_body),
] {
let request = MessagesCall {
body: call.body.clone(),
api_key: Some("same-key".into()),
api_base: Some(upstream.uri()),
..super::call()
};
let override_options = CachePolicy {
ttl: Some(Duration::from_secs(30)),
..CachePolicy::default()
};
let MessagesCallResponse::Complete(response) =
route.execute(request, &(), override_options).await.unwrap()
else {
panic!("expected a completed message");
};
assert_eq!(response.id, expected["id"].as_str().unwrap());
}
assert_eq!(received(&upstream).await.len(), 2);
}

View file

@ -44,11 +44,11 @@ pub fn provider_http(
pub fn messages_route(secrets: Arc<dyn SecretSource>) -> litellm_core::messages::MessagesRoute {
let resources = resources();
litellm_core::messages::MessagesRoute::builder()
.with_http(provider_http(&resources, &http_config()))
.with_auth(resources.auth)
.with_secrets(secrets)
.build()
litellm_core::messages::MessagesRoute::new(
provider_http(&resources, &http_config()),
resources.auth,
secrets,
)
}
pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletionsRoute {

View file

@ -1,6 +1,6 @@
use std::time::Duration;
use litellm_cache_response::{CacheOptions, CacheScope};
use litellm_cache_response::{CacheOptions, CachePolicy, CacheScope};
use litellm_gateway_auth::AuthenticatedRequest;
use serde::Deserialize;
use serde_json::{Map, Value};
@ -38,11 +38,13 @@ pub(crate) fn prepare(
.map_err(|error| Error::InvalidBody(error.to_string()))?;
let caller = identity.caller();
let options = CacheOptions {
caching,
no_cache: controls.no_cache,
no_store: controls.no_store,
ttl: controls.ttl.map(duration).transpose()?,
max_age: controls.max_age.map(duration).transpose()?,
policy: CachePolicy {
caching,
no_cache: controls.no_cache,
no_store: controls.no_store,
ttl: controls.ttl.map(duration).transpose()?,
max_age: controls.max_age.map(duration).transpose()?,
},
scope: CacheScope::Isolated(
serde_json::json!([
caller.principal().authority(),

View file

@ -69,7 +69,7 @@ async fn handle(
extra_headers: None,
timeout: deployment.timeout,
},
cache_options,
cache_options.policy,
),
(),
headers.clone(),

View file

@ -68,11 +68,7 @@ impl Gateway {
auth.clone(),
secrets.clone(),
),
messages: MessagesRoute::builder()
.with_http(provider.clone())
.with_auth(auth.clone())
.with_secrets(secrets.clone())
.build(),
messages: MessagesRoute::new(provider.clone(), auth.clone(), secrets.clone()),
responses: ResponsesRoute::new(provider, auth.clone(), secrets.clone()),
ocr: OcrRoute::new(OcrClient::new(
&resources.pool,

View file

@ -54,7 +54,7 @@ async fn handle(
};
let call = project(deployment, body, headers)?;
let machine = route.machine(call, cache_options);
let machine = route.machine(call, cache_options.policy);
let stream =
Sse::<Messages, _, _>::new(Json, |error| Bytes::from(Error::from(error).sse_frame()));
let headers = crate::caching::CacheHeaders::default();

View file

@ -38,7 +38,7 @@ pub(crate) async fn create(
extra_headers: None,
timeout: deployment.timeout,
};
let machine = route.machine(call, cache_options);
let machine = route.machine(call, cache_options.policy);
let stream = Sse::<Responses, _, _>::new(Json, |error| {
let error = Error::from(error);
Bytes::from(format!(

View file

@ -330,15 +330,17 @@ pub(in crate::cache) fn configured(
Ok((
Some(cache.service.clone()),
litellm_cache_response::CacheOptions {
caching: kwargs
.get_item("caching")?
.filter(|value| !value.is_none())
.map(|value| value.extract())
.transpose()?,
no_cache: boolean("no-cache")?,
no_store: boolean("no-store")?,
ttl: seconds("ttl")?,
max_age: seconds("s-max-age")?.or(seconds("s-maxage")?),
policy: litellm_cache_response::CachePolicy {
caching: kwargs
.get_item("caching")?
.filter(|value| !value.is_none())
.map(|value| value.extract())
.transpose()?,
no_cache: boolean("no-cache")?,
no_store: boolean("no-store")?,
ttl: seconds("ttl")?,
max_age: seconds("s-max-age")?.or(seconds("s-maxage")?),
},
scope: litellm_cache_response::CacheScope::Shared,
},
))

View file

@ -1,5 +1,7 @@
use super::{native, python};
use litellm_cache_response::{CacheOptions, CacheScope, ResponseCacheService, ScopedCache};
use litellm_cache_response::{
CacheOptions, CachePolicy, CacheScope, ResponseCacheService, ScopedCache,
};
use litellm_host::{
machine::{HostServices, MachineFault},
protocol::Protocol,
@ -154,8 +156,11 @@ pub(crate) fn configure(
.map(|value| value.unwrap_or(false))
};
let options = CacheOptions {
no_cache: boolean("no-cache")?,
no_store: boolean("no-store")?,
policy: CachePolicy {
no_cache: boolean("no-cache")?,
no_store: boolean("no-store")?,
..CachePolicy::default()
},
..CacheOptions::new(CacheScope::Shared)
};
let namespace = cache

View file

@ -175,7 +175,7 @@ fn run_public(
)),
None => route,
};
Ok(route.machine(request, cache_options))
Ok(route.machine(request, cache_options.policy))
},
host::ChatCompletionsPythonHost(host),
hooks,

View file

@ -26,14 +26,12 @@ fn run_messages(
py,
arguments,
move |py, arguments, request| {
let builder = litellm_core::messages::MessagesRoute::builder()
.with_http(
crate::http::provider_client(py, arguments, asynchronous)?
.map_err(crate::http::client_error)?,
)
.with_auth(crate::http::resources().auth.clone())
.with_secrets(crate::secrets::source(py)?);
let route = builder.build();
let route = litellm_core::messages::MessagesRoute::new(
crate::http::provider_client(py, arguments, asynchronous)?
.map_err(crate::http::client_error)?,
crate::http::resources().auth.clone(),
crate::secrets::source(py)?,
);
Ok(litellm_host::call::hosted_call(
request,
None,
@ -51,7 +49,7 @@ fn run_messages(
call,
&interceptors,
litellm_core::CallOptions {
cache: Some(options),
cache: Some(options.policy),
observers,
},
)

View file

@ -96,7 +96,7 @@ fn run_public(
)),
None => route,
};
Ok(route.machine(request, cache_options))
Ok(route.machine(request, cache_options.policy))
},
host::ResponsesPythonHost(host),
hooks,