mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
501ef23f4a
commit
29d44131a9
12 changed files with 82 additions and 37 deletions
|
|
@ -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)?;
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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>>,
|
||||
},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"]);
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue