use super::body::{HyperResponseBody, StreamBody};
use super::handle::{
ConnCtx, MethodNotAllowed, method_not_allowed_response, refuse_head, run_head_gate,
to_hyper_full,
};
use super::record::record_request;
use super::request::{RequestOrigin, method_is_head};
use super::response::HeaderPair;
use super::server_lifecycle::ConnectionLifecycle;
use super::sse::SseWriter;
use super::{Request, Response};
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,
is_head: bool,
) -> hyper::Response<HyperResponseBody> {
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 is_head {
true => StreamBody::Drained,
false => body,
};
streaming_response_or_empty(builder, body, hyper::StatusCode::BAD_GATEWAY)
}
fn finish_upstream_stream(
forwarded: Result<super::async_proxy::StreamingProxyResponse, crate::RuntimeError>,
ctx: &ConnCtx,
method: &'static str,
path: &str,
is_head: bool,
start: std::time::Instant,
) -> hyper::Response<HyperResponseBody> {
let upstream = match forwarded {
Ok(upstream) => upstream,
Err(e) => {
tracing::warn!(error = %e, "streaming proxy upstream failed");
record_request(ctx, method, path, 502, start);
return to_hyper_full(Response::text_raw(502, "proxy upstream failed"));
}
};
let response = build_streaming_response(
upstream.status,
&upstream.headers,
StreamBody::Proxy(upstream.rx),
is_head,
);
record_request(ctx, method, path, response.status().as_u16(), start);
response
}
pub(super) fn handle_stream_response(
stream_resp: super::stream::StreamResponse,
req: Request,
ctx: &ConnCtx,
start: std::time::Instant,
) -> Result<hyper::Response<HyperResponseBody>, std::convert::Infallible> {
let is_head = req.is_head();
let parts = stream_resp.into_parts();
let response = build_streaming_response(
parts.status,
&parts.headers,
StreamBody::Channel(parts.rx),
is_head,
);
record_request(
ctx,
req.method(),
req.path(),
response.status().as_u16(),
start,
);
Ok(response)
}
pub(super) async fn handle_proxy_stream_response(
req: Request,
backend: &str,
prefix: &str,
ctx: &ConnCtx,
start: std::time::Instant,
) -> Result<hyper::Response<HyperResponseBody>, std::convert::Infallible> {
let proxy_req = super::async_proxy::ProxyRequest::from_request(&req);
let is_head = req.is_head();
let forwarded = super::async_proxy::forward_request_streaming(proxy_req, backend, prefix).await;
Ok(finish_upstream_stream(
forwarded,
ctx,
req.method(),
req.path(),
is_head,
start,
))
}
pub(super) async fn dispatch_streaming_proxy(
hyper_req: hyper::Request<hyper::body::Incoming>,
dispatch: &super::router::ServerDispatch,
ctx: &ConnCtx,
target: super::dispatch::StreamingProxyTarget,
origin: RequestOrigin<'_>,
method: super::method::Method,
start: std::time::Instant,
) -> Result<hyper::Response<HyperResponseBody>, std::convert::Infallible> {
let super::dispatch::StreamingProxyTarget {
backend,
prefix,
params,
} = target;
let method_str = method.as_str();
let is_head = method_is_head(method);
let gate_blocked = match run_head_gate(&hyper_req, dispatch, origin, Some(params)).await {
Ok(blocked) => blocked,
Err(MethodNotAllowed) => {
return Ok(refuse_head(
ctx,
Some(method),
hyper_req.uri().path(),
method_not_allowed_response(),
start,
));
}
};
if let Some(blocked) = gate_blocked {
return Ok(refuse_head(
ctx,
Some(method),
hyper_req.uri().path(),
blocked,
start,
));
}
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,
method_str,
hyper_parts.uri.path(),
is_head,
start,
))
}