mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
d7f5b40ab4
commit
6a55e0a8aa
24 changed files with 726 additions and 373 deletions
26
litellm-rust/Cargo.lock
generated
26
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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" }
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
29
litellm-rust/crates/inference-chat/Cargo.toml
Normal file
29
litellm-rust/crates/inference-chat/Cargo.toml
Normal 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
|
||||
6
litellm-rust/crates/inference-chat/src/constants.rs
Normal file
6
litellm-rust/crates/inference-chat/src/constants.rs
Normal 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";
|
||||
|
|
@ -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,
|
||||
};
|
||||
|
|
@ -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,
|
||||
|
|
@ -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,
|
||||
|
|
@ -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";
|
||||
}
|
||||
337
litellm-rust/crates/inference-chat/tests/caching.rs
Normal file
337
litellm-rust/crates/inference-chat/tests/caching.rs
Normal 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;
|
||||
}
|
||||
|
|
@ -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();
|
||||
281
litellm-rust/crates/inference-chat/tests/support/mod.rs
Normal file
281
litellm-rust/crates/inference-chat/tests/support/mod.rs
Normal 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()
|
||||
}
|
||||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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::*,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue