diff --git a/litellm-rust/crates/core/src/responses/handler.rs b/litellm-rust/crates/core/src/responses/handler.rs new file mode 100644 index 00000000000..ad85a750b4c --- /dev/null +++ b/litellm-rust/crates/core/src/responses/handler.rs @@ -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, +) -> Result { + 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() +} diff --git a/litellm-rust/crates/core/src/responses/mod.rs b/litellm-rust/crates/core/src/responses/mod.rs index 464a81fe89c..35c79a7e801 100644 --- a/litellm-rust/crates/core/src/responses/mod.rs +++ b/litellm-rust/crates/core/src/responses/mod.rs @@ -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, + secrets: Arc, +} + +impl ResponsesRoute { + pub fn new( + http: litellm_http::Client, + auth: Arc, + secrets: Arc, + ) -> Self { + Self { + http, + auth, + secrets, + } + } + + pub async fn execute( + &self, + call: ResponsesCall, + hooks: &impl RouteHooks, + ) -> Result { + litellm_host::lifecycle::observe_call(hooks.observer(), self.run(call, hooks)).await + } + + async fn run( + &self, + call: ResponsesCall, + hooks: &impl RouteHooks, + ) -> Result { + let request = prepare::prepare(call, self.secrets.as_ref()).await?; + handler::execute(&self.http, &self.auth, request, hooks).await + } +} diff --git a/litellm-rust/crates/core/src/responses/prepare.rs b/litellm-rust/crates/core/src/responses/prepare.rs new file mode 100644 index 00000000000..6e00ecd577b --- /dev/null +++ b/litellm-rust/crates/core/src/responses/prepare.rs @@ -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 { + 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, + }) +} diff --git a/litellm-rust/crates/core/src/responses/route.rs b/litellm-rust/crates/core/src/responses/route.rs new file mode 100644 index 00000000000..a2d2788d519 --- /dev/null +++ b/litellm-rust/crates/core/src/responses/route.rs @@ -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 { + hosted_call(call, move |call, _, hooks| async move { + self.run(call, &hooks).await + }) + } +} diff --git a/litellm-rust/crates/core/src/responses/types.rs b/litellm-rust/crates/core/src/responses/types.rs new file mode 100644 index 00000000000..ee93f28ac3a --- /dev/null +++ b/litellm-rust/crates/core/src/responses/types.rs @@ -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, + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, + pub extra_headers: Option>, + pub timeout: Option, +} + +pub struct ResponsesStreamHead { + pub headers: Vec<(String, String)>, +} + +pub type ResponsesOutput = CallOutput; + +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, +} diff --git a/litellm-rust/crates/core/tests/responses.rs b/litellm-rust/crates/core/tests/responses.rs new file mode 100644 index 00000000000..a2652fdd610 --- /dev/null +++ b/litellm-rust/crates/core/tests/responses.rs @@ -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::::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::::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::>().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::::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()); +} diff --git a/litellm-rust/crates/core/tests/support/mod.rs b/litellm-rust/crates/core/tests/support/mod.rs index d45a537b19d..baf21d2970a 100644 --- a/litellm-rust/crates/core/tests/support/mod.rs +++ b/litellm-rust/crates/core/tests/support/mod.rs @@ -57,6 +57,15 @@ pub fn chat_completions_route() -> litellm_core::chat_completions::ChatCompletio ) } +pub fn responses_route(secrets: Arc) -> 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( diff --git a/litellm-rust/crates/gateway-inference/src/lib.rs b/litellm-rust/crates/gateway-inference/src/lib.rs index a0ef7563b8a..67d00c20b02 100644 --- a/litellm-rust/crates/gateway-inference/src/lib.rs +++ b/litellm-rust/crates/gateway-inference/src/lib.rs @@ -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, 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) -> 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)) diff --git a/litellm-rust/crates/gateway-inference/src/responses.rs b/litellm-rust/crates/gateway-inference/src/responses.rs new file mode 100644 index 00000000000..97040ce6656 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/src/responses.rs @@ -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>, + JsonObject(body): JsonObject, +) -> Result { + 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::::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?) +} diff --git a/litellm-rust/crates/gateway-inference/tests/responses.rs b/litellm-rust/crates/gateway-inference/tests/responses.rs new file mode 100644 index 00000000000..68b40a16472 --- /dev/null +++ b/litellm-rust/crates/gateway-inference/tests/responses.rs @@ -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") + ); +} diff --git a/litellm-rust/crates/gateway-inference/tests/routes.rs b/litellm-rust/crates/gateway-inference/tests/routes.rs index 48f73358ce4..69450255086 100644 --- a/litellm-rust/crates/gateway-inference/tests/routes.rs +++ b/litellm-rust/crates/gateway-inference/tests/routes.rs @@ -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)] diff --git a/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs b/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs index 31eaca932ab..3263672edee 100644 --- a/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/responses/transformation.rs @@ -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, + ) -> Result; + + fn get_complete_url( + &self, + api_base: Option<&str>, + lookup: &dyn Fn(&str) -> Option, + ) -> String; + + fn transform_responses_api_request( + &self, + model: &str, + input: Value, + params: Map, + ) -> Result; + + fn transform_response_api_response(&self, body: Value) -> Result; +} #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct ResponsesWsTransformResult { diff --git a/litellm-rust/crates/llms/src/openai/responses/transformation.rs b/litellm-rust/crates/llms/src/openai/responses/transformation.rs index 1aa172ab967..ecb5f2f3a65 100644 --- a/litellm-rust/crates/llms/src/openai/responses/transformation.rs +++ b/litellm-rust/crates/llms/src/openai/responses/transformation.rs @@ -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, + ) -> Result { + 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 { + 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, + ) -> Result { + 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 { + serde_json::from_value(body) + .map_err(|error| Error::InvalidResponse(error.to_string().into())) + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/litellm-rust/crates/types/src/responses/main.rs b/litellm-rust/crates/types/src/responses/main.rs new file mode 100644 index 00000000000..548dcd8d75e --- /dev/null +++ b/litellm-rust/crates/types/src/responses/main.rs @@ -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, + #[serde(flatten)] + pub extra: Map, +} diff --git a/litellm-rust/crates/types/src/responses/mod.rs b/litellm-rust/crates/types/src/responses/mod.rs index 02493c5f6ed..578373421e6 100644 --- a/litellm-rust/crates/types/src/responses/mod.rs +++ b/litellm-rust/crates/types/src/responses/mod.rs @@ -1 +1,2 @@ +pub mod main; pub mod streaming_websocket;