mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
refactor(rust): extract inference-responses crate (#44811)
Co-authored-by: Yujong Lee <yujong@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
63babf23e6
commit
785cb12f37
22 changed files with 718 additions and 161 deletions
29
litellm-rust/Cargo.lock
generated
29
litellm-rust/Cargo.lock
generated
|
|
@ -3836,6 +3836,7 @@ dependencies = [
|
|||
"litellm-host-http",
|
||||
"litellm-http",
|
||||
"litellm-inference",
|
||||
"litellm-inference-responses",
|
||||
"litellm-inference-transcription",
|
||||
"litellm-llms",
|
||||
"litellm-llms-types",
|
||||
|
|
@ -4040,6 +4041,33 @@ dependencies = [
|
|||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-inference-responses"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"litellm-auth",
|
||||
"litellm-cache",
|
||||
"litellm-cache-memory",
|
||||
"litellm-cache-response",
|
||||
"litellm-host",
|
||||
"litellm-host-native",
|
||||
"litellm-http",
|
||||
"litellm-inference",
|
||||
"litellm-llms",
|
||||
"litellm-llms-types",
|
||||
"litellm-secrets",
|
||||
"litellm-tracing",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"tracing",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-inference-transcription"
|
||||
version = "0.1.0"
|
||||
|
|
@ -4142,6 +4170,7 @@ dependencies = [
|
|||
"litellm-host-python",
|
||||
"litellm-http",
|
||||
"litellm-inference",
|
||||
"litellm-inference-responses",
|
||||
"litellm-inference-transcription",
|
||||
"litellm-llms",
|
||||
"litellm-llms-types",
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ litellm-traces-clickhouse = { path = "crates/traces-clickhouse" }
|
|||
litellm-storage-clickhouse = { path = "crates/storage-clickhouse" }
|
||||
litellm-inference = { path = "crates/inference" }
|
||||
litellm-inference-transcription = { path = "crates/inference-transcription" }
|
||||
litellm-inference-responses = { path = "crates/inference-responses" }
|
||||
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
|
||||
litellm-gateway = { path = "crates/gateway" }
|
||||
litellm-gateway-inference = { path = "crates/gateway-inference" }
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ litellm-auth.workspace = true
|
|||
litellm-gateway-auth.workspace = true
|
||||
litellm-inference.workspace = true
|
||||
litellm-inference-transcription.workspace = true
|
||||
litellm-inference-responses.workspace = true
|
||||
litellm-host-http.workspace = true
|
||||
litellm-host.workspace = true
|
||||
litellm-http.workspace = true
|
||||
|
|
|
|||
|
|
@ -18,8 +18,9 @@ use axum::{Router, routing::post};
|
|||
use litellm_http::{ClientVariant, HttpClientConfig, media::UrlPolicy};
|
||||
use litellm_inference::{
|
||||
chat_completions::ChatCompletionsRoute, messages::MessagesRoute, ocr::OcrRoute,
|
||||
resources::CoreResources, responses::ResponsesRoute,
|
||||
resources::CoreResources,
|
||||
};
|
||||
use litellm_inference_responses::ResponsesRoute;
|
||||
use litellm_inference_transcription::AudioTranscriptionRoute;
|
||||
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ use std::sync::Arc;
|
|||
use axum::{Json, body::Bytes, extract::State, response::Response};
|
||||
use litellm_gateway_auth::AuthenticatedRequest;
|
||||
use litellm_host_http::Sse;
|
||||
use litellm_inference::responses::{route::Responses, types::ResponsesCall};
|
||||
use litellm_inference_responses::{route::Responses, types::ResponsesCall};
|
||||
use serde_json::json;
|
||||
|
||||
use crate::{Error, Gateway, JsonObject, request};
|
||||
|
|
|
|||
36
litellm-rust/crates/inference-responses/Cargo.toml
Normal file
36
litellm-rust/crates/inference-responses/Cargo.toml
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
[package]
|
||||
name = "litellm-inference-responses"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
bytes.workspace = true
|
||||
futures-util.workspace = true
|
||||
litellm-auth = { workspace = true, features = ["aws", "azure", "gcp"] }
|
||||
litellm-cache-response.workspace = true
|
||||
litellm-host.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-inference.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-llms-types.workspace = true
|
||||
litellm-secrets.workspace = true
|
||||
reqwest.workspace = true
|
||||
serde_json = { workspace = true, features = ["preserve_order"] }
|
||||
tokio = { workspace = true, features = ["sync"] }
|
||||
tokio-tungstenite.workspace = true
|
||||
tracing.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
futures-util.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
litellm-cache-memory.workspace = true
|
||||
litellm-cache-response.workspace = true
|
||||
litellm-host-native.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-inference = { workspace = true, features = ["test-support"] }
|
||||
litellm-tracing.workspace = true
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock.workspace = true
|
||||
|
|
@ -36,9 +36,11 @@ pub(super) async fn execute(
|
|||
)
|
||||
.await?;
|
||||
let cache = cache.filter(|_| authenticated.signer.is_none());
|
||||
let cache_request =
|
||||
crate::caching::CacheRequest::from_wire(identity, cache.as_ref().map(|_| &wire));
|
||||
crate::caching::execute_streaming::<super::route::Responses, _, _>(
|
||||
let cache_request = litellm_inference::caching::CacheRequest::from_wire(
|
||||
identity,
|
||||
cache.as_ref().map(|_| &wire),
|
||||
);
|
||||
litellm_inference::caching::execute_streaming::<super::route::Responses, _, _>(
|
||||
cache_request,
|
||||
cache.as_ref().map(|cache| cache.service.clone()),
|
||||
cache.as_ref().map(|cache| cache.options(cache_options)),
|
||||
|
|
@ -50,7 +52,7 @@ pub(super) async fn execute(
|
|||
Some(serde_json::Value::Bool(value)) => *value,
|
||||
Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())),
|
||||
};
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
let outbound = litellm_inference::outbound::outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
|
|
@ -59,7 +61,7 @@ pub(super) async fn execute(
|
|||
&wire.body,
|
||||
Some(request.timeout.unwrap_or(Duration::from_secs(600))),
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
let response = litellm_inference::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(network)?;
|
||||
let status = response.status().as_u16();
|
||||
|
|
@ -1,11 +1,12 @@
|
|||
pub use crate::error::RouteError as Error;
|
||||
use litellm_host::observation::ObservationSender;
|
||||
pub use litellm_inference::RouteError as Error;
|
||||
|
||||
pub mod route;
|
||||
pub mod types;
|
||||
pub mod websocket;
|
||||
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
pub mod types;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
|
|
@ -47,9 +48,9 @@ impl ResponsesRoute {
|
|||
&self,
|
||||
call: ResponsesCall,
|
||||
interceptors: &impl Interceptors<Error>,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
options: impl Into<litellm_inference::CallOptions>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let crate::CallOptions {
|
||||
let litellm_inference::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
|
|
@ -75,7 +76,7 @@ impl ResponsesRoute {
|
|||
interceptors: &impl litellm_host::interceptors::Interceptors<Error>,
|
||||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
crate::diagnostic::call(async {
|
||||
litellm_inference::diagnostic::call(async {
|
||||
self.run_provider(call, cache_options, interceptors, observers)
|
||||
.await
|
||||
})
|
||||
|
|
@ -90,7 +91,10 @@ impl ResponsesRoute {
|
|||
observers: Option<&ObservationSender>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.context.model, &request.context.custom_llm_provider);
|
||||
litellm_inference::diagnostic::provider(
|
||||
&request.context.model,
|
||||
&request.context.custom_llm_provider,
|
||||
);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<ResponsesOutput, Error>> =
|
||||
Box::pin(handler::execute(
|
||||
&self.http,
|
||||
|
|
@ -27,9 +27,9 @@ impl ResponsesRoute {
|
|||
pub fn machine(
|
||||
self,
|
||||
call: ResponsesCall,
|
||||
options: impl Into<crate::CallOptions>,
|
||||
options: impl Into<litellm_inference::CallOptions>,
|
||||
) -> HostedMachine<Responses> {
|
||||
let crate::CallOptions {
|
||||
let litellm_inference::CallOptions {
|
||||
cache: cache_options,
|
||||
observers,
|
||||
} = options.into();
|
||||
|
|
@ -44,7 +44,7 @@ impl ResponsesRoute {
|
|||
}
|
||||
}
|
||||
|
||||
impl crate::caching::Cachable for Responses {
|
||||
impl litellm_inference::caching::Cachable for Responses {
|
||||
const SURFACE: &'static str = "responses";
|
||||
|
||||
fn reusable(response: &Self::Response) -> bool {
|
||||
|
|
@ -56,7 +56,7 @@ impl crate::caching::Cachable for Responses {
|
|||
}
|
||||
}
|
||||
|
||||
impl crate::caching::StreamCachable for Responses {
|
||||
impl litellm_inference::caching::StreamCachable for Responses {
|
||||
const TERMINAL_EVENT: &'static str = "response.completed";
|
||||
|
||||
fn replay(data: bytes::Bytes) -> Option<litellm_host::call::OutputOf<Self>> {
|
||||
|
|
@ -40,7 +40,7 @@ impl ResponsesWebSocketConnection {
|
|||
headers: &HashMap<String, String>,
|
||||
timeout: Option<Duration>,
|
||||
) -> Result<Self, Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.connect_url", async {
|
||||
litellm_inference::diagnostic::operation("litellm.websocket.connect_url", async {
|
||||
let mut request = url.into_client_request().map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
})?;
|
||||
|
|
@ -86,7 +86,7 @@ impl ResponsesWebSocketConnection {
|
|||
fields(outcome)
|
||||
)]
|
||||
pub async fn send_text(&self, text: String) -> Result<(), Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.send_text", async {
|
||||
litellm_inference::diagnostic::operation("litellm.websocket.send_text", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Err(Error::Transport(litellm_http::transport::Error::Network(
|
||||
|
|
@ -107,7 +107,7 @@ impl ResponsesWebSocketConnection {
|
|||
fields(outcome)
|
||||
)]
|
||||
pub async fn recv_text(&self) -> Result<Option<String>, Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.recv_text", async {
|
||||
litellm_inference::diagnostic::operation("litellm.websocket.recv_text", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
let Some(socket) = socket.as_mut() else {
|
||||
return Ok(None);
|
||||
|
|
@ -134,7 +134,7 @@ impl ResponsesWebSocketConnection {
|
|||
fields(outcome)
|
||||
)]
|
||||
pub async fn close(&self) -> Result<(), Error> {
|
||||
crate::diagnostic::operation("litellm.websocket.close", async {
|
||||
litellm_inference::diagnostic::operation("litellm.websocket.close", async {
|
||||
let mut socket = self.socket.lock().await;
|
||||
if let Some(socket) = socket.as_mut() {
|
||||
socket.close(None).await.map_err(|error| {
|
||||
327
litellm-rust/crates/inference-responses/tests/caching.rs
Normal file
327
litellm-rust/crates/inference-responses/tests/caching.rs
Normal file
|
|
@ -0,0 +1,327 @@
|
|||
mod support;
|
||||
|
||||
use std::{
|
||||
sync::{
|
||||
Arc,
|
||||
atomic::{AtomicUsize, Ordering},
|
||||
},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{
|
||||
CacheOptions, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService,
|
||||
ResponseEnvelope, ScopedCache,
|
||||
};
|
||||
use litellm_host::interceptors::{
|
||||
ExecutionFacts, Interceptors, ProviderIdentity, RawResponse, RequestContext, ResultSource,
|
||||
WireRequest,
|
||||
};
|
||||
use litellm_inference::{
|
||||
RouteError,
|
||||
caching::{CacheRequest, execute_unary},
|
||||
};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Value, json};
|
||||
|
||||
fn cache_request(input: Value) -> CacheRequest {
|
||||
CacheRequest {
|
||||
identity: ProviderIdentity {
|
||||
model: "test-model".into(),
|
||||
provider: "test-provider".into(),
|
||||
},
|
||||
input,
|
||||
}
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn cache() -> Arc<dyn ResponseCacheService> {
|
||||
cache_with_limit(4096)
|
||||
}
|
||||
|
||||
fn cache_with_limit(max_entry_bytes: usize) -> Arc<dyn ResponseCacheService> {
|
||||
Arc::new(
|
||||
ResponseCache::new(Arc::new(InMemoryCache::new(
|
||||
Some(100),
|
||||
Some(Duration::from_secs(60)),
|
||||
)))
|
||||
.with_config(ResponseCacheConfig {
|
||||
namespace: "test".into(),
|
||||
max_entry_bytes,
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
struct InvalidEntryCache(
|
||||
ResponseCache<InMemoryCache<litellm_cache_response::CacheEntry>>,
|
||||
Value,
|
||||
);
|
||||
|
||||
impl ResponseCacheService for InvalidEntryCache {
|
||||
fn config(&self) -> &ResponseCacheConfig {
|
||||
self.0.config()
|
||||
}
|
||||
|
||||
fn lookup<'a>(
|
||||
&'a self,
|
||||
request: &'a litellm_cache_response::ResponseCacheRequest,
|
||||
now: Duration,
|
||||
) -> futures_util::future::BoxFuture<'a, Result<Option<Value>, litellm_cache::Error>> {
|
||||
Box::pin(async move {
|
||||
Ok(self
|
||||
.0
|
||||
.async_lookup(request, now)
|
||||
.await?
|
||||
.or_else(|| Some(self.1.clone())))
|
||||
})
|
||||
}
|
||||
|
||||
fn store<'a>(
|
||||
&'a self,
|
||||
request: &'a litellm_cache_response::ResponseCacheRequest,
|
||||
response: Value,
|
||||
now: Duration,
|
||||
) -> futures_util::future::BoxFuture<'a, Result<(), litellm_cache::Error>> {
|
||||
Box::pin(self.0.async_store(request, response, now))
|
||||
}
|
||||
}
|
||||
|
||||
struct ChangingSecrets {
|
||||
revision: AtomicUsize,
|
||||
endpoints: [String; 2],
|
||||
change_credentials: bool,
|
||||
}
|
||||
|
||||
impl litellm_secrets::source::SecretSource for ChangingSecrets {
|
||||
fn get_secret_str<'a>(
|
||||
&'a self,
|
||||
name: &'a str,
|
||||
) -> futures_util::future::BoxFuture<
|
||||
'a,
|
||||
Result<Option<litellm_secrets::SecretValue>, litellm_secrets::Error>,
|
||||
> {
|
||||
Box::pin(async move {
|
||||
let revision = self.revision.load(Ordering::SeqCst);
|
||||
let value = if name.ends_with("_API_KEY") {
|
||||
Some(format!(
|
||||
"key-{}",
|
||||
if self.change_credentials { revision } else { 0 }
|
||||
))
|
||||
} else if name.ends_with("_API_BASE") {
|
||||
Some(self.endpoints[revision].clone())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
Ok(value.map(litellm_secrets::SecretValue::new))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct ChangingHooks {
|
||||
calls: AtomicUsize,
|
||||
rewrite: bool,
|
||||
facts: std::sync::Mutex<Vec<ExecutionFacts>>,
|
||||
}
|
||||
|
||||
impl Interceptors<RouteError> for ChangingHooks {
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
mut wire: WireRequest,
|
||||
_: RequestContext,
|
||||
) -> Result<WireRequest, RouteError> {
|
||||
let call = self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
if self.rewrite {
|
||||
wire.body["temperature"] = json!(if call < 2 { 0.1 } else { 0.8 });
|
||||
}
|
||||
Ok(wire)
|
||||
}
|
||||
|
||||
async fn after_provider_response(&self, _: RawResponse) -> Result<(), RouteError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn result_ready(&self, facts: ExecutionFacts) -> Result<(), RouteError> {
|
||||
self.facts.lock().unwrap().push(facts);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::chat_completion(json!({"kind":"Response","value":{"id":"chat-1","model":"test","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2}}}))]
|
||||
#[case::wrong_envelope(json!({"kind":"Stream","value":"data: [DONE]\n\n"}))]
|
||||
#[tokio::test]
|
||||
async fn responses_refetches_instead_of_deserializing_another_api_response(
|
||||
#[case] poisoned: Value,
|
||||
) {
|
||||
use litellm_inference_responses::route::Responses;
|
||||
use litellm_llms_types::formats::responses::ResponsesApiResponse;
|
||||
|
||||
let cache: Arc<dyn ResponseCacheService> = Arc::new(InvalidEntryCache(
|
||||
ResponseCache::new(Arc::new(InMemoryCache::default())),
|
||||
serde_json::to_value(ResponseEnvelope::new("responses", poisoned)).unwrap(),
|
||||
));
|
||||
let calls = AtomicUsize::new(0);
|
||||
for _ in 0..2 {
|
||||
let response = execute_unary::<Responses, _, _>(
|
||||
cache_request(json!({"input":"hello"})),
|
||||
Some(cache.clone()),
|
||||
Some(CacheOptions::new(CacheScope::Shared)),
|
||||
&(),
|
||||
None,
|
||||
|| async {
|
||||
calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(ResponsesApiResponse {
|
||||
id: "fresh-response".into(),
|
||||
model: "test".into(),
|
||||
output: vec![
|
||||
json!({"type":"message","content":[{"type":"output_text","text":"fresh"}]}),
|
||||
],
|
||||
extra: [("status".into(), json!("completed"))]
|
||||
.into_iter()
|
||||
.collect(),
|
||||
})
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.id, "fresh-response");
|
||||
assert_eq!(response.output[0]["content"][0]["text"], "fresh");
|
||||
}
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::completed("completed", 1)]
|
||||
#[case::incomplete("incomplete", 2)]
|
||||
#[tokio::test]
|
||||
async fn responses_cache_only_reuses_completed_responses(
|
||||
cache: Arc<dyn ResponseCacheService>,
|
||||
#[case] status: &str,
|
||||
#[case] expected_calls: usize,
|
||||
) {
|
||||
use litellm_inference_responses::route::Responses;
|
||||
use litellm_llms_types::formats::responses::ResponsesApiResponse;
|
||||
|
||||
let calls = AtomicUsize::new(0);
|
||||
for _ in 0..2 {
|
||||
let response = execute_unary::<Responses, _, _>(
|
||||
cache_request(json!({"input":"hello"})),
|
||||
Some(cache.clone()),
|
||||
Some(CacheOptions::new(CacheScope::Shared)),
|
||||
&(),
|
||||
None,
|
||||
|| async {
|
||||
let call = calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(ResponsesApiResponse {
|
||||
id: call.to_string(),
|
||||
model: "test".into(),
|
||||
output: Vec::new(),
|
||||
extra: [("status".into(), json!(status))].into_iter().collect(),
|
||||
})
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.extra.get("status"), Some(&json!(status)));
|
||||
}
|
||||
assert_eq!(calls.load(Ordering::SeqCst), expected_calls);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::credentials("credentials")]
|
||||
#[case::endpoint("endpoint")]
|
||||
#[case::callback("callback")]
|
||||
#[tokio::test]
|
||||
async fn responses_cache_identity_follows_resolved_configuration_and_request_callbacks(
|
||||
cache: Arc<dyn ResponseCacheService>,
|
||||
#[case] change: &str,
|
||||
) {
|
||||
use litellm_inference_responses::types::ResponsesCall;
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
|
||||
|
||||
let first = MockServer::start().await;
|
||||
let second = MockServer::start().await;
|
||||
let response = json!({"id":"response-test", "model":"test", "output":[], "status":"completed"});
|
||||
Mock::given(method("POST"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(response.clone()))
|
||||
.expect(if change == "endpoint" { 1 } else { 2 })
|
||||
.mount(&first)
|
||||
.await;
|
||||
Mock::given(method("POST"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(response))
|
||||
.expect(if change == "endpoint" { 1 } else { 0 })
|
||||
.mount(&second)
|
||||
.await;
|
||||
let secrets = Arc::new(ChangingSecrets {
|
||||
revision: AtomicUsize::new(0),
|
||||
endpoints: [
|
||||
first.uri(),
|
||||
if change == "endpoint" {
|
||||
second.uri()
|
||||
} else {
|
||||
first.uri()
|
||||
},
|
||||
],
|
||||
change_credentials: change == "credentials",
|
||||
});
|
||||
let hooks = ChangingHooks {
|
||||
rewrite: change == "callback",
|
||||
..Default::default()
|
||||
};
|
||||
for call in 0..4 {
|
||||
secrets
|
||||
.revision
|
||||
.store(usize::from(call >= 2), Ordering::SeqCst);
|
||||
let cache = ScopedCache::new(cache.clone(), CacheScope::Shared);
|
||||
support::responses_route(secrets.clone())
|
||||
.with_cache(cache)
|
||||
.execute(
|
||||
ResponsesCall {
|
||||
model: "test".into(),
|
||||
input: json!("hello"),
|
||||
optional_params: Default::default(),
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
},
|
||||
&hooks,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
assert_eq!(hooks.calls.load(Ordering::SeqCst), 4);
|
||||
{
|
||||
let facts = hooks.facts.lock().unwrap();
|
||||
assert_eq!(facts[0].source, ResultSource::Provider);
|
||||
assert_eq!(facts[2].source, ResultSource::Provider);
|
||||
let (ResultSource::Cache { key: first_key }, ResultSource::Cache { key: second_key }) =
|
||||
(&facts[1].source, &facts[3].source)
|
||||
else {
|
||||
panic!("unchanged effective requests must hit the cache");
|
||||
};
|
||||
assert_ne!(first_key, second_key);
|
||||
}
|
||||
let requests = first.received_requests().await.unwrap();
|
||||
if change == "credentials" {
|
||||
assert_ne!(
|
||||
requests[0].headers["authorization"],
|
||||
requests[1].headers["authorization"]
|
||||
);
|
||||
}
|
||||
if change == "callback" {
|
||||
assert_eq!(
|
||||
serde_json::from_slice::<Value>(&requests[0].body).unwrap()["temperature"],
|
||||
0.1
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::from_slice::<Value>(&requests[1].body).unwrap()["temperature"],
|
||||
0.8
|
||||
);
|
||||
}
|
||||
first.verify().await;
|
||||
second.verify().await;
|
||||
}
|
||||
|
|
@ -6,11 +6,11 @@ use std::sync::Arc;
|
|||
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_host::{call::HostedCompletion, lifecycle::CallEvent};
|
||||
use litellm_inference::responses::{
|
||||
use litellm_inference::test_support::{RecordingSecrets, no_secrets};
|
||||
use litellm_inference_responses::{
|
||||
route::Responses,
|
||||
types::{ResponsesCall, ResponsesOutput},
|
||||
};
|
||||
use litellm_inference::test_support::{RecordingSecrets, no_secrets};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
use wiremock::ResponseTemplate;
|
||||
|
|
@ -418,7 +418,7 @@ async fn websocket_operations_trace_outcomes_without_capturing_frames_or_credent
|
|||
traces: TraceCapture,
|
||||
) {
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use litellm_inference::responses::websocket::ResponsesWebSocketConnection;
|
||||
use litellm_inference_responses::websocket::ResponsesWebSocketConnection;
|
||||
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
279
litellm-rust/crates/inference-responses/tests/support/mod.rs
Normal file
279
litellm-rust/crates/inference-responses/tests/support/mod.rs
Normal file
|
|
@ -0,0 +1,279 @@
|
|||
//! Shared fixtures for route integration tests: a scripted upstream and a recording
|
||||
//! secret source.
|
||||
|
||||
#![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset
|
||||
|
||||
use std::{
|
||||
ops::ControlFlow,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
|
||||
use litellm_inference::test_support::{http_config, provider_http, resources};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use serde_json::Value;
|
||||
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
|
||||
|
||||
pub fn responses_route(
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> litellm_inference_responses::ResponsesRoute {
|
||||
let resources = resources();
|
||||
litellm_inference_responses::ResponsesRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
secrets,
|
||||
)
|
||||
}
|
||||
|
||||
pub async fn upstream(responses: impl IntoIterator<Item = ResponseTemplate>) -> MockServer {
|
||||
let server = MockServer::start().await;
|
||||
respond_in_order(&server, responses).await;
|
||||
server
|
||||
}
|
||||
|
||||
pub async fn respond_in_order(
|
||||
server: &MockServer,
|
||||
responses: impl IntoIterator<Item = ResponseTemplate>,
|
||||
) {
|
||||
for response in responses {
|
||||
Mock::given(any())
|
||||
.respond_with(response)
|
||||
.up_to_n_times(1)
|
||||
.mount(server)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn received(server: &MockServer) -> Vec<Request> {
|
||||
server
|
||||
.received_requests()
|
||||
.await
|
||||
.expect("request recording is on")
|
||||
}
|
||||
|
||||
pub async fn only_request(server: &MockServer) -> Request {
|
||||
let [request] = <[Request; 1]>::try_from(received(server).await)
|
||||
.unwrap_or_else(|requests| panic!("expected one request, got {}", requests.len()));
|
||||
request
|
||||
}
|
||||
|
||||
pub fn json_response(body: Value) -> ResponseTemplate {
|
||||
ResponseTemplate::new(200).set_body_json(body)
|
||||
}
|
||||
|
||||
pub fn status_response(status: u16, body: Value) -> ResponseTemplate {
|
||||
ResponseTemplate::new(status).set_body_json(body)
|
||||
}
|
||||
|
||||
pub trait ReceivedRequest {
|
||||
fn header(&self, name: &str) -> Option<&str>;
|
||||
fn header_values(&self, name: &str) -> Vec<&str>;
|
||||
fn json(&self) -> Value;
|
||||
fn body_text(&self) -> String;
|
||||
/// The path and query, as the request line carried them.
|
||||
fn target(&self) -> String;
|
||||
fn query(&self, name: &str) -> Option<String>;
|
||||
}
|
||||
|
||||
impl ReceivedRequest for Request {
|
||||
fn header(&self, name: &str) -> Option<&str> {
|
||||
self.headers.get(name).and_then(|value| value.to_str().ok())
|
||||
}
|
||||
|
||||
fn header_values(&self, name: &str) -> Vec<&str> {
|
||||
self.headers
|
||||
.get_all(name)
|
||||
.iter()
|
||||
.filter_map(|value| value.to_str().ok())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn json(&self) -> Value {
|
||||
serde_json::from_slice(&self.body).expect("request body is json")
|
||||
}
|
||||
|
||||
fn body_text(&self) -> String {
|
||||
String::from_utf8_lossy(&self.body).into_owned()
|
||||
}
|
||||
|
||||
fn target(&self) -> String {
|
||||
match self.url.query() {
|
||||
Some(query) => format!("{}?{query}", self.url.path()),
|
||||
None => self.url.path().to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
fn query(&self, name: &str) -> Option<String> {
|
||||
self.url
|
||||
.query_pairs()
|
||||
.find_map(|(key, value)| (key == name).then(|| value.into_owned()))
|
||||
}
|
||||
}
|
||||
|
||||
pub struct RecordingCall<P: litellm_host::protocol::Protocol> {
|
||||
pub request: Mutex<Option<P::Request>>,
|
||||
pub events: Arc<CallEvents>,
|
||||
pub chunks: Mutex<Vec<P::Chunk>>,
|
||||
pub head: Mutex<Option<P::StreamHead>>,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct CallEvents(pub Observations);
|
||||
pub struct Observations {
|
||||
pub sender: litellm_host::observation::ObservationSender,
|
||||
receiver: Mutex<tokio::sync::mpsc::Receiver<litellm_host::lifecycle::CallEvent>>,
|
||||
recorded: Mutex<Vec<litellm_host::lifecycle::CallEvent>>,
|
||||
}
|
||||
|
||||
impl Default for Observations {
|
||||
fn default() -> Self {
|
||||
let (sender, receiver) = litellm_host::observation::observation_channel(
|
||||
std::num::NonZeroUsize::new(128).unwrap(),
|
||||
);
|
||||
Self {
|
||||
sender,
|
||||
receiver: Mutex::new(receiver),
|
||||
recorded: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Observations {
|
||||
pub fn lock(
|
||||
&self,
|
||||
) -> std::sync::LockResult<std::sync::MutexGuard<'_, Vec<litellm_host::lifecycle::CallEvent>>>
|
||||
{
|
||||
let mut events = self.recorded.lock()?;
|
||||
let mut receiver = self.receiver.lock().unwrap();
|
||||
while let Ok(event) = receiver.try_recv() {
|
||||
events.push(event);
|
||||
}
|
||||
Ok(events)
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_host::lifecycle::CallObserver for CallEvents {
|
||||
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
|
||||
self.0.sender.emit(event);
|
||||
}
|
||||
}
|
||||
|
||||
impl<P: litellm_host::protocol::Protocol> RecordingCall<P> {
|
||||
pub fn new(request: P::Request) -> Self {
|
||||
Self {
|
||||
request: Mutex::new(Some(request)),
|
||||
events: Arc::new(CallEvents::default()),
|
||||
chunks: Mutex::new(Vec::new()),
|
||||
head: Mutex::new(None),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<P: litellm_host::protocol::Protocol> litellm_host::interceptors::Interceptors<P::Error>
|
||||
for RecordingCall<P>
|
||||
{
|
||||
async fn before_provider_request(
|
||||
&self,
|
||||
wire: litellm_host::interceptors::WireRequest,
|
||||
_: litellm_host::interceptors::RequestContext,
|
||||
) -> Result<litellm_host::interceptors::WireRequest, P::Error> {
|
||||
Ok(litellm_host::interceptors::WireRequest {
|
||||
headers: wire
|
||||
.headers
|
||||
.into_iter()
|
||||
.chain([("x-hook".into(), "called".into())])
|
||||
.collect(),
|
||||
..wire
|
||||
})
|
||||
}
|
||||
|
||||
async fn after_provider_response(
|
||||
&self,
|
||||
_: litellm_host::interceptors::RawResponse,
|
||||
) -> Result<(), P::Error> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl<P> RecordingCall<P>
|
||||
where
|
||||
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
|
||||
P::Error: From<litellm_host::machine::MachineFault>,
|
||||
{
|
||||
pub fn request(&self) -> Result<P::Request, P::Error> {
|
||||
self.request
|
||||
.lock()
|
||||
.unwrap()
|
||||
.take()
|
||||
.ok_or_else(|| litellm_host::machine::MachineFault::Abandoned.into())
|
||||
}
|
||||
pub fn runtime(&self) -> litellm_host_native::in_process::Host<'_, (), Self, Self> {
|
||||
litellm_host_native::in_process::Host {
|
||||
services: &(),
|
||||
interceptors: self,
|
||||
stream: self,
|
||||
observers: Some(&self.events.0.sender),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<P> litellm_host_native::in_process::StreamConsumer<P> for RecordingCall<P>
|
||||
where
|
||||
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
|
||||
P::Error: From<litellm_host::machine::MachineFault>,
|
||||
{
|
||||
async fn open_stream(&self, head: P::StreamHead) -> Result<ControlFlow<()>, P::Error> {
|
||||
*self.head.lock().unwrap() = Some(head);
|
||||
Ok(ControlFlow::Continue(()))
|
||||
}
|
||||
async fn send_chunk(&self, chunk: P::Chunk) -> Result<ControlFlow<()>, P::Error> {
|
||||
self.chunks.lock().unwrap().push(chunk);
|
||||
Ok(ControlFlow::Continue(()))
|
||||
}
|
||||
}
|
||||
impl<P> litellm_host::lifecycle::CallObserver for RecordingCall<P>
|
||||
where
|
||||
P: litellm_host::protocol::Protocol<HostCall = std::convert::Infallible>,
|
||||
P::Error: From<litellm_host::machine::MachineFault>,
|
||||
{
|
||||
fn observe(&self, event: litellm_host::lifecycle::CallEvent) {
|
||||
self.events.0.sender.emit(event);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct TraceCapture(Arc<Mutex<Vec<Value>>>);
|
||||
|
||||
impl TraceCapture {
|
||||
pub fn logger(&self) -> litellm_tracing::Logger {
|
||||
litellm_tracing::Logger::new(self.clone())
|
||||
}
|
||||
|
||||
pub fn records(&self) -> Vec<Value> {
|
||||
self.0.lock().unwrap().clone()
|
||||
}
|
||||
|
||||
pub fn summaries(&self, name: &str) -> Vec<Value> {
|
||||
self.records()
|
||||
.into_iter()
|
||||
.filter(|record| record["span_name"] == name)
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_tracing::Sink for TraceCapture {
|
||||
fn enabled(&self, metadata: &litellm_tracing::Metadata<'_>) -> bool {
|
||||
metadata.target().starts_with("litellm_inference")
|
||||
}
|
||||
|
||||
fn emit(&self, record: &litellm_tracing::Record) {
|
||||
self.0
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(Value::Object(record.fields.clone()));
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest::fixture]
|
||||
pub fn traces() -> TraceCapture {
|
||||
TraceCapture::default()
|
||||
}
|
||||
|
|
@ -52,7 +52,7 @@ impl From<litellm_host::machine::MachineFault> for RouteError {
|
|||
}
|
||||
|
||||
impl RouteError {
|
||||
pub(crate) fn post_call(error: Self) -> Self {
|
||||
pub fn post_call(error: Self) -> Self {
|
||||
Self::PostCallHook(Arc::new(error))
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -10,7 +10,6 @@ pub mod ocr;
|
|||
pub mod outbound;
|
||||
pub mod provider;
|
||||
pub mod resources;
|
||||
pub mod responses;
|
||||
#[cfg(feature = "test-support")]
|
||||
pub mod test_support;
|
||||
|
||||
|
|
|
|||
|
|
@ -12,8 +12,7 @@ use bytes::Bytes;
|
|||
use futures_util::{StreamExt, TryStreamExt, stream};
|
||||
use litellm_cache_memory::InMemoryCache;
|
||||
use litellm_cache_response::{
|
||||
CacheOptions, CachePolicy, CacheScope, ResponseCache, ResponseCacheConfig,
|
||||
ResponseCacheService, ResponseEnvelope,
|
||||
CacheOptions, CachePolicy, CacheScope, ResponseCache, ResponseCacheConfig, ResponseCacheService,
|
||||
};
|
||||
use litellm_host::{
|
||||
call::{CallOutput, OutputOf},
|
||||
|
|
@ -403,50 +402,6 @@ async fn an_invalid_cached_envelope_is_replaced_by_a_provider_result(#[case] poi
|
|||
assert_eq!(calls.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::chat_completion(json!({"kind":"Response","value":{"id":"chat-1","model":"test","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2}}}))]
|
||||
#[case::wrong_envelope(json!({"kind":"Stream","value":"data: [DONE]\n\n"}))]
|
||||
#[tokio::test]
|
||||
async fn responses_refetches_instead_of_deserializing_another_api_response(
|
||||
#[case] poisoned: Value,
|
||||
) {
|
||||
use litellm_inference::responses::route::Responses;
|
||||
use litellm_llms_types::formats::responses::ResponsesApiResponse;
|
||||
|
||||
let cache: Arc<dyn ResponseCacheService> = Arc::new(InvalidEntryCache(
|
||||
ResponseCache::new(Arc::new(InMemoryCache::default())),
|
||||
serde_json::to_value(ResponseEnvelope::new("responses", poisoned)).unwrap(),
|
||||
));
|
||||
let calls = AtomicUsize::new(0);
|
||||
for _ in 0..2 {
|
||||
let response = execute_unary::<Responses, _, _>(
|
||||
cache_request(json!({"input":"hello"})),
|
||||
Some(cache.clone()),
|
||||
Some(CacheOptions::new(CacheScope::Shared)),
|
||||
&(),
|
||||
None,
|
||||
|| async {
|
||||
calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(ResponsesApiResponse {
|
||||
id: "fresh-response".into(),
|
||||
model: "test".into(),
|
||||
output: vec![
|
||||
json!({"type":"message","content":[{"type":"output_text","text":"fresh"}]}),
|
||||
],
|
||||
extra: [("status".into(), json!("completed"))]
|
||||
.into_iter()
|
||||
.collect(),
|
||||
})
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.id, "fresh-response");
|
||||
assert_eq!(response.output[0]["content"][0]["text"], "fresh");
|
||||
}
|
||||
assert_eq!(calls.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::system("system", json!("answer ALPHA"), json!("answer BETA"))]
|
||||
#[case::stop_sequences("stop_sequences", json!(["STOP"]), json!(["END"]))]
|
||||
|
|
@ -833,43 +788,6 @@ async fn namespaces_and_surfaces_isolate_entries_on_shared_storage() {
|
|||
assert_eq!(calls.load(Ordering::SeqCst), 3);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::completed("completed", 1)]
|
||||
#[case::incomplete("incomplete", 2)]
|
||||
#[tokio::test]
|
||||
async fn responses_cache_only_reuses_completed_responses(
|
||||
cache: Arc<dyn ResponseCacheService>,
|
||||
#[case] status: &str,
|
||||
#[case] expected_calls: usize,
|
||||
) {
|
||||
use litellm_inference::responses::route::Responses;
|
||||
use litellm_llms_types::formats::responses::ResponsesApiResponse;
|
||||
|
||||
let calls = AtomicUsize::new(0);
|
||||
for _ in 0..2 {
|
||||
let response = execute_unary::<Responses, _, _>(
|
||||
cache_request(json!({"input":"hello"})),
|
||||
Some(cache.clone()),
|
||||
Some(CacheOptions::new(CacheScope::Shared)),
|
||||
&(),
|
||||
None,
|
||||
|| async {
|
||||
let call = calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(ResponsesApiResponse {
|
||||
id: call.to_string(),
|
||||
model: "test".into(),
|
||||
output: Vec::new(),
|
||||
extra: [("status".into(), json!(status))].into_iter().collect(),
|
||||
})
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.extra.get("status"), Some(&json!(status)));
|
||||
}
|
||||
assert_eq!(calls.load(Ordering::SeqCst), expected_calls);
|
||||
}
|
||||
|
||||
mod support;
|
||||
use support::traces;
|
||||
|
||||
|
|
@ -1028,9 +946,6 @@ impl Interceptors<RouteError> for ChangingHooks {
|
|||
#[case::messages_credentials("messages", "credentials")]
|
||||
#[case::messages_endpoint("messages", "endpoint")]
|
||||
#[case::messages_callback("messages", "callback")]
|
||||
#[case::responses_credentials("responses", "credentials")]
|
||||
#[case::responses_endpoint("responses", "endpoint")]
|
||||
#[case::responses_callback("responses", "callback")]
|
||||
#[tokio::test]
|
||||
async fn cache_identity_follows_resolved_configuration_and_request_callbacks(
|
||||
cache: Arc<dyn ResponseCacheService>,
|
||||
|
|
@ -1041,19 +956,14 @@ async fn cache_identity_follows_resolved_configuration_and_request_callbacks(
|
|||
use litellm_inference::{
|
||||
chat_completions::{ChatCompletionsRoute, types::ChatCompletionsRequest},
|
||||
messages::MessagesCall,
|
||||
responses::types::ResponsesCall,
|
||||
};
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
|
||||
|
||||
let first = MockServer::start().await;
|
||||
let second = MockServer::start().await;
|
||||
let response = if surface == "responses" {
|
||||
json!({"id":"response-test", "model":"test", "output":[], "status":"completed"})
|
||||
} else {
|
||||
json!({"id":"message-test", "type":"message", "role":"assistant", "model":"test",
|
||||
"content":[{"type":"text", "text":"answer"}], "stop_reason":"end_turn", "stop_sequence":null,
|
||||
"usage":{"input_tokens":3,"output_tokens":2}})
|
||||
};
|
||||
let response = json!({"id":"message-test", "type":"message", "role":"assistant", "model":"test",
|
||||
"content":[{"type":"text", "text":"answer"}], "stop_reason":"end_turn", "stop_sequence":null,
|
||||
"usage":{"input_tokens":3,"output_tokens":2}});
|
||||
Mock::given(method("POST"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(response.clone()))
|
||||
.expect(if change == "endpoint" { 1 } else { 2 })
|
||||
|
|
@ -1116,26 +1026,6 @@ async fn cache_identity_follows_resolved_configuration_and_request_callbacks(
|
|||
api_key:None,api_base:None,custom_llm_provider:None,extra_headers:None,provider_specific_header:None,timeout:None,shaping:Default::default(),
|
||||
}, &hooks, None).await.unwrap();
|
||||
}
|
||||
"responses" => {
|
||||
support::responses_route(secrets.clone())
|
||||
.with_cache(cache)
|
||||
.execute(
|
||||
ResponsesCall {
|
||||
model: "test".into(),
|
||||
input: json!("hello"),
|
||||
optional_params: Default::default(),
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
},
|
||||
&hooks,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
_ => unreachable!(),
|
||||
}
|
||||
}
|
||||
|
|
@ -1153,12 +1043,10 @@ async fn cache_identity_follows_resolved_configuration_and_request_callbacks(
|
|||
}
|
||||
let requests = first.received_requests().await.unwrap();
|
||||
if change == "credentials" {
|
||||
let header = if surface == "responses" {
|
||||
"authorization"
|
||||
} else {
|
||||
"x-api-key"
|
||||
};
|
||||
assert_ne!(requests[0].headers[header], requests[1].headers[header]);
|
||||
assert_ne!(
|
||||
requests[0].headers["x-api-key"],
|
||||
requests[1].headers["x-api-key"]
|
||||
);
|
||||
}
|
||||
if change == "callback" {
|
||||
assert_eq!(
|
||||
|
|
|
|||
|
|
@ -37,17 +37,6 @@ pub fn chat_completions_route() -> litellm_inference::chat_completions::ChatComp
|
|||
)
|
||||
}
|
||||
|
||||
pub fn responses_route(
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> litellm_inference::responses::ResponsesRoute {
|
||||
let resources = resources();
|
||||
litellm_inference::responses::ResponsesRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
secrets,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_ocr_route(
|
||||
resources: &litellm_inference::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ litellm-auth-aws.workspace = true
|
|||
litellm-callbacks-legacy-python.workspace = true
|
||||
litellm-inference.workspace = true
|
||||
litellm-inference-transcription.workspace = true
|
||||
litellm-inference-responses.workspace = true
|
||||
litellm-core-utils.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
mod host;
|
||||
|
||||
use litellm_inference::responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
|
||||
use litellm_inference_responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
|
||||
use pyo3::{
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple},
|
||||
|
|
@ -81,7 +81,7 @@ fn run_public(
|
|||
py,
|
||||
arguments,
|
||||
move |py, arguments, request| {
|
||||
let route = litellm_inference::responses::ResponsesRoute::new(
|
||||
let route = litellm_inference_responses::ResponsesRoute::new(
|
||||
crate::http::provider_client(py, arguments, asynchronous)?
|
||||
.map_err(crate::http::client_error)?,
|
||||
crate::http::resources().auth.clone(),
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ use std::convert::Infallible;
|
|||
|
||||
use super::super::inference::InferenceHost;
|
||||
use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned};
|
||||
use litellm_inference::responses::{Error, route::Responses, types::ResponsesCall};
|
||||
use litellm_inference_responses::{Error, route::Responses, types::ResponsesCall};
|
||||
use pyo3::{
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue