use bytes::Bytes;
use futures::StreamExt;
use super::{Exchange, Opened, Opening, Transport};
use crate::error::ProviderError;
use crate::http_client::framing::{Framing, NdjsonFramer, SseFramer};
use crate::http_client::{self, HttpClientExt};
use crate::observe::{AdapterContext, AdapterErrorBoundary, AdapterSlot};
use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
use crate::wire::{Body, Encoded, Mode, Projector, Wire, WireFrame};
impl<W, H> Transport<W> for H
where
W: Wire<Payload = Encoded, Frame = WireFrame>,
H: HttpClientExt + Clone + WasmCompatSend + WasmCompatSync + 'static,
{
fn send(&self, payload: Encoded, exchange: Exchange) -> Opening<WireFrame> {
let Encoded {
mut request,
framing,
request_id_header,
relaxed_content_type,
route,
project,
analysis_only,
} = payload;
let Exchange { mode, observation } = exchange;
let streamed = mode == Mode::Streaming && framing != Framing::Whole;
if streamed && matches!(request.body(), Body::Multipart(_)) {
return Opening::failed(ProviderError::request(
"a multipart request cannot open a streamed reply",
));
}
accept_header(&mut request, framing);
let path = request.uri().path().to_owned();
let declared = route.map_or_else(|| path.clone(), str::to_owned);
let exchange = HttpExchange {
framing,
request_id_header,
relaxed_content_type,
path,
observation: observation.map(|context| (context, AdapterSlot::default())),
project,
};
let http = self.clone();
Opening::new(async move {
exchange.install(&request, &declared);
let slot = exchange.slot().cloned();
let mut opened = if streamed {
match byte_request(request) {
Ok(request) => exchange.streaming(&http, request).await,
Err(error) => Opened::failed(error),
}
} else {
exchange.unary(&http, request).await
};
opened.slot = slot;
opened.analysis_only = analysis_only;
Ok(opened)
})
}
}
struct HttpExchange {
framing: Framing,
request_id_header: Option<&'static str>,
relaxed_content_type: bool,
path: String,
observation: Option<(AdapterContext, AdapterSlot)>,
project: Option<Projector>,
}
impl HttpExchange {
fn slot(&self) -> Option<&AdapterSlot> {
self.observation.as_ref().map(|(_, slot)| slot)
}
fn project(&self, payload: &[u8]) {
project(self.slot(), self.project, payload);
}
fn install<B>(&self, request: &http::Request<B>, declared: &str) {
if let Some((context, slot)) = &self.observation {
slot.install(context.attempt_for(request, declared));
}
}
fn failed(self, error: ProviderError, request_id: Option<String>) -> Opened<WireFrame> {
Opened::failed(error)
.with_request_id(request_id)
.with_route(self.path)
}
async fn unary<H: HttpClientExt>(
self,
http: &H,
request: http::Request<Body>,
) -> Opened<WireFrame> {
let sent = match send(http, request, self.request_id_header, self.slot()).await {
Ok(sent) => sent,
Err(error) => {
if let Some(body) = error.provider_response_body() {
self.project(body.as_bytes());
}
return self.failed(error, None);
}
};
if let Some(rejected) =
wrong_content_type(&sent.headers, self.framing, self.relaxed_content_type)
{
let error = ProviderError::from_transport_error(rejected)
.with_provider_status(Some(sent.status))
.with_provider_request_id(sent.provider_request_id.clone())
.with_response_headers(Some(sent.headers.clone()));
self.project(&sent.body);
return self.failed(error, sent.provider_request_id);
}
let document = serde_json::from_slice(&sent.body).ok();
let mut framer = Framer::new(self.framing);
let payloads: Vec<Framed> = framer
.push(&sent.body)
.into_iter()
.chain(framer.finish())
.collect();
let slot = self.slot().cloned();
let projector = self.project;
let frames = futures::stream::iter(payloads).filter_map(move |payload| {
project(slot.as_ref(), projector, payload.payload());
futures::future::ready(payload.into_frame().map(Ok))
});
Opened {
document,
..Opened::new(frames)
.with_request_id(sent.provider_request_id)
.with_http(sent.status, sent.headers)
.with_route(self.path)
}
}
async fn streaming<H: HttpClientExt>(
self,
http: &H,
request: http::Request<Vec<u8>>,
) -> Opened<WireFrame> {
let response = match http.send_streaming(request).await {
Ok(response) if response.status() != http::StatusCode::OK => {
Err(reject_response(response).await)
}
Ok(response) => {
match wrong_content_type(
response.headers(),
self.framing,
self.relaxed_content_type,
) {
Some(error) => Err(error),
None => Ok(response),
}
}
other => other,
};
let response = match response {
Ok(response) => response,
Err(error) => {
if let Some(slot) = self.slot() {
slot.error_boundary(AdapterErrorBoundary::from_http(&error));
if let Some(status) = error.non_success_status() {
slot.response_with_headers(status, error.non_success_headers());
}
if let Some(body) = error.non_success_body() {
self.project(body.as_bytes());
}
}
let request_id = error
.non_success_headers()
.and_then(|headers| request_id_from(headers, self.request_id_header));
let error = ProviderError::from_transport_error(error)
.with_provider_request_id(request_id.clone());
return self.failed(error, request_id);
}
};
if let Some(slot) = self.slot() {
slot.response_with_headers(response.status(), Some(response.headers()));
}
let request_id = request_id_from(response.headers(), self.request_id_header);
let status = response.status();
let headers = response.headers().clone();
let slot = self.slot().cloned();
let Self {
framing,
path,
project: projector,
..
} = self;
let mut body = response.into_body();
let frames = async_stream::stream! {
let mut framer = Framer::new(framing);
while let Some(chunk) = body.next().await {
let chunk = match chunk {
Ok(chunk) => chunk,
Err(error) => {
if let Some(slot) = &slot {
slot.error_boundary(AdapterErrorBoundary::Transport);
}
yield Err(ProviderError::from_transport_error(error));
return;
}
};
if let Some(slot) = &slot {
slot.bytes(&chunk);
}
for payload in framer.push(&chunk) {
project(slot.as_ref(), projector, payload.payload());
if let Some(frame) = payload.into_frame() {
yield Ok(frame);
}
}
}
for payload in framer.finish() {
project(slot.as_ref(), projector, payload.payload());
if let Some(frame) = payload.into_frame() {
yield Ok(frame);
}
}
};
Opened::new(frames)
.with_request_id(request_id)
.with_http(status, headers)
.with_route(path)
}
}
fn project(slot: Option<&AdapterSlot>, project: Option<Projector>, payload: &[u8]) {
if let (Some(slot), Some(project)) = (slot, project) {
slot.project(|sink| project(payload, sink));
}
}
struct Sent {
status: http::StatusCode,
headers: http::HeaderMap,
body: Bytes,
provider_request_id: Option<String>,
}
async fn send<H>(
http: &H,
request: http::Request<Body>,
request_id_header: Option<&'static str>,
observation: Option<&AdapterSlot>,
) -> Result<Sent, ProviderError>
where
H: HttpClientExt,
{
let (parts, body) = request.into_parts();
let response = match body {
Body::Bytes(bytes) => {
http.send::<_, Bytes>(http::Request::from_parts(parts, bytes))
.await
}
Body::Multipart(form) => {
http.send_multipart::<Bytes>(http::Request::from_parts(parts, form))
.await
}
};
let response = match response {
Ok(response) => response,
Err(error) => {
if let Some(observation) = observation
&& let Some(status) = error.non_success_status()
{
observation.response_with_headers(status, error.non_success_headers());
}
let request_id = error
.non_success_headers()
.and_then(|headers| request_id_from(headers, request_id_header));
return Err(
ProviderError::from_transport_error(error).with_provider_request_id(request_id)
);
}
};
let (parts, body) = response.into_parts();
let status = parts.status;
if let Some(observation) = observation {
observation.response_with_headers(status, Some(&parts.headers));
}
let provider_request_id = request_id_from(&parts.headers, request_id_header);
let body = body.await.map_err(ProviderError::from_transport_error)?;
if !status.is_success() {
return Err(
ProviderError::from_http_response(status, String::from_utf8_lossy(&body))
.with_provider_request_id(provider_request_id)
.with_response_headers(Some(parts.headers)),
);
}
Ok(Sent {
status,
headers: parts.headers,
body,
provider_request_id,
})
}
pub(crate) fn request_id_from(headers: &http::HeaderMap, header: Option<&str>) -> Option<String> {
header.and_then(|header| {
headers
.get(header)
.and_then(|value| value.to_str().ok())
.filter(|value| !value.is_empty())
.map(str::to_string)
})
}
fn content_type(request: &mut http::Request<Body>) {
if matches!(request.body(), Body::Bytes(_)) {
request
.headers_mut()
.entry(http::header::CONTENT_TYPE)
.or_insert(http::HeaderValue::from_static("application/json"));
}
}
fn accept_header(request: &mut http::Request<Body>, framing: Framing) {
content_type(request);
if framing == Framing::Sse {
request
.headers_mut()
.entry("Accept")
.or_insert(http::HeaderValue::from_static("text/event-stream"));
}
}
fn wrong_content_type(
headers: &http::HeaderMap,
framing: Framing,
relaxed: bool,
) -> Option<http_client::Error> {
if framing != Framing::Sse {
return None;
}
let Some(content_type) = headers.get(&http::header::CONTENT_TYPE) else {
return (!relaxed)
.then(|| http_client::Error::InvalidContentType(http::HeaderValue::from_static("")));
};
let event_stream = content_type
.to_str()
.ok()
.and_then(|value| value.parse::<mime::Mime>().ok())
.is_some_and(|mime_type| {
matches!(
(mime_type.type_(), mime_type.subtype()),
(mime::TEXT, mime::EVENT_STREAM)
)
});
(!event_stream).then(|| http_client::Error::InvalidContentType(content_type.clone()))
}
fn byte_request(request: http::Request<Body>) -> Result<http::Request<Vec<u8>>, ProviderError> {
let (parts, body) = request.into_parts();
match body {
Body::Bytes(bytes) => Ok(http::Request::from_parts(parts, bytes)),
Body::Multipart(_) => Err(ProviderError::request(
"a multipart request cannot open a streamed reply",
)),
}
}
struct Framed {
payload: Vec<u8>,
frame: bool,
}
impl Framed {
fn payload(&self) -> &[u8] {
&self.payload
}
fn into_frame(self) -> Option<WireFrame> {
self.frame.then(|| match String::from_utf8(self.payload) {
Ok(text) => WireFrame::Text(text),
Err(error) => WireFrame::Bytes(error.into_bytes()),
})
}
}
enum Framer {
Sse(SseFramer),
Ndjson(NdjsonFramer),
Whole(Vec<u8>),
}
impl Framer {
fn new(framing: Framing) -> Self {
match framing {
Framing::Sse => Self::Sse(SseFramer::new()),
Framing::Ndjson => Self::Ndjson(NdjsonFramer::new()),
Framing::Whole => Self::Whole(Vec::new()),
}
}
fn push(&mut self, chunk: &[u8]) -> Vec<Framed> {
match self {
Self::Sse(framer) => framer
.push(chunk)
.map(|event| Framed {
frame: !event.data.trim().is_empty(),
payload: event.data.into_bytes(),
})
.collect(),
Self::Ndjson(framer) => framer
.push(chunk)
.map(|line| Framed {
frame: true,
payload: line,
})
.collect(),
Self::Whole(buffer) => {
buffer.extend_from_slice(chunk);
Vec::new()
}
}
}
fn finish(&mut self) -> Vec<Framed> {
match self {
Self::Sse(_) => Vec::new(),
Self::Ndjson(framer) => framer
.finish()
.map(|line| Framed {
frame: true,
payload: line,
})
.into_iter()
.collect(),
Self::Whole(buffer) => {
let payload = std::mem::take(buffer);
if payload.is_empty() {
Vec::new()
} else {
vec![Framed {
frame: true,
payload,
}]
}
}
}
}
}
const REJECTED_BODY_LIMIT: usize = 1 << 20;
const REJECTED_CHUNK_LIMIT: usize = 4096;
async fn reject_response(
response: http::Response<crate::http_client::BoxedStream>,
) -> http_client::Error {
let status = response.status();
let headers = response.headers().clone();
let mut body = response.into_body();
let mut bytes: Vec<u8> = Vec::new();
let mut chunks = 0usize;
while let Some(chunk) = body.next().await {
chunks += 1;
if let Ok(chunk) = chunk {
let room = REJECTED_BODY_LIMIT.saturating_sub(bytes.len());
bytes.extend_from_slice(chunk.get(..chunk.len().min(room)).unwrap_or_default());
}
if bytes.len() >= REJECTED_BODY_LIMIT || chunks >= REJECTED_CHUNK_LIMIT {
break;
}
}
http_client::Error::InvalidStatusCodeWithDetails {
status,
body: String::from_utf8_lossy(&bytes).into_owned(),
headers,
}
}