use std::net::SocketAddr;
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime};
use hyper::body::Body;
use hyper::{Response, StatusCode};
use plecto_control::otlp::SpanRecord;
use plecto_control::{
ChainOutcome, ConfigSnapshot, HashInput, HashKeySource, HttpRequest, HttpResponse,
RateLimitDecision, RequestBodyOutcome, RequestTrace, ResponseOutcome, RouteInfo,
};
use tokio::sync::OwnedSemaphorePermit;
use crate::body::{INBOUND_BODY_READ_TIMEOUT, MAX_REQUEST_BODY_BUFFER, buffer_request_body};
use crate::error::ServerError;
use crate::forward::{ForwardBody, ForwardOutcome, ForwardRequest, forward_with_retry};
use crate::headers::{
copy_headers, copy_headers_direct, headers_to_vec, set_forwarded, to_http_request,
};
use crate::respond::{
discard_upstream_body, fault, http_response, stream_response, stream_response_direct, synth,
synth_retry_after, with_error_code,
};
use crate::{ReqBody, ResponseBody, ServerState, access_log};
pub(crate) async fn proxy_core(
state: Arc<ServerState>,
scheme: &'static str,
peer: SocketAddr,
parts: hyper::http::request::Parts,
body: ReqBody,
) -> Result<Response<ResponseBody>, ServerError> {
struct InFlight<'a>(&'a crate::metrics::ServerMetrics);
impl Drop for InFlight<'_> {
fn drop(&mut self) {
self.0.dec_in_flight();
}
}
let start = Instant::now();
state.metrics.inc_in_flight();
let in_flight = InFlight(&state.metrics);
let access = state
.control
.access_log_enabled()
.then(|| access_log::Access {
method: parts.method.as_str().to_string(),
authority: crate::headers::request_authority(&parts),
path: parts.uri.path().to_string(),
});
let trace = parts
.headers
.get("traceparent")
.and_then(|v| v.to_str().ok())
.and_then(RequestTrace::from_traceparent)
.unwrap_or_else(RequestTrace::root);
let otlp_request = state.otlp.as_ref().filter(|_| trace.is_sampled()).map(|_| {
(
parts.method.as_str().to_string(),
parts.uri.path().to_string(),
SystemTime::now(),
)
});
let result = proxy_core_inner(state.clone(), scheme, peer, trace, parts, body).await;
drop(in_flight);
let status = match &result {
Ok(resp) => resp.status().as_u16(),
Err(_) => StatusCode::BAD_GATEWAY.as_u16(),
};
let elapsed = start.elapsed();
state.metrics.record_request(status, elapsed);
if let Some(access) = access {
access_log::record(scheme, peer, &access, status, elapsed);
}
if let (Some(buffer), Some((method, path, started))) = (state.otlp.as_ref(), otlp_request) {
buffer.push(SpanRecord::request_span(
&trace, &method, &path, scheme, status, started, elapsed,
));
}
result
}
async fn proxy_core_inner(
state: Arc<ServerState>,
scheme: &'static str,
peer: SocketAddr,
trace: RequestTrace,
mut parts: hyper::http::request::Parts,
body: ReqBody,
) -> Result<Response<ResponseBody>, ServerError> {
set_forwarded(&mut parts.headers, peer.ip(), scheme);
let mut http_req = to_http_request(&parts, scheme);
let bodyless = body.size_hint().exact() == Some(0);
match plecto_control::normalize_path(&http_req.path) {
Some(std::borrow::Cow::Owned(path)) => http_req.path = path,
Some(std::borrow::Cow::Borrowed(_)) => {}
None => {
return Ok(with_error_code(
synth(
StatusCode::BAD_REQUEST,
&fault::BAD_PATH,
b"bad request path",
),
&plecto_control::PATH_NORMALIZATION_REJECTED,
));
}
}
let snapshot = state.control.snapshot_with_trace(trace);
let Some(route) = snapshot.find_route(&http_req) else {
return Ok(synth(StatusCode::NOT_FOUND, &fault::NO_ROUTE, b"no route"));
};
let idx = route.index;
if let RateLimitDecision::Limit { retry_after_ms } = route.check_rate_limit(peer.ip()) {
state.metrics.inc_rate_limited();
return Ok(with_error_code(
synth_retry_after(
StatusCode::TOO_MANY_REQUESTS,
&fault::RATE_LIMITED,
b"rate limit exceeded",
retry_after_ms.div_ceil(1000),
),
&plecto_control::QUOTA_EXCEEDED,
));
}
let upgrade = upgrade_intent(&route, &mut parts, bodyless);
let chain = ChainRef {
snapshot: &snapshot,
idx,
has_filters: route.has_filters,
};
let forward = if route.has_filters {
let snap_req = snapshot.clone();
match tokio::task::spawn_blocking(move || snap_req.dispatch_request(idx, http_req)).await? {
ChainOutcome::Respond(resp) => return Ok(http_response(resp)),
ChainOutcome::Forward(req) => req,
}
} else {
http_req
};
let Some(group) = route.pick_upstream() else {
return Ok(synth(
StatusCode::SERVICE_UNAVAILABLE,
&fault::NO_HEALTHY_UPSTREAM,
b"no healthy upstream",
));
};
let hash_key: Option<HashInput<'_>> = group.hash_key_source().and_then(|src| match src {
HashKeySource::Header(name) => parts
.headers
.get(name)
.map(|v| HashInput::Bytes(v.as_bytes())),
HashKeySource::SourceIp => Some(HashInput::Ip(peer.ip())),
});
let upstream_path = route.rewrite_path(&forward.path);
let timeout = group.request_timeout();
let overall = group.overall_timeout();
let overall_deadline = (!overall.is_zero()).then(|| Instant::now() + overall);
let mut real_body = if bodyless {
ForwardBody::Bodyless
} else {
ForwardBody::OneShot(body)
};
let _buf_permit =
match request_body_hook(&state, &chain, route.reads_body, &mut real_body).await? {
BodyHookOutcome::Proceed(permit) => permit,
BodyHookOutcome::Respond(resp) => return Ok(resp),
};
let permit = match group.try_acquire() {
Some(permit) => permit,
None => {
state.metrics.inc_circuit_open();
return Ok(synth(
StatusCode::SERVICE_UNAVAILABLE,
&fault::CIRCUIT_OPEN,
b"upstream overloaded",
));
}
};
let Some(pick) = group.pick(hash_key) else {
return Ok(synth(
StatusCode::SERVICE_UNAVAILABLE,
&fault::NO_HEALTHY_UPSTREAM,
b"no healthy upstream",
));
};
let forward_req = ForwardRequest {
method: forward.method.as_str(),
headers: if route.has_filters {
crate::forward::AttemptHeaders::Chain(&forward.headers)
} else {
crate::forward::AttemptHeaders::Direct(&parts.headers)
},
original_headers: &parts.headers,
authority: &forward.authority,
upstream_path: &upstream_path,
traceparent: &snapshot.traceparent(),
upgrade_token: upgrade.as_ref().map(|(t, _, _)| t.as_str()),
};
let client = state.clients.for_group(&group);
let (upstream_resp, upstream_pick) = match forward_with_retry(
&client,
&state.metrics,
&group,
pick,
hash_key,
forward_req,
real_body,
timeout,
overall_deadline,
group.max_retries(),
)
.await
{
ForwardOutcome::Response(resp, pick) => (resp, pick),
ForwardOutcome::OverallTimeout => {
return Ok(synth(
StatusCode::GATEWAY_TIMEOUT,
&fault::REQUEST_TIMEOUT,
b"request timeout",
));
}
ForwardOutcome::PerTryTimeout => {
return Ok(synth(
StatusCode::GATEWAY_TIMEOUT,
&fault::UPSTREAM_TIMEOUT,
b"upstream timeout",
));
}
ForwardOutcome::SendFailed(e) => return Err(e.into()),
ForwardOutcome::BuildFailed(e) => return Err(e.into()),
};
if upstream_resp.status() == StatusCode::SWITCHING_PROTOCOLS {
return upgrade_switch(
&state,
&chain,
forward,
upgrade,
upstream_resp,
permit,
upstream_pick,
)
.await;
}
respond_through_chain(&chain, &route, &parts, forward, upstream_resp).await
}
struct ChainRef<'a> {
snapshot: &'a ConfigSnapshot,
idx: usize,
has_filters: bool,
}
type UpgradeIntent = (String, hyper::upgrade::OnUpgrade, Option<Duration>);
fn upgrade_intent(
route: &RouteInfo,
parts: &mut hyper::http::request::Parts,
bodyless: bool,
) -> Option<UpgradeIntent> {
if !bodyless {
return None;
}
let cfg = route.upgrade.as_ref()?;
let header = crate::headers::upgrade_request_header(&parts.headers)?;
let token = cfg.allowed_token(header)?.to_string();
let on_upgrade = parts.extensions.remove::<hyper::upgrade::OnUpgrade>()?;
Some((token, on_upgrade, cfg.idle_timeout()))
}
enum BodyHookOutcome {
Proceed(Option<OwnedSemaphorePermit>),
Respond(Response<ResponseBody>),
}
async fn request_body_hook(
state: &ServerState,
chain: &ChainRef<'_>,
reads_body: bool,
real_body: &mut ForwardBody,
) -> Result<BodyHookOutcome, ServerError> {
if !reads_body {
return Ok(BodyHookOutcome::Proceed(None));
}
let Some(b) = real_body.take_oneshot() else {
return Ok(BodyHookOutcome::Proceed(None));
};
let permit = match state.body_buffer_limit.clone().acquire_owned().await {
Ok(permit) => permit,
Err(_) => {
return Ok(BodyHookOutcome::Respond(synth(
StatusCode::SERVICE_UNAVAILABLE,
&fault::BODY_BUFFER_UNAVAILABLE,
b"body buffer unavailable",
)));
}
};
let buffered = match tokio::time::timeout(
INBOUND_BODY_READ_TIMEOUT,
buffer_request_body(b, MAX_REQUEST_BODY_BUFFER),
)
.await
{
Ok(crate::body::BufferOutcome::Buffered(buf)) => buf,
Ok(crate::body::BufferOutcome::TooLarge) => {
return Ok(BodyHookOutcome::Respond(synth(
StatusCode::PAYLOAD_TOO_LARGE,
&fault::BODY_TOO_LARGE,
b"request body too large",
)));
}
Ok(crate::body::BufferOutcome::ReadError) => {
return Ok(BodyHookOutcome::Respond(synth(
StatusCode::BAD_REQUEST,
&fault::BODY_READ_ERROR,
b"request body read error",
)));
}
Err(_) => {
return Ok(BodyHookOutcome::Respond(synth(
StatusCode::REQUEST_TIMEOUT,
&fault::BODY_TIMEOUT,
b"request body read timeout",
)));
}
};
let snap_body = chain.snapshot.clone();
let idx = chain.idx;
match tokio::task::spawn_blocking(move || snap_body.dispatch_request_body(idx, buffered))
.await?
{
RequestBodyOutcome::Respond(resp) => Ok(BodyHookOutcome::Respond(http_response(resp))),
RequestBodyOutcome::Forward(edited) => {
*real_body = ForwardBody::Replayable(bytes::Bytes::from(edited));
Ok(BodyHookOutcome::Proceed(Some(permit)))
}
}
}
fn upstream_switched_to(headers: &hyper::HeaderMap, token: &str) -> bool {
headers
.get(hyper::header::UPGRADE)
.and_then(|v| v.to_str().ok())
.is_some_and(|v| v.split(',').any(|t| t.trim().eq_ignore_ascii_case(token)))
}
async fn upgrade_switch(
state: &Arc<ServerState>,
chain: &ChainRef<'_>,
forward: HttpRequest,
upgrade: Option<UpgradeIntent>,
mut upstream_resp: Response<ResponseBody>,
permit: impl Send + 'static,
pick: plecto_control::Pick,
) -> Result<Response<ResponseBody>, ServerError> {
let Some((token, downstream_on, idle)) = upgrade else {
return Ok(synth(
StatusCode::BAD_GATEWAY,
&fault::BAD_UPGRADE,
b"unsolicited upgrade",
));
};
if !upstream_switched_to(upstream_resp.headers(), &token) {
return Ok(synth(
StatusCode::BAD_GATEWAY,
&fault::BAD_UPGRADE,
b"upgrade token mismatch",
));
}
let upstream_on = hyper::upgrade::on(&mut upstream_resp);
let mut builder = Response::builder().status(StatusCode::SWITCHING_PROTOCOLS);
if chain.has_filters {
let http_resp = HttpResponse {
status: StatusCode::SWITCHING_PROTOCOLS.as_u16(),
headers: headers_to_vec(upstream_resp.headers()),
body: Vec::new(),
};
let snap_resp = chain.snapshot.clone();
let idx = chain.idx;
let outcome = tokio::task::spawn_blocking(move || {
snap_resp.dispatch_response(idx, &forward, http_resp)
})
.await?;
let edited = match outcome {
ResponseOutcome::Respond(resp) => return Ok(http_response(resp)),
ResponseOutcome::Forward(edited) => edited,
};
if edited.status != StatusCode::SWITCHING_PROTOCOLS.as_u16() {
return Ok(http_response(edited));
}
copy_headers(builder.headers_mut(), &edited.headers);
} else {
copy_headers_direct(builder.headers_mut(), upstream_resp.headers());
}
let Ok(upgrade_value) = hyper::header::HeaderValue::from_str(&token) else {
return Ok(synth(
StatusCode::BAD_GATEWAY,
&fault::BAD_UPGRADE,
b"bad upgrade token",
));
};
if let Some(h) = builder.headers_mut() {
h.insert(hyper::header::UPGRADE, upgrade_value);
h.insert(
hyper::header::CONNECTION,
hyper::header::HeaderValue::from_static("upgrade"),
);
}
let resp101 = builder
.body(crate::body::full(Vec::new()))
.map_err(ServerError::from)?;
let drain = state.drain.clone();
let metrics = state.metrics.clone();
let tunnel_active = crate::metrics::TunnelActive::new(metrics.clone());
tokio::spawn(async move {
let _permit = permit;
let _pick = pick;
let _active = tunnel_active;
let (down, up) = crate::tunnel::run(downstream_on, upstream_on, idle, drain).await;
metrics.add_tunnel_bytes(down, up);
});
Ok(resp101)
}
async fn respond_through_chain(
chain: &ChainRef<'_>,
route: &RouteInfo,
parts: &hyper::http::request::Parts,
forward: HttpRequest,
upstream_resp: Response<ResponseBody>,
) -> Result<Response<ResponseBody>, ServerError> {
let (uparts, ubody) = upstream_resp.into_parts();
if !chain.has_filters {
let resp = stream_response_direct(uparts.status, &uparts.headers, ubody);
return Ok(crate::compression::apply(resp, route, parts));
}
let http_resp = HttpResponse {
status: uparts.status.as_u16(),
headers: headers_to_vec(&uparts.headers),
body: Vec::new(), };
let snap_resp = chain.snapshot.clone();
let idx = chain.idx;
let outcome =
tokio::task::spawn_blocking(move || snap_resp.dispatch_response(idx, &forward, http_resp))
.await?;
match outcome {
ResponseOutcome::Forward(edited) => {
let resp = stream_response(edited.status, &edited.headers, &uparts.headers, ubody);
Ok(crate::compression::apply(resp, route, parts))
}
ResponseOutcome::Respond(resp) => {
discard_upstream_body(ubody);
Ok(http_response(resp))
}
}
}