fix(messages): preserve upstream response headers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Yujong Lee 2026-09-27 01:15:56 +00:00
parent 501ef23f4a
commit 29d44131a9
12 changed files with 82 additions and 37 deletions

View file

@ -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<StreamDecoder>,
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)?;

View file

@ -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<AnthropicMessagesResponse>),
Message {
headers: reqwest::header::HeaderMap,
message: Box<AnthropicMessagesResponse>,
},
/// 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,

View file

@ -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<AnthropicMessagesResponse>),
Message {
headers: HeaderMap,
message: Box<AnthropicMessagesResponse>,
},
Stream {
headers: Vec<(String, String)>,
headers: HeaderMap,
chunks: BoxStream<'static, Result<Bytes, Error>>,
},
}

View file

@ -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::<Value>(&emitted).unwrap(), raw);

View file

@ -114,7 +114,7 @@ async fn run(call: MessagesCall) -> Result<MessagesOutput, Error> {
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"),
}
}

View file

@ -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"));

View file

@ -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"));

View file

@ -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<u8> = 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::<Vec<_>>().await.unwrap().concat();
assert_eq!(delivered, SSE_BODY.as_bytes());

View file

@ -60,8 +60,10 @@ async fn handle(gateway: &Gateway, headers: &HeaderMap, body: &[u8]) -> Result<R
)
.await?
{
MessagesResponse::Message(message) => 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<ProviderSpecificHeaders>
/// 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<Bytes, RouteError>>) -> Response {
fn stream(headers: HeaderMap, chunks: BoxStream<'static, Result<Bytes, RouteError>>) -> 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()
}

View file

@ -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"]);

View file

@ -241,10 +241,13 @@ impl ProtocolHost for MessagesPythonHost {
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {
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<PyAny>> {
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;

View file

@ -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
)