diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index 6832a59c7bd..bba06149c6f 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -86,6 +86,7 @@ pub(super) async fn execute( provider_name, )); } + let headers = response.headers().clone(); let text = response.text().await.map_err(network)?; debug!(body = text.as_str(), "provider response body"); hooks @@ -93,8 +94,10 @@ pub(super) async fn execute( raw: RawResponse { body: text.clone() }, }) .await?; - decode_response(config, &body.model, &text) - .map(|message| MessagesResponse::Message(Box::new(message))) + decode_response(config, &body.model, &text).map(|message| MessagesResponse::Message { + headers, + message: Box::new(message), + }) } fn serialize_failure(err: serde_json::Error) -> Error { @@ -154,11 +157,7 @@ fn streaming_response( decoder: Option, provider: &'static str, ) -> MessagesResponse { - let headers = response - .headers() - .iter() - .filter_map(|(name, value)| Some((name.to_string(), value.to_str().ok()?.to_string()))) - .collect(); + let headers = response.headers().clone(); let chunks = match decoder { None => futures_util::stream::try_unfold(response, move |mut response| async move { let chunk = response.chunk().await.map_err(network)?; diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 5e5bbc927b6..adadfb5dc34 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -17,14 +17,17 @@ use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessag use super::{Error, MessagesCall, MessagesResponse, handler::execute, prepare::prepare}; pub enum MessagesOutput { - Message(Box), + Message { + headers: reqwest::header::HeaderMap, + message: Box, + }, /// Every chunk already reached the host through `Deliver`. Streamed, } /// The upstream response as the caller sees it at stream hand-off, before any chunk. pub struct MessagesStreamHead { - pub headers: Vec<(String, String)>, + pub headers: reqwest::header::HeaderMap, } pub struct Messages; @@ -101,7 +104,9 @@ async fn drive( let call = host.project().await?; let request = prepare(call, secrets.as_ref()).await?; match execute(&http, &auth, request, &host).await? { - MessagesResponse::Message(message) => Ok(MessagesOutput::Message(message)), + MessagesResponse::Message { headers, message } => { + Ok(MessagesOutput::Message { headers, message }) + } MessagesResponse::Stream { headers, mut chunks, diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index b09cb96a919..7fd00e1ba62 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -9,6 +9,7 @@ use litellm_types::{ }, utils::ProviderSpecificHeaders, }; +use reqwest::header::HeaderMap; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; @@ -34,9 +35,12 @@ pub(super) fn invalid_request(err: serde_json::Error) -> Error { } pub enum MessagesResponse { - Message(Box), + Message { + headers: HeaderMap, + message: Box, + }, Stream { - headers: Vec<(String, String)>, + headers: HeaderMap, chunks: BoxStream<'static, Result>, }, } diff --git a/litellm-rust/crates/core/tests/messages/host.rs b/litellm-rust/crates/core/tests/messages/host.rs index b19ecf11f09..2048b9666a2 100644 --- a/litellm-rust/crates/core/tests/messages/host.rs +++ b/litellm-rust/crates/core/tests/messages/host.rs @@ -146,7 +146,7 @@ async fn the_raw_upstream_text_is_emitted_once_for_a_message(call: MessagesCall) let output = run_through(&host).await.expect("messages call succeeds"); - assert!(matches!(output, MessagesOutput::Message(_))); + assert!(matches!(output, MessagesOutput::Message { .. })); let [emitted] = <[String; 1]>::try_from(host.raw_responses()) .unwrap_or_else(|raws| panic!("expected one raw response, got {}", raws.len())); assert_eq!(serde_json::from_str::(&emitted).unwrap(), raw); diff --git a/litellm-rust/crates/core/tests/messages/main.rs b/litellm-rust/crates/core/tests/messages/main.rs index 534af6d7d06..bde7ab662eb 100644 --- a/litellm-rust/crates/core/tests/messages/main.rs +++ b/litellm-rust/crates/core/tests/messages/main.rs @@ -114,7 +114,7 @@ async fn run(call: MessagesCall) -> Result { async fn run_message(call: MessagesCall) -> AnthropicMessagesResponse { match run(call).await.expect("messages call succeeds") { - MessagesOutput::Message(message) => *message, + MessagesOutput::Message { message, .. } => *message, MessagesOutput::Streamed => panic!("a non-streaming call returned a stream"), } } diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 14c8eb6b7c6..87b5a587275 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -181,7 +181,8 @@ async fn a_provider_slower_than_the_timeout_fails_the_call(call: MessagesCall) { #[rstest] #[tokio::test] async fn the_facade_sends_through_the_injected_http_pool_configuration(call: MessagesCall) { - let upstream = upstream([message_response()]).await; + let upstream = + upstream([message_response().insert_header("x-upstream-request-id", "req-123")]).await; let base = upstream.uri(); let settings = HttpSettings { user_agent: Some("host-owned/1".into()), @@ -201,10 +202,11 @@ async fn the_facade_sends_through_the_injected_http_pool_configuration(call: Mes .await .expect("messages request succeeds"); - let MessagesResponse::Message(message) = response else { + let MessagesResponse::Message { headers, message } = response else { panic!("a non-streaming request returns a message"); }; assert_eq!(message.id, "msg_1"); + assert_eq!(headers.get("x-upstream-request-id").unwrap(), "req-123"); let sent = only_request(&upstream).await; assert_eq!(sent.header("x-api-key"), Some("sk-ant")); assert_eq!(sent.header("user-agent"), Some("host-owned/1")); diff --git a/litellm-rust/crates/core/tests/messages/secrets.rs b/litellm-rust/crates/core/tests/messages/secrets.rs index 55e510d00d3..ab378149621 100644 --- a/litellm-rust/crates/core/tests/messages/secrets.rs +++ b/litellm-rust/crates/core/tests/messages/secrets.rs @@ -43,7 +43,7 @@ async fn the_credential_and_base_come_from_the_secret_source( .await .expect("messages call succeeds"); - assert!(matches!(output, MessagesOutput::Message(_))); + assert!(matches!(output, MessagesOutput::Message { .. })); let request = only_request(&upstream).await; assert_eq!(request.url.path(), path); assert_eq!(request.header("x-api-key"), Some("sk-from-manager")); diff --git a/litellm-rust/crates/core/tests/messages/stream.rs b/litellm-rust/crates/core/tests/messages/stream.rs index f0e55eca8dd..8e455f3143d 100644 --- a/litellm-rust/crates/core/tests/messages/stream.rs +++ b/litellm-rust/crates/core/tests/messages/stream.rs @@ -28,7 +28,7 @@ const UPSTREAM_HEADERS: [(&str, &str); 2] = [ const SSE_BODY: &str = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; enum Seen { - Open(Vec<(String, String)>), + Open(reqwest::header::HeaderMap), Deliver(Bytes), } @@ -128,9 +128,9 @@ async fn upstream_headers_are_on_the_stream_head_before_the_first_chunk(call: Me .filter(|(name, _)| { UPSTREAM_HEADERS .iter() - .any(|(upstream, _)| upstream == name) + .any(|(upstream, _)| upstream == &name.as_str()) }) - .map(|(name, value)| (name.as_str(), value.as_str())) + .map(|(name, value)| (name.as_str(), value.to_str().unwrap())) .collect(); assert_eq!(surfaced, UPSTREAM_HEADERS); let delivered: Vec = chunks @@ -321,7 +321,10 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte( panic!("a streaming request returns a stream"); }; for (name, value) in UPSTREAM_HEADERS { - assert!(headers.contains(&(name.into(), value.into()))); + assert_eq!( + headers.get(name).and_then(|value| value.to_str().ok()), + Some(value) + ); } let delivered = chunks.try_collect::>().await.unwrap().concat(); assert_eq!(delivered, SSE_BODY.as_bytes()); diff --git a/litellm-rust/crates/gateway-inference/src/messages/mod.rs b/litellm-rust/crates/gateway-inference/src/messages/mod.rs index 5d6a8faa0e8..683a6f35811 100644 --- a/litellm-rust/crates/gateway-inference/src/messages/mod.rs +++ b/litellm-rust/crates/gateway-inference/src/messages/mod.rs @@ -60,8 +60,10 @@ async fn handle(gateway: &Gateway, headers: &HeaderMap, body: &[u8]) -> Result Ok(Json(message).into_response()), - MessagesResponse::Stream { chunks, .. } => Ok(stream(chunks)), + MessagesResponse::Message { headers, message } => { + Ok((headers, Json(message)).into_response()) + } + MessagesResponse::Stream { headers, chunks } => Ok(stream(headers, chunks)), } } @@ -107,16 +109,18 @@ fn anthropic_api_headers(headers: &HeaderMap) -> Option /// A chunk that fails after the stream opened is delivered as an SSE error frame, since /// the status line already went out; the stream ends on it. -fn stream(chunks: BoxStream<'static, Result>) -> Response { +fn stream(headers: HeaderMap, chunks: BoxStream<'static, Result>) -> Response { let body = chunks.map(|chunk| { Ok::<_, Infallible>( chunk.unwrap_or_else(|error| Bytes::from(Error::Route(error).sse_frame())), ) }); - ( - StatusCode::OK, - [(header::CONTENT_TYPE, "text/event-stream")], - Body::from_stream(body), - ) - .into_response() + let mut response_headers = headers; + response_headers.insert( + header::CONTENT_TYPE, + "text/event-stream" + .parse() + .expect("static content type is valid"), + ); + (StatusCode::OK, response_headers, Body::from_stream(body)).into_response() } diff --git a/litellm-rust/crates/gateway-inference/tests/messages.rs b/litellm-rust/crates/gateway-inference/tests/messages.rs index 30836498d49..ecfc2f333f7 100644 --- a/litellm-rust/crates/gateway-inference/tests/messages.rs +++ b/litellm-rust/crates/gateway-inference/tests/messages.rs @@ -24,9 +24,13 @@ async fn messages_reaches_the_provider_and_preserves_json_or_sse(#[case] streami "stop_reason": "end_turn", "usage": {"input_tokens": 1, "output_tokens": 1}}); let sse = "event: message_start\ndata: {\"type\":\"message_start\"}\n\nevent: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; let template = if streaming { - ResponseTemplate::new(200).set_body_raw(sse, "text/event-stream") + ResponseTemplate::new(200) + .insert_header("x-upstream-request-id", "req-stream") + .set_body_raw(sse, "text/event-stream") } else { - ResponseTemplate::new(200).set_body_json(message.clone()) + ResponseTemplate::new(200) + .insert_header("x-upstream-request-id", "req-message") + .set_body_json(message.clone()) }; let messages = json!([{"role": "user", "content": "hi"}]); Mock::given(method("POST")).and(path("/v1/messages")) @@ -44,8 +48,10 @@ async fn messages_reaches_the_provider_and_preserves_json_or_sse(#[case] streami assert_eq!(response.status(), 200); if streaming { assert_eq!(response.headers()["content-type"], "text/event-stream"); + assert_eq!(response.headers()["x-upstream-request-id"], "req-stream"); assert_eq!(to_bytes(response.into_body(), 4096).await.unwrap(), sse); } else { + assert_eq!(response.headers()["x-upstream-request-id"], "req-message"); let body = support::json(response).await; assert_eq!(body["content"], message["content"]); assert_eq!(body["usage"], message["usage"]); diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index cf634ab3fc3..70d136d4819 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -241,10 +241,13 @@ impl ProtocolHost for MessagesPythonHost { fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult> { match response { - MessagesOutput::Message(message) => py + MessagesOutput::Message { headers, message } => py .import(ROUTE_HOST_MODULE)? .getattr("response")? - .call1((to_py(py, message.as_ref())?,)) + .call1(( + to_py(py, message.as_ref())?, + to_py(py, &headers_to_pairs(&headers))?, + )) .map(Bound::unbind), MessagesOutput::Streamed => Ok(py.None()), } @@ -253,7 +256,7 @@ impl ProtocolHost for MessagesPythonHost { fn head(&mut self, py: Python<'_>, head: MessagesStreamHead) -> PyResult> { py.import(ROUTE_HOST_MODULE)? .getattr("stream_hidden_params")? - .call1((to_py(py, &head.headers)?,)) + .call1((to_py(py, &headers_to_pairs(&head.headers))?,)) .map(Bound::unbind) } @@ -281,6 +284,18 @@ impl ProtocolHost for MessagesPythonHost { } } +fn headers_to_pairs(headers: &reqwest::header::HeaderMap) -> Vec<(String, String)> { + headers + .iter() + .filter_map(|(name, value)| { + value + .to_str() + .ok() + .map(|value| (name.to_string(), value.to_owned())) + }) + .collect() +} + #[cfg(test)] mod tests { use rstest::rstest; diff --git a/litellm/rust_bridge/messages/route_host.py b/litellm/rust_bridge/messages/route_host.py index caae9916ffa..3a5ea284147 100644 --- a/litellm/rust_bridge/messages/route_host.py +++ b/litellm/rust_bridge/messages/route_host.py @@ -10,6 +10,7 @@ from pydantic import TypeAdapter, ValidationError import litellm from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.llms.anthropic.pass_through.utils import is_reasoning_auto_summary_enabled +from litellm.router_utils.add_retry_fallback_headers import _add_headers_to_response from litellm.rust_bridge import failures from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse @@ -47,10 +48,16 @@ class MessagesShaping: additional_drop_params: Sequence[str] -def response(value: Mapping[str, object]) -> AnthropicMessagesResponse: +def response( + value: Mapping[str, object], + headers: Sequence[tuple[str, str]] = (), +) -> AnthropicMessagesResponse: + response_value = dict(value) + if headers: + _add_headers_to_response(response_value, dict(httpx.Headers(list(headers)).items())) return cast( # cast-ok: AnthropicMessagesResponse is a TypedDict over the normalized native payload AnthropicMessagesResponse, - dict(value), # mutable-ok: the public Messages response is a TypedDict the caller may annotate in place + response_value, # mutable-ok: the public Messages response is a TypedDict the caller may annotate in place )