mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
feat(rust): support the HTTP Responses API (#43464)
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
0d6b8b5ab4
commit
0a21f24d50
15 changed files with 800 additions and 7 deletions
91
litellm-rust/crates/core/src/responses/handler.rs
Normal file
91
litellm-rust/crates/core/src/responses/handler.rs
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use futures_util::StreamExt;
|
||||
use litellm_host::{
|
||||
event::{MachineEvent, RawResponse, WireRequest},
|
||||
hooks::RouteHooks,
|
||||
};
|
||||
use litellm_llms::base_llm::auth::{Authenticated, resolve_auth};
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
types::{ProviderResponsesRequest, ResponsesOutput, ResponsesStreamHead},
|
||||
};
|
||||
|
||||
pub(super) async fn execute(
|
||||
http: &litellm_http::Client,
|
||||
auth: &litellm_auth::AuthServices,
|
||||
request: ProviderResponsesRequest,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let authenticated = resolve_auth(auth, request.environment, &|_| None).await?;
|
||||
let wire = hooks
|
||||
.before_provider_request(
|
||||
WireRequest {
|
||||
url: request.url,
|
||||
headers: authenticated.headers,
|
||||
body: request.body,
|
||||
},
|
||||
request.context,
|
||||
)
|
||||
.await?;
|
||||
let stream = match wire.body.get("stream") {
|
||||
None => false,
|
||||
Some(serde_json::Value::Bool(value)) => *value,
|
||||
Some(_) => return Err(Error::InvalidRequest("stream must be a boolean".into())),
|
||||
};
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
Authenticated {
|
||||
headers: wire.headers,
|
||||
signer: authenticated.signer,
|
||||
},
|
||||
wire.url,
|
||||
&wire.body,
|
||||
Some(request.timeout.unwrap_or(Duration::from_secs(600))),
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(network)?;
|
||||
let status = response.status().as_u16();
|
||||
if !response.status().is_success() {
|
||||
let body = response.text().await.map_err(network)?;
|
||||
return Err(litellm_http::transport::Error::Http {
|
||||
status,
|
||||
body: litellm_http::request::truncate_error_body(&body),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
if stream {
|
||||
let headers = response
|
||||
.headers()
|
||||
.iter()
|
||||
.filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_owned())))
|
||||
.collect();
|
||||
let chunks = response
|
||||
.bytes_stream()
|
||||
.map(|chunk| chunk.map_err(network))
|
||||
.boxed();
|
||||
return Ok(ResponsesOutput::Stream {
|
||||
head: ResponsesStreamHead { headers },
|
||||
chunks,
|
||||
});
|
||||
}
|
||||
let body = response.text().await.map_err(network)?;
|
||||
hooks
|
||||
.on_event(MachineEvent::ResponseReceived {
|
||||
raw: RawResponse { body: body.clone() },
|
||||
})
|
||||
.await
|
||||
.map_err(Error::post_call)?;
|
||||
let value = serde_json::from_str(&body)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into()))?;
|
||||
request
|
||||
.config
|
||||
.transform_response_api_response(value)
|
||||
.map(ResponsesOutput::Complete)
|
||||
.map_err(Error::from)
|
||||
}
|
||||
|
||||
fn network(error: reqwest::Error) -> Error {
|
||||
litellm_http::transport::Error::Network(error.to_string()).into()
|
||||
}
|
||||
|
|
@ -1,2 +1,52 @@
|
|||
pub use crate::error::RouteError as Error;
|
||||
pub mod websocket;
|
||||
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub mod route;
|
||||
pub mod types;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_host::hooks::RouteHooks;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use types::{ResponsesCall, ResponsesOutput};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ResponsesRoute {
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
}
|
||||
|
||||
impl ResponsesRoute {
|
||||
pub fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self {
|
||||
http,
|
||||
auth,
|
||||
secrets,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
litellm_host::lifecycle::observe_call(hooks.observer(), self.run(call, hooks)).await
|
||||
}
|
||||
|
||||
async fn run(
|
||||
&self,
|
||||
call: ResponsesCall,
|
||||
hooks: &impl RouteHooks<Error>,
|
||||
) -> Result<ResponsesOutput, Error> {
|
||||
let request = prepare::prepare(call, self.secrets.as_ref()).await?;
|
||||
handler::execute(&self.http, &self.auth, request, hooks).await
|
||||
}
|
||||
}
|
||||
|
|
|
|||
57
litellm-rust/crates/core/src/responses/prepare.rs
Normal file
57
litellm-rust/crates/core/src/responses/prepare.rs
Normal file
|
|
@ -0,0 +1,57 @@
|
|||
use litellm_host::event::RequestContext;
|
||||
use litellm_llms::{
|
||||
base_llm::responses::transformation::BaseResponsesApiConfig,
|
||||
openai::responses::transformation::OpenAiResponsesApiConfig,
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
Error,
|
||||
types::{ProviderResponsesRequest, ResponsesCall},
|
||||
};
|
||||
|
||||
pub(super) async fn prepare(
|
||||
call: ResponsesCall,
|
||||
secrets: &dyn SecretSource,
|
||||
) -> Result<ProviderResponsesRequest, Error> {
|
||||
let provider = call.custom_llm_provider.as_deref().unwrap_or("openai");
|
||||
if provider != "openai" {
|
||||
return Err(Error::Unsupported("native HTTP responses provider"));
|
||||
}
|
||||
let model = call.model.strip_prefix("openai/").unwrap_or(&call.model);
|
||||
if model.is_empty() || model.contains('/') {
|
||||
return Err(Error::InvalidProvider(call.model));
|
||||
}
|
||||
let config: &'static dyn BaseResponsesApiConfig = &OpenAiResponsesApiConfig;
|
||||
let snapshot = secrets
|
||||
.resolve(config.secret_names(call.api_key.as_deref(), call.api_base.as_deref()))
|
||||
.await?;
|
||||
let lookup = |name: &str| snapshot.get(name);
|
||||
let environment = config.validate_environment(
|
||||
litellm_http::request::string_headers("responses", call.extra_headers)?,
|
||||
call.api_key.as_deref(),
|
||||
&lookup,
|
||||
)?;
|
||||
let context = RequestContext {
|
||||
model: model.into(),
|
||||
custom_llm_provider: provider.into(),
|
||||
optional_params: Value::Object(call.optional_params.clone()),
|
||||
secret_fields: Vec::new(),
|
||||
api_key: match &environment.auth {
|
||||
litellm_llms::base_llm::auth::AuthScheme::Credential { secret, .. } => {
|
||||
Some(secret.clone())
|
||||
}
|
||||
_ => None,
|
||||
},
|
||||
};
|
||||
let body = config.transform_responses_api_request(model, call.input, call.optional_params)?;
|
||||
Ok(ProviderResponsesRequest {
|
||||
url: config.get_complete_url(call.api_base.as_deref(), &lookup),
|
||||
config,
|
||||
environment,
|
||||
body,
|
||||
context,
|
||||
timeout: call.timeout,
|
||||
})
|
||||
}
|
||||
32
litellm-rust/crates/core/src/responses/route.rs
Normal file
32
litellm-rust/crates/core/src/responses/route.rs
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
use std::convert::Infallible;
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_host::{
|
||||
call::{HostedMachine, hosted_call},
|
||||
protocol::Protocol,
|
||||
};
|
||||
use litellm_types::responses::main::ResponsesApiResponse;
|
||||
|
||||
use super::{
|
||||
Error, ResponsesRoute,
|
||||
types::{ResponsesCall, ResponsesStreamHead},
|
||||
};
|
||||
|
||||
pub struct Responses;
|
||||
|
||||
impl Protocol for Responses {
|
||||
type Response = ResponsesApiResponse;
|
||||
type Error = Error;
|
||||
type Request = ResponsesCall;
|
||||
type HostCall = Infallible;
|
||||
type Chunk = Bytes;
|
||||
type StreamHead = ResponsesStreamHead;
|
||||
}
|
||||
|
||||
impl ResponsesRoute {
|
||||
pub fn machine(self, call: ResponsesCall) -> HostedMachine<Responses> {
|
||||
hosted_call(call, move |call, _, hooks| async move {
|
||||
self.run(call, &hooks).await
|
||||
})
|
||||
}
|
||||
}
|
||||
37
litellm-rust/crates/core/src/responses/types.rs
Normal file
37
litellm-rust/crates/core/src/responses/types.rs
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_host::call::CallOutput;
|
||||
use litellm_llms::base_llm::{
|
||||
auth::ValidatedEnvironment, responses::transformation::BaseResponsesApiConfig,
|
||||
};
|
||||
use litellm_types::responses::main::ResponsesApiResponse;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::Error;
|
||||
|
||||
pub struct ResponsesCall {
|
||||
pub model: String,
|
||||
pub input: Value,
|
||||
pub optional_params: Map<String, Value>,
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub struct ResponsesStreamHead {
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
pub type ResponsesOutput = CallOutput<ResponsesApiResponse, ResponsesStreamHead, Bytes, Error>;
|
||||
|
||||
pub(super) struct ProviderResponsesRequest {
|
||||
pub config: &'static dyn BaseResponsesApiConfig,
|
||||
pub environment: ValidatedEnvironment,
|
||||
pub url: String,
|
||||
pub body: Value,
|
||||
pub context: litellm_host::event::RequestContext,
|
||||
pub timeout: Option<Duration>,
|
||||
}
|
||||
233
litellm-rust/crates/core/tests/responses.rs
Normal file
233
litellm-rust/crates/core/tests/responses.rs
Normal file
|
|
@ -0,0 +1,233 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_core::responses::{
|
||||
route::Responses,
|
||||
types::{ResponsesCall, ResponsesOutput},
|
||||
};
|
||||
use litellm_host::{call::HostedCompletion, event::CallEvent};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
use wiremock::ResponseTemplate;
|
||||
|
||||
mod support;
|
||||
use support::*;
|
||||
|
||||
#[fixture]
|
||||
fn call() -> ResponsesCall {
|
||||
ResponsesCall {
|
||||
model: "openai/test-model".into(),
|
||||
input: json!("hello"),
|
||||
optional_params: Default::default(),
|
||||
api_key: Some("test-key".into()),
|
||||
api_base: None,
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::direct(false)]
|
||||
#[case::hosted(true)]
|
||||
#[tokio::test]
|
||||
async fn http_responses_share_execution_and_hooks(call: ResponsesCall, #[case] hosted: bool) {
|
||||
let body = json!({"id": "response-1", "model": "test-model", "output": [{"type":"message", "content":[]}], "usage":{"total_tokens":7}, "provider_extra": true});
|
||||
let upstream = upstream([json_response(body.clone())]).await;
|
||||
let host = RecordingCall::<Responses>::new(ResponsesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
});
|
||||
let response = if hosted {
|
||||
let HostedCompletion::Complete(response) = litellm_host::in_process::run_hosted(
|
||||
responses_route(no_secrets()).machine(host.request().unwrap()),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.unwrap() else {
|
||||
panic!()
|
||||
};
|
||||
response
|
||||
} else {
|
||||
let call = host.request.lock().unwrap().take().unwrap();
|
||||
let ResponsesOutput::Complete(response) = responses_route(no_secrets())
|
||||
.execute(call, &host)
|
||||
.await
|
||||
.unwrap()
|
||||
else {
|
||||
panic!()
|
||||
};
|
||||
response
|
||||
};
|
||||
assert_eq!(serde_json::to_value(response).unwrap(), body);
|
||||
let sent = only_request(&upstream).await;
|
||||
assert_eq!(sent.url.path(), "/responses");
|
||||
assert_eq!(sent.header("authorization"), Some("Bearer test-key"));
|
||||
assert_eq!(sent.header("x-hook"), Some("called"));
|
||||
assert_eq!(sent.json(), json!({"model":"test-model", "input":"hello"}));
|
||||
assert!(matches!(
|
||||
&host.events.0.lock().unwrap()[..],
|
||||
[
|
||||
CallEvent::Started { .. },
|
||||
CallEvent::Machine(_),
|
||||
CallEvent::Succeeded { .. }
|
||||
]
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::direct(false)]
|
||||
#[case::hosted(true)]
|
||||
#[tokio::test]
|
||||
async fn streaming_keeps_headers_and_bytes_and_finishes_after_consumption(
|
||||
call: ResponsesCall,
|
||||
#[case] hosted: bool,
|
||||
) {
|
||||
let body = "event: response.completed\ndata: {\"type\":\"response.completed\"}\n\n";
|
||||
let upstream = upstream([ResponseTemplate::new(200)
|
||||
.insert_header("x-request-id", "response-stream")
|
||||
.set_body_raw(body, "text/event-stream")])
|
||||
.await;
|
||||
let host = RecordingCall::<Responses>::new(ResponsesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
optional_params: json!({"stream":true}).as_object().unwrap().clone(),
|
||||
..call
|
||||
});
|
||||
let (headers, bytes) = if hosted {
|
||||
assert_eq!(
|
||||
litellm_host::in_process::run_hosted(
|
||||
responses_route(no_secrets()).machine(host.request().unwrap()),
|
||||
host.runtime(),
|
||||
)
|
||||
.await
|
||||
.unwrap(),
|
||||
HostedCompletion::StreamEnded
|
||||
);
|
||||
(
|
||||
host.head.lock().unwrap().take().unwrap().headers,
|
||||
host.chunks.lock().unwrap().concat(),
|
||||
)
|
||||
} else {
|
||||
let call = host.request.lock().unwrap().take().unwrap();
|
||||
let ResponsesOutput::Stream { head, chunks } = responses_route(no_secrets())
|
||||
.execute(call, &host)
|
||||
.await
|
||||
.unwrap()
|
||||
else {
|
||||
panic!()
|
||||
};
|
||||
assert_eq!(host.events.0.lock().unwrap().len(), 1);
|
||||
(
|
||||
head.headers,
|
||||
chunks.try_collect::<Vec<_>>().await.unwrap().concat(),
|
||||
)
|
||||
};
|
||||
assert!(headers.contains(&("x-request-id".into(), "response-stream".into())));
|
||||
assert_eq!(bytes, body.as_bytes());
|
||||
assert!(matches!(
|
||||
&host.events.0.lock().unwrap()[..],
|
||||
[CallEvent::Started { .. }, CallEvent::Succeeded { .. }]
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::http(429, json!({"error":"limited"}))]
|
||||
#[case::invalid_response(200, json!({"unexpected":true}))]
|
||||
#[tokio::test]
|
||||
async fn provider_failures_emit_failure_once(
|
||||
call: ResponsesCall,
|
||||
#[case] status: u16,
|
||||
#[case] body: serde_json::Value,
|
||||
) {
|
||||
let upstream = upstream([ResponseTemplate::new(status).set_body_json(body)]).await;
|
||||
let host = RecordingCall::<Responses>::new(ResponsesCall {
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
});
|
||||
let call = host.request.lock().unwrap().take().unwrap();
|
||||
assert!(
|
||||
responses_route(no_secrets())
|
||||
.execute(call, &host)
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
assert_eq!(received(&upstream).await.len(), 1);
|
||||
let events = host.events.0.lock().unwrap();
|
||||
assert!(matches!(events.last(), Some(CallEvent::Failed { .. })));
|
||||
assert_eq!(
|
||||
events
|
||||
.iter()
|
||||
.filter(|event| matches!(
|
||||
event,
|
||||
CallEvent::Failed { .. } | CallEvent::Succeeded { .. }
|
||||
))
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::explicit(true)]
|
||||
#[case::from_secrets(false)]
|
||||
#[tokio::test]
|
||||
async fn credentials_and_endpoint_are_resolved_only_when_needed(
|
||||
call: ResponsesCall,
|
||||
#[case] explicit: bool,
|
||||
) {
|
||||
let upstream = upstream([json_response(
|
||||
json!({"id":"response", "model":"test-model", "output":[]}),
|
||||
)])
|
||||
.await;
|
||||
let base = upstream.uri();
|
||||
let key = "resolved-test-key";
|
||||
let secrets = Arc::new(RecordingSecrets::new([
|
||||
("OPENAI_API_KEY", key),
|
||||
("OPENAI_BASE_URL", base.as_str()),
|
||||
]));
|
||||
let call = ResponsesCall {
|
||||
api_key: explicit.then(|| key.into()),
|
||||
api_base: explicit.then(|| base.clone()),
|
||||
..call
|
||||
};
|
||||
responses_route(secrets.clone())
|
||||
.execute(call, &())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
only_request(&upstream).await.header("authorization"),
|
||||
Some(format!("Bearer {key}").as_str())
|
||||
);
|
||||
if explicit {
|
||||
assert!(secrets.requested().is_empty());
|
||||
} else {
|
||||
assert!(secrets.requested().contains(&"OPENAI_API_KEY".into()));
|
||||
assert!(secrets.requested().contains(&"OPENAI_BASE_URL".into()));
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::provider("other/model", "other")]
|
||||
#[case::conflicting_prefix("other/model", "openai")]
|
||||
#[tokio::test]
|
||||
async fn unsupported_providers_fail_before_secrets_or_transport(
|
||||
call: ResponsesCall,
|
||||
#[case] model: &str,
|
||||
#[case] provider: &str,
|
||||
) {
|
||||
let upstream = upstream([]).await;
|
||||
let secrets = Arc::new(RecordingSecrets::failing());
|
||||
let call = ResponsesCall {
|
||||
model: model.into(),
|
||||
custom_llm_provider: Some(provider.into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
..call
|
||||
};
|
||||
assert!(
|
||||
responses_route(secrets.clone())
|
||||
.execute(call, &())
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
assert!(secrets.requested().is_empty());
|
||||
assert!(received(&upstream).await.is_empty());
|
||||
}
|
||||
|
|
@ -57,6 +57,15 @@ pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletio
|
|||
)
|
||||
}
|
||||
|
||||
pub fn responses_route(secrets: Arc<dyn SecretSource>) -> litellm_core::responses::ResponsesRoute {
|
||||
let resources = resources();
|
||||
litellm_core::responses::ResponsesRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
secrets,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn audio_transcription_route() -> litellm_core::audio_transcription::AudioTranscriptionRoute {
|
||||
let resources = resources();
|
||||
litellm_core::audio_transcription::AudioTranscriptionRoute::new(
|
||||
|
|
|
|||
|
|
@ -9,13 +9,14 @@ mod error;
|
|||
pub mod messages;
|
||||
mod ocr;
|
||||
mod request;
|
||||
mod responses;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use axum::{Router, routing::post};
|
||||
use litellm_core::{
|
||||
audio_transcription::AudioTranscriptionRoute, chat_completions::ChatCompletionsRoute,
|
||||
messages::MessagesRoute, ocr::OcrRoute, resources::CoreResources,
|
||||
messages::MessagesRoute, ocr::OcrRoute, resources::CoreResources, responses::ResponsesRoute,
|
||||
};
|
||||
use litellm_http::{ClientVariant, HttpClientConfig, media::UrlPolicy};
|
||||
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
|
||||
|
|
@ -30,6 +31,7 @@ pub struct Gateway {
|
|||
pub chat_completions: ChatCompletionsRoute,
|
||||
pub messages: MessagesRoute,
|
||||
pub ocr: OcrRoute,
|
||||
pub responses: ResponsesRoute,
|
||||
pub models: ModelList,
|
||||
pub secrets: Arc<dyn SecretSource>,
|
||||
pub resources: CoreResources,
|
||||
|
|
@ -56,7 +58,8 @@ impl Gateway {
|
|||
auth.clone(),
|
||||
secrets.clone(),
|
||||
),
|
||||
messages: MessagesRoute::new(provider, auth.clone(), secrets.clone()),
|
||||
messages: MessagesRoute::new(provider.clone(), auth.clone(), secrets.clone()),
|
||||
responses: ResponsesRoute::new(provider, auth.clone(), secrets.clone()),
|
||||
ocr: OcrRoute::new(OcrClient::new(
|
||||
&resources.pool,
|
||||
&http,
|
||||
|
|
@ -87,8 +90,8 @@ pub fn router(gateway: Arc<Gateway>) -> Router {
|
|||
"/v1/audio/transcriptions",
|
||||
post(audio_transcription::create),
|
||||
)
|
||||
.route("/responses", post(request::unsupported))
|
||||
.route("/v1/responses", post(request::unsupported))
|
||||
.route("/responses", post(responses::create))
|
||||
.route("/v1/responses", post(responses::create))
|
||||
.route("/embeddings", post(request::unsupported))
|
||||
.route("/v1/embeddings", post(request::unsupported))
|
||||
.route("/completions", post(request::unsupported))
|
||||
|
|
|
|||
37
litellm-rust/crates/gateway-inference/src/responses.rs
Normal file
37
litellm-rust/crates/gateway-inference/src/responses.rs
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use axum::{Json, body::Bytes, extract::State, response::Response};
|
||||
use litellm_core::responses::{route::Responses, types::ResponsesCall};
|
||||
use litellm_host_http::Sse;
|
||||
use serde_json::json;
|
||||
|
||||
use crate::{Error, Gateway, JsonObject, request};
|
||||
|
||||
pub(crate) async fn create(
|
||||
State(gateway): State<Arc<Gateway>>,
|
||||
JsonObject(body): JsonObject,
|
||||
) -> Result<Response, Error> {
|
||||
let deployment = request::resolve_deployment(&gateway, &body)?;
|
||||
let call = ResponsesCall {
|
||||
model: deployment.model.clone(),
|
||||
input: body.get("input").cloned().unwrap_or_default(),
|
||||
optional_params: body
|
||||
.into_iter()
|
||||
.filter(|(name, _)| !matches!(name.as_str(), "model" | "input"))
|
||||
.collect(),
|
||||
api_key: deployment.api_key.clone(),
|
||||
api_base: deployment.api_base.clone(),
|
||||
custom_llm_provider: deployment.custom_llm_provider.clone(),
|
||||
extra_headers: None,
|
||||
timeout: deployment.timeout,
|
||||
};
|
||||
let machine = gateway.responses.clone().machine(call);
|
||||
let stream = Sse::<Responses, _, _>::new(Json, |error| {
|
||||
let error = Error::from(error);
|
||||
Bytes::from(format!(
|
||||
"event: error\ndata: {}\n\n",
|
||||
json!({"type": "error", "code": error.status().as_u16().to_string(), "message": error.to_string(), "param": null})
|
||||
))
|
||||
});
|
||||
Ok(litellm_host_http::serve(machine, (), (), stream).await?)
|
||||
}
|
||||
99
litellm-rust/crates/gateway-inference/tests/responses.rs
Normal file
99
litellm-rust/crates/gateway-inference/tests/responses.rs
Normal file
|
|
@ -0,0 +1,99 @@
|
|||
mod support;
|
||||
|
||||
use axum::body::to_bytes;
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{body_json, header, method, path},
|
||||
};
|
||||
|
||||
#[rstest]
|
||||
#[case::completed(false)]
|
||||
#[case::streaming(true)]
|
||||
#[tokio::test]
|
||||
async fn responses_aliases_run_the_core_route(
|
||||
#[case] stream: bool,
|
||||
#[values("/responses", "/v1/responses")] route: &str,
|
||||
) {
|
||||
let upstream = MockServer::start().await;
|
||||
let completed =
|
||||
json!({"id": "response-1", "model": "test-model", "output": [], "provider_extra": true});
|
||||
let events = format!(
|
||||
"event: response.completed\ndata: {}\n\n",
|
||||
json!({"type": "response.completed", "response": completed})
|
||||
);
|
||||
let template = if stream {
|
||||
ResponseTemplate::new(200)
|
||||
.insert_header("content-type", "text/event-stream")
|
||||
.set_body_string(&events)
|
||||
} else {
|
||||
ResponseTemplate::new(200).set_body_json(&completed)
|
||||
};
|
||||
Mock::given(method("POST"))
|
||||
.and(path("/responses"))
|
||||
.and(header("authorization", "Bearer test-key"))
|
||||
.and(body_json(json!({"model": "test-model", "input": "hello", "stream": stream, "metadata": {"caller": "test"}})))
|
||||
.respond_with(template)
|
||||
.expect(1)
|
||||
.mount(&upstream).await;
|
||||
let response = support::post(support::app("openai/test-model", &upstream.uri()), route,
|
||||
json!({"model": "public/model", "input": "hello", "stream": stream, "metadata": {"caller": "test"}})).await;
|
||||
assert_eq!(response.status(), 200);
|
||||
if stream {
|
||||
assert_eq!(response.headers()["content-type"], "text/event-stream");
|
||||
assert_eq!(
|
||||
to_bytes(response.into_body(), 4096).await.unwrap().as_ref(),
|
||||
events.as_bytes()
|
||||
);
|
||||
} else {
|
||||
assert_eq!(support::json(response).await, completed);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::missing_input(json!({"model": "public/model"}), 400)]
|
||||
#[case::invalid_stream(json!({"model": "public/model", "input": "hello", "stream": "yes"}), 400)]
|
||||
#[case::unknown_model(json!({"model": "unknown", "input": "hello"}), 400)]
|
||||
#[tokio::test]
|
||||
async fn invalid_responses_requests_never_reach_the_provider(
|
||||
#[case] body: serde_json::Value,
|
||||
#[case] status: u16,
|
||||
) {
|
||||
let upstream = MockServer::start().await;
|
||||
let response = support::post(
|
||||
support::app("openai/test-model", &upstream.uri()),
|
||||
"/v1/responses",
|
||||
body,
|
||||
)
|
||||
.await;
|
||||
assert_eq!(response.status(), status);
|
||||
assert_eq!(support::json(response).await["error"]["code"], status);
|
||||
assert!(upstream.received_requests().await.unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn responses_preserve_upstream_errors_without_retry() {
|
||||
let upstream = MockServer::start().await;
|
||||
Mock::given(method("POST"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(429).set_body_json(json!({"error": {"message": "slow down"}})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&upstream)
|
||||
.await;
|
||||
let response = support::post(
|
||||
support::app("openai/test-model", &upstream.uri()),
|
||||
"/v1/responses",
|
||||
json!({"model": "public/model", "input": "hello"}),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(response.status(), 429);
|
||||
assert!(
|
||||
support::json(response).await["error"]["message"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("slow down")
|
||||
);
|
||||
}
|
||||
|
|
@ -227,6 +227,7 @@ async fn upload_audio_format_validation_matches_core(#[case] filename: &str) {
|
|||
|
||||
#[rstest]
|
||||
#[case::chat("/v1/chat/completions", false)]
|
||||
#[case::responses("/v1/responses", false)]
|
||||
#[case::deployment("/openai/deployments/public%2Fmodel/chat/completions", false)]
|
||||
#[case::messages("/v1/messages", true)]
|
||||
#[case::ocr("/v1/ocr", false)]
|
||||
|
|
|
|||
|
|
@ -1,7 +1,39 @@
|
|||
use litellm_types::responses::main::ResponsesApiResponse;
|
||||
use litellm_types::responses::streaming_websocket::ResponsesWsEvent;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::Error;
|
||||
use crate::{Error, base_llm::auth::ValidatedEnvironment};
|
||||
|
||||
pub trait BaseResponsesApiConfig: Sync {
|
||||
fn secret_names(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
api_base: Option<&str>,
|
||||
) -> &'static [&'static str];
|
||||
|
||||
fn validate_environment(
|
||||
&self,
|
||||
headers: Vec<(String, String)>,
|
||||
api_key: Option<&str>,
|
||||
lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error>;
|
||||
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String;
|
||||
|
||||
fn transform_responses_api_request(
|
||||
&self,
|
||||
model: &str,
|
||||
input: Value,
|
||||
params: Map<String, Value>,
|
||||
) -> Result<Value, Error>;
|
||||
|
||||
fn transform_response_api_response(&self, body: Value) -> Result<ResponsesApiResponse, Error>;
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ResponsesWsTransformResult {
|
||||
|
|
|
|||
|
|
@ -1,9 +1,17 @@
|
|||
use litellm_types::responses::main::ResponsesApiResponse;
|
||||
use litellm_types::responses::streaming_websocket::ResponsesWsEvent;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use litellm_auth::{CredentialPlacement, SecretValue};
|
||||
|
||||
use crate::{
|
||||
Error,
|
||||
base_llm::responses::transformation::{
|
||||
ResponsesWebSocketProviderConfig, ResponsesWsTransformResult, enforce_model,
|
||||
base_llm::{
|
||||
auth::{AuthScheme, ValidatedEnvironment},
|
||||
responses::transformation::{
|
||||
BaseResponsesApiConfig, ResponsesWebSocketProviderConfig, ResponsesWsTransformResult,
|
||||
enforce_model,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
|
|
@ -92,6 +100,98 @@ fn percent_encode(value: &str) -> String {
|
|||
.collect()
|
||||
}
|
||||
|
||||
impl BaseResponsesApiConfig for OpenAiResponsesApiConfig {
|
||||
fn secret_names(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
api_base: Option<&str>,
|
||||
) -> &'static [&'static str] {
|
||||
match (
|
||||
api_key.is_some_and(|key| !key.is_empty()),
|
||||
api_base.is_some_and(|base| !base.is_empty()),
|
||||
) {
|
||||
(true, true) => &[],
|
||||
(true, false) => &["OPENAI_BASE_URL", "OPENAI_API_BASE"],
|
||||
(false, true) => &["OPENAI_API_KEY"],
|
||||
(false, false) => &["OPENAI_API_KEY", "OPENAI_BASE_URL", "OPENAI_API_BASE"],
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
&self,
|
||||
headers: Vec<(String, String)>,
|
||||
api_key: Option<&str>,
|
||||
lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<ValidatedEnvironment, Error> {
|
||||
let key = api_key
|
||||
.filter(|key| !key.is_empty())
|
||||
.map(str::to_owned)
|
||||
.or_else(|| lookup("OPENAI_API_KEY"))
|
||||
.filter(|key| !key.is_empty())
|
||||
.ok_or(litellm_auth::Error::MissingApiKey {
|
||||
provider: "OpenAI",
|
||||
environment_variable: "OPENAI_API_KEY",
|
||||
})?;
|
||||
Ok(ValidatedEnvironment {
|
||||
headers: litellm_http::request::with_default_headers(
|
||||
headers,
|
||||
&[("content-type", "application/json")],
|
||||
),
|
||||
auth: AuthScheme::Credential {
|
||||
placement: CredentialPlacement::Bearer,
|
||||
secret: SecretValue::new(key),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
let base = api_base
|
||||
.filter(|base| !base.is_empty())
|
||||
.map(str::to_owned)
|
||||
.or_else(|| lookup("OPENAI_BASE_URL"))
|
||||
.or_else(|| lookup("OPENAI_API_BASE"))
|
||||
.unwrap_or_else(|| OPENAI_RESPONSES_DEFAULT_API_BASE.into());
|
||||
format!("{}/responses", base.trim_end_matches('/'))
|
||||
}
|
||||
|
||||
fn transform_responses_api_request(
|
||||
&self,
|
||||
model: &str,
|
||||
input: Value,
|
||||
params: Map<String, Value>,
|
||||
) -> Result<Value, Error> {
|
||||
if !input.is_string() && !input.is_array() {
|
||||
return Err(Error::InvalidRequest(
|
||||
"responses input must be a string or an array".into(),
|
||||
));
|
||||
}
|
||||
if params
|
||||
.get("stream")
|
||||
.is_some_and(|stream| !stream.is_boolean())
|
||||
{
|
||||
return Err(Error::InvalidRequest("stream must be a boolean".into()));
|
||||
}
|
||||
Ok(Value::Object(
|
||||
params
|
||||
.into_iter()
|
||||
.chain([
|
||||
("model".into(), Value::String(model.into())),
|
||||
("input".into(), input),
|
||||
])
|
||||
.collect(),
|
||||
))
|
||||
}
|
||||
|
||||
fn transform_response_api_response(&self, body: Value) -> Result<ResponsesApiResponse, Error> {
|
||||
serde_json::from_value(body)
|
||||
.map_err(|error| Error::InvalidResponse(error.to_string().into()))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
|
|||
11
litellm-rust/crates/types/src/responses/main.rs
Normal file
11
litellm-rust/crates/types/src/responses/main.rs
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ResponsesApiResponse {
|
||||
pub id: String,
|
||||
pub model: String,
|
||||
pub output: Vec<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra: Map<String, Value>,
|
||||
}
|
||||
|
|
@ -1 +1,2 @@
|
|||
pub mod main;
|
||||
pub mod streaming_websocket;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue