use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
use tracing::Instrument as _;
use tracing::field::Empty;
use crate::backend::Backend;
use crate::body::{self, Buffered};
use crate::credential::Credential;
use crate::framing;
use crate::head_read;
use crate::headers;
use crate::http_head;
use crate::outcome::Outcome;
use crate::request_body::{self, Carried};
const HEAD_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
use crate::refusal;
pub(crate) async fn serve_exchange<S, B>(
stream: &mut S,
credential: &Credential,
backend: &B,
peer: &str,
) -> std::io::Result<Outcome>
where
S: AsyncRead + AsyncWrite + Unpin + Send,
B: Backend + Sync,
{
let span = tracing::info_span!(
"exchange",
method = Empty,
path = Empty,
status = Empty,
device = Empty
);
async {
let started = std::time::Instant::now();
let result = run(stream, credential, backend, peer).await;
let elapsed_ms = u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX);
match &result {
Ok(outcome) => tracing::info!(outcome = outcome.as_str(), elapsed_ms, "exchange"),
Err(error) => tracing::warn!(%error, elapsed_ms, "exchange failed"),
}
result
}
.instrument(span)
.await
}
async fn run<S, B>(
stream: &mut S,
credential: &Credential,
backend: &B,
peer: &str,
) -> std::io::Result<Outcome>
where
S: AsyncRead + AsyncWrite + Unpin + Send,
B: Backend + Sync,
{
let asked = tokio::time::timeout(HEAD_TIMEOUT, head_read::request(stream, Vec::new())).await;
let Ok(asked) = asked else {
return Ok(Outcome::TimedOut);
};
let Ok((mut head, leftover)) = asked? else {
return refusal::send(stream, refusal::bad_request(), Outcome::BadRequest).await;
};
let span = tracing::Span::current();
span.record("method", head.method.as_str());
span.record("path", http_head::path_only(&head.target));
let Ok(request_framing) = framing::framing(&head.headers, false) else {
return refusal::send(stream, refusal::bad_request(), Outcome::BadRequest).await;
};
let Some(admitted) = credential.admits(http_head::authorization(&head.headers)) else {
return refusal::send(stream, refusal::unauthorized(), Outcome::Unauthorized).await;
};
if let Some(device) = admitted.device() {
span.record("device", device);
}
let method = head.method.clone();
let expects_continue = http_head::expects_continue(&head.headers);
let forward = credential.forward(&admitted);
http_head::rewrite_for_backend(&mut head, backend.authority(), peer, &forward);
let Ok(upstream) = backend.connect().await else {
return refusal::send(stream, refusal::backend_unreachable(), Outcome::BadGateway).await;
};
let (mut up_read, mut up_write) = tokio::io::split(upstream);
up_write
.write_all(&http_head::serialize_request(&head, request_framing))
.await?;
up_write.flush().await?;
if expects_continue {
stream.write_all(b"HTTP/1.1 100 Continue\r\n\r\n").await?;
stream.flush().await?;
}
let carried = request_body::carry(
stream,
leftover,
request_framing,
&mut up_write,
&mut up_read,
)
.await;
let (mut response, response_leftover) = match carried {
Carried::Answered(response, rest) => (response, rest),
Carried::Unreadable => {
return refusal::send(stream, refusal::bad_gateway(), Outcome::BadGateway).await;
}
Carried::Unfinished => {
return refusal::send(stream, refusal::incomplete_request(), Outcome::Unfinished).await;
}
};
span.record("status", response.status);
let Ok(response_framing) =
framing::response_framing(response.status, &method, &response.headers)
else {
return refusal::send(stream, refusal::bad_gateway(), Outcome::BadGateway).await;
};
headers::strip_hop_by_hop(&mut response.headers);
response
.headers
.push(("Connection".to_owned(), "close".to_owned()));
stream
.write_all(&http_head::serialize_response(&response, response_framing))
.await?;
stream.flush().await?;
let mut from_backend = Buffered::new(&mut up_read, response_leftover);
body::forward(&mut from_backend, stream, response_framing).await?;
Ok(Outcome::Forwarded)
}
#[cfg(test)]
#[path = "exchange_tests.rs"]
mod exchange_tests;