use super::Request;
use super::body::{HyperResponseBody, StreamBody};
use super::handle::{ConnCtx, answer, answer_rejected, run_head_gate};
use super::record::record_scoped;
use super::rejection::{Rejected, RejectionScope, RequestIdentity};
use super::request::{RequestHead, RequestOrigin};
use super::response::HeaderPair;
use super::server_lifecycle::ConnectionLifecycle;
use super::sse::SseWriter;
pub(super) async fn handle_sse(
handler: super::router::SseHandler,
req: Request,
buffer_size: usize,
lifecycle: &ConnectionLifecycle,
) -> Result<hyper::Response<HyperResponseBody>, std::convert::Infallible> {
super::mock::LifecycleScript::pause_at(
lifecycle.script().as_deref(),
super::mock::LifecycleCheckpoint::SseBufferConfigured(buffer_size),
)
.await;
let body = match req.is_head() {
true => StreamBody::Drained,
false => spawn_sse_producer(handler, req, buffer_size),
};
let builder = hyper::Response::builder()
.status(200)
.header("Content-Type", "text/event-stream")
.header("Cache-Control", "no-cache");
Ok(streaming_response_or_empty(
builder,
body,
hyper::StatusCode::OK,
))
}
fn spawn_sse_producer(
handler: super::router::SseHandler,
req: Request,
buffer_size: usize,
) -> StreamBody {
let (tx, rx) = tokio::sync::mpsc::channel::<bytes::Bytes>(buffer_size);
let uri = req.uri_owned();
crate::task::spawn_internal_blocking("sse", uri.path(), move || {
let mut writer = SseWriter::new(tx);
if let Err(e) = handler(&req, &mut writer) {
tracing::warn!(error = %e, "SSE handler returned error");
}
});
StreamBody::Channel(rx)
}
fn streaming_response_or_empty(
builder: hyper::http::response::Builder,
body: StreamBody,
fallback: hyper::StatusCode,
) -> hyper::Response<HyperResponseBody> {
builder
.body(HyperResponseBody::Streaming(body))
.unwrap_or_else(|err| {
tracing::error!("failed to build streaming response: {err}");
let mut response =
hyper::Response::new(HyperResponseBody::Streaming(StreamBody::Drained));
*response.status_mut() = fallback;
response
})
}
fn build_streaming_response(
status: u16,
headers: &[HeaderPair],
body: StreamBody,
scope: &RejectionScope,
) -> Result<hyper::Response<HyperResponseBody>, hyper::http::Error> {
let mut builder = hyper::Response::builder().status(status);
for (name, value) in headers {
builder = builder.header(name.as_ref(), value.as_ref());
}
let body = match scope.is_head() {
true => StreamBody::Drained,
false => body,
};
builder.body(HyperResponseBody::Streaming(body))
}
fn finish_upstream_stream(
forwarded: Result<super::async_proxy::StreamingProxyResponse, super::async_proxy::ProxyFailure>,
ctx: &ConnCtx,
scope: &RejectionScope,
start: std::time::Instant,
) -> hyper::Response<HyperResponseBody> {
let upstream = match forwarded {
Ok(upstream) => upstream,
Err(failure) => {
return answer_rejected(ctx, scope, Rejected::from_proxy_failure(failure), start);
}
};
let response = match build_streaming_response(
upstream.status,
&upstream.headers,
StreamBody::Proxy(upstream.rx),
scope,
) {
Ok(response) => response,
Err(error) => {
return answer_rejected(ctx, scope, Rejected::unrepresentable(error), start);
}
};
record_scoped(ctx, scope, response.status().as_u16(), start);
response
}
pub(super) fn handle_stream_response(
stream_resp: super::stream::StreamResponse,
held_request: Request,
ctx: &ConnCtx,
scope: &RejectionScope,
start: std::time::Instant,
) -> Result<hyper::Response<HyperResponseBody>, std::convert::Infallible> {
let parts = stream_resp.into_parts();
let response = match build_streaming_response(
parts.status,
&parts.headers,
StreamBody::Channel(parts.rx),
scope,
) {
Ok(response) => response,
Err(error) => {
return Ok(answer_rejected(
ctx,
scope,
Rejected::unrepresentable(error),
start,
));
}
};
record_scoped(ctx, scope, response.status().as_u16(), start);
drop(held_request);
Ok(response)
}
pub(super) async fn handle_proxy_stream_response(
req: Request,
backend: &str,
prefix: &str,
ctx: &ConnCtx,
scope: &RejectionScope,
start: std::time::Instant,
) -> Result<hyper::Response<HyperResponseBody>, std::convert::Infallible> {
let proxy_req = super::async_proxy::ProxyRequest::from_request(&req);
let forwarded = super::async_proxy::forward_request_streaming(proxy_req, backend, prefix).await;
Ok(finish_upstream_stream(forwarded, ctx, scope, start))
}
pub(super) async fn dispatch_streaming_proxy(
hyper_req: hyper::Request<hyper::body::Incoming>,
ctx: &ConnCtx,
target: super::dispatch::StreamingProxyTarget,
origin: RequestOrigin<'_>,
router: Option<&super::dispatch::FrozenRouter>,
pre_body: &super::dispatch::PreBodyScope,
start: std::time::Instant,
) -> Result<hyper::Response<HyperResponseBody>, std::convert::Infallible> {
let super::dispatch::StreamingProxyTarget {
backend,
prefix,
params,
method,
} = target;
let scope = pre_body.scope(RequestIdentity::from_head(
&origin,
hyper_req.method(),
hyper_req.uri(),
));
let head = RequestHead::from_hyper_request(&hyper_req, origin);
if let Some(blocked) = run_head_gate(&head, router, Some(params), &scope).await {
return Ok(answer(ctx, blocked, start, &scope));
}
let scheme = match origin.is_tls {
true => "https",
false => "http",
};
let (hyper_parts, body) = hyper_req.into_parts();
let proxy_parts = super::async_proxy::IncomingProxyParts {
method,
path_and_query: hyper_parts
.uri
.path_and_query()
.map_or("/", |pq| pq.as_str())
.into(),
headers: hyper_parts.headers,
remote_addr: origin.remote_addr,
scheme,
};
let forwarded =
super::async_proxy::forward_incoming_streaming(proxy_parts, body, &backend, &prefix).await;
Ok(finish_upstream_stream(forwarded, ctx, &scope, start))
}