use bytes::Bytes;
use serde::de::DeserializeOwned;
use super::envelope::ProviderEnvelope;
use crate::completion::CompletionError;
use crate::http_client::HttpClientExt;
pub(crate) async fn send_completion<C, A, F>(
client: &C,
request: crate::http_client::Request<Vec<u8>>,
label: &str,
request_id_header: Option<&str>,
record_telemetry: F,
) -> Result<(A::Payload, Option<String>), CompletionError>
where
C: HttpClientExt,
A: DeserializeOwned + ProviderEnvelope,
A::Payload: serde::Serialize,
F: FnOnce(&A::Payload),
{
let response = match client.send::<_, Bytes>(request).await {
Ok(response) => response,
Err(crate::http_client::Error::InvalidStatusCodeWithDetails {
status,
body,
headers,
}) => {
return Err(match request_id_header {
Some(header) => {
let provider_request_id = headers
.get(header)
.and_then(|value| value.to_str().ok())
.filter(|value| !value.is_empty())
.map(str::to_string);
CompletionError::from_http_response_with_request_id(
status,
body,
provider_request_id,
)
.with_response_headers(Some(headers))
}
None => CompletionError::HttpError(
crate::http_client::Error::InvalidStatusCodeWithDetails {
status,
body,
headers,
},
),
});
}
Err(crate::http_client::Error::InvalidStatusCodeWithMessage(status, body))
if request_id_header.is_some() =>
{
return Err(CompletionError::from_http_response_with_request_id(
status, body, None,
));
}
Err(other) => return Err(other.into()),
};
let (parts, body) = response.into_parts();
let status = parts.status;
let provider_request_id = request_id_header.and_then(|header| {
parts
.headers
.get(header)
.and_then(|value| value.to_str().ok())
.filter(|value| !value.is_empty())
.map(str::to_string)
});
let response_headers = Some(Box::new(parts.headers));
let body = body.await.map_err(CompletionError::HttpError)?;
if !status.is_success() {
return Err(match request_id_header {
Some(_) => CompletionError::from_http_response_with_request_id(
status,
String::from_utf8_lossy(&body),
provider_request_id,
),
None => CompletionError::from_http_response(status, String::from_utf8_lossy(&body)),
}
.with_response_headers(response_headers));
}
let envelope: A = serde_json::from_slice(&body).map_err(|err| {
tracing::error!(
error = %err,
body = %String::from_utf8_lossy(&body),
"failed to deserialize {label} response"
);
CompletionError::JsonError(err)
})?;
match envelope.into_payload() {
Ok(payload) => {
record_telemetry(&payload);
super::trace_json(
crate::providers::internal::LogTarget::Completions,
&format!("{label} response"),
&payload,
);
Ok((payload, provider_request_id))
}
Err(message) => {
tracing::warn!(message = %message, "provider returned an error response");
Err(match request_id_header {
Some(_) => CompletionError::from_http_response_with_request_id(
status,
String::from_utf8_lossy(&body),
provider_request_id,
),
None => CompletionError::from_http_response(status, String::from_utf8_lossy(&body)),
}
.with_response_headers(response_headers))
}
}
}
#[cfg(test)]
mod header_preservation_tests {
use super::*;
use crate::test_utils::RecordingHttpClient;
const CONTRACT: Option<&str> = Some("x-request-id");
const NO_CONTRACT: Option<&str> = None;
const BODY: &str = r#"{"error":{"message":"rate limited"}}"#;
fn rate_limited_headers() -> http::HeaderMap {
let mut headers = http::HeaderMap::new();
headers.insert(
http::header::RETRY_AFTER,
http::HeaderValue::from_static("20"),
);
headers.insert("x-ratelimit-remaining", http::HeaderValue::from_static("0"));
headers.insert("x-request-id", http::HeaderValue::from_static("req_abc"));
headers
}
struct RejectingEnvelope;
impl<'de> serde::Deserialize<'de> for RejectingEnvelope {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
serde::de::IgnoredAny::deserialize(deserializer)?;
Ok(Self)
}
}
impl super::super::envelope::ProviderEnvelope for RejectingEnvelope {
type Payload = serde_json::Value;
fn into_payload(self) -> Result<Self::Payload, String> {
Err("envelope".to_string())
}
}
async fn drive(
client: RecordingHttpClient,
request_id_header: Option<&str>,
) -> CompletionError {
drive_as::<super::super::envelope::DirectPayload<serde_json::Value>>(
client,
request_id_header,
)
.await
}
async fn drive_as<A>(
client: RecordingHttpClient,
request_id_header: Option<&str>,
) -> CompletionError
where
A: DeserializeOwned + ProviderEnvelope<Payload = serde_json::Value>,
{
let request = crate::http_client::Request::builder()
.method(http::Method::POST)
.uri("https://example.test/v1/chat")
.body(Vec::new())
.expect("valid request");
send_completion::<_, A, _>(&client, request, "test provider", request_id_header, |_| {})
.await
.expect_err("the scripted response is a failure")
}
fn assert_rate_limit_metadata_survived(error: &CompletionError, cell: &str) {
let headers = error
.provider_response_headers()
.unwrap_or_else(|| panic!("{cell}: response headers were dropped by the driver"));
assert_eq!(
headers
.get(http::header::RETRY_AFTER)
.and_then(|value| value.to_str().ok()),
Some("20"),
"{cell}: Retry-After not recoverable",
);
assert_eq!(
headers
.get("x-ratelimit-remaining")
.and_then(|value| value.to_str().ok()),
Some("0"),
"{cell}: x-ratelimit-remaining not recoverable",
);
assert_eq!(
error.provider_response_status(),
Some(http::StatusCode::TOO_MANY_REQUESTS),
"{cell}: status lost",
);
assert_eq!(
error.provider_response_body(),
Some(BODY),
"{cell}: body lost"
);
}
#[tokio::test]
async fn transport_error_preserves_headers_for_a_contract_provider() {
let client = RecordingHttpClient::with_error_headers(
http::StatusCode::TOO_MANY_REQUESTS,
BODY,
rate_limited_headers(),
);
let error = drive(client, CONTRACT).await;
assert_rate_limit_metadata_survived(&error, "transport-error/contract");
assert!(matches!(error, CompletionError::ProviderResponse(_)));
assert_eq!(error.provider_request_id(), Some("req_abc"));
}
#[tokio::test]
async fn transport_error_preserves_headers_for_a_contract_less_provider() {
let client = RecordingHttpClient::with_error_headers(
http::StatusCode::TOO_MANY_REQUESTS,
BODY,
rate_limited_headers(),
);
let error = drive(client, NO_CONTRACT).await;
assert_rate_limit_metadata_survived(&error, "transport-error/contract-less");
assert!(matches!(error, CompletionError::HttpError(_)));
assert_eq!(error.provider_request_id(), None);
}
#[tokio::test]
async fn non_success_response_preserves_headers_for_a_contract_provider() {
let client = RecordingHttpClient::with_error_response_headers(
http::StatusCode::TOO_MANY_REQUESTS,
BODY,
rate_limited_headers(),
);
let error = drive(client, CONTRACT).await;
assert_rate_limit_metadata_survived(&error, "response/contract");
assert!(matches!(error, CompletionError::ProviderResponse(_)));
assert_eq!(error.provider_request_id(), Some("req_abc"));
}
#[tokio::test]
async fn non_success_response_preserves_headers_for_a_contract_less_provider() {
let client = RecordingHttpClient::with_error_response_headers(
http::StatusCode::TOO_MANY_REQUESTS,
BODY,
rate_limited_headers(),
);
let error = drive(client, NO_CONTRACT).await;
assert_rate_limit_metadata_survived(&error, "response/contract-less");
assert!(matches!(error, CompletionError::HttpError(_)));
}
#[tokio::test]
async fn header_less_transport_reports_no_headers() {
for (contract, expect_provider_response) in [(CONTRACT, true), (NO_CONTRACT, false)] {
let client = RecordingHttpClient::with_error(http::StatusCode::TOO_MANY_REQUESTS, BODY);
let error = drive(client, contract).await;
assert!(error.provider_response_headers().is_none());
assert_eq!(
matches!(error, CompletionError::ProviderResponse(_)),
expect_provider_response,
"classification must not depend on header capture",
);
assert_eq!(error.provider_response_body(), Some(BODY));
}
}
#[tokio::test]
async fn success_status_error_envelope_preserves_headers() {
let client = RecordingHttpClient::with_error_response_headers(
http::StatusCode::OK,
r#"{"error":{"message":"envelope"}}"#,
rate_limited_headers(),
);
let error = drive_as::<RejectingEnvelope>(client, CONTRACT).await;
assert!(matches!(error, CompletionError::ProviderResponse(_)));
assert_eq!(
error
.provider_response_headers()
.and_then(|headers| headers.get(http::header::RETRY_AFTER))
.and_then(|value| value.to_str().ok()),
Some("20"),
"a 200-with-error-envelope must expose its Retry-After too",
);
assert_eq!(error.provider_response_status(), Some(http::StatusCode::OK));
assert_eq!(error.provider_request_id(), Some("req_abc"));
}
#[tokio::test]
async fn successful_response_is_unaffected_by_header_capture() {
let client = RecordingHttpClient::with_error_response_headers(
http::StatusCode::OK,
r#"{"ok":true}"#,
rate_limited_headers(),
);
let request = crate::http_client::Request::builder()
.method(http::Method::POST)
.uri("https://example.test/v1/chat")
.body(Vec::new())
.expect("valid request");
let (payload, request_id) = send_completion::<
_,
super::super::envelope::DirectPayload<serde_json::Value>,
_,
>(&client, request, "test provider", CONTRACT, |_| {})
.await
.expect("a 2xx payload should decode");
assert_eq!(payload["ok"], true);
assert_eq!(request_id.as_deref(), Some("req_abc"));
}
}