use std::collections::BTreeMap;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use axum::body::Body;
use axum::extract::Request;
use axum::http::StatusCode;
use axum::response::Response;
use hyper::upgrade::Upgraded;
use hyper_util::rt::TokioIo;
use super::ca;
use super::{proxy, ModelRelayState};
use crate::error::OlError;
#[derive(Clone, Debug)]
pub(crate) struct ConnectAuthority {
pub(crate) host: String,
pub(crate) port: u16,
}
const INTERCEPT_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
enum Plan {
Intercept(Arc<ca::Interceptor>),
Tunnel(crate::egress::tunnel::TunnelStream),
}
pub async fn handle_connect(st: Arc<ModelRelayState>, mut req: Request<Body>) -> Response {
let Some(authority) = req.uri().authority().cloned() else {
let mut resp = Response::new(Body::empty());
*resp.status_mut() = StatusCode::BAD_REQUEST;
return resp;
};
let host = authority
.host()
.trim_start_matches('[')
.trim_end_matches(']')
.to_ascii_lowercase();
let port = authority.port_u16().unwrap_or(443);
let preflight = proxy::is_preflight(req.headers());
let slot = st
.intercept
.read()
.unwrap_or_else(|e| e.into_inner())
.clone();
let plan = match slot.filter(|ic| ic.intercepts(&host)) {
Some(ic) => Plan::Intercept(ic),
None => match open_upstream(&st, &host, port).await {
Ok(up) => Plan::Tunnel(up),
Err(e) => {
tracing::warn!(
code = %e.code,
host = %host,
port,
"CONNECT tunnel could not be opened"
);
return proxy::synth_502(&st);
}
},
};
let upgrade = hyper::upgrade::on(&mut req);
tokio::spawn(async move {
let upgraded = match upgrade.await {
Ok(u) => u,
Err(e) => {
tracing::debug!(error = %e, "model relay CONNECT upgrade failed");
return;
}
};
let io = TokioIo::new(upgraded);
match plan {
Plan::Intercept(ic) => terminate(st, ic, io, host, port, preflight).await,
Plan::Tunnel(mut upstream) => {
let mut io = io;
let _ = tokio::io::copy_bidirectional(&mut io, &mut upstream).await;
}
}
});
Response::new(Body::empty())
}
async fn open_upstream(
st: &Arc<ModelRelayState>,
host: &str,
port: u16,
) -> Result<crate::egress::tunnel::TunnelStream, OlError> {
if !st.resolve_override.is_empty() {
if let Some(addr) = st.resolve_override.get(&(host.to_string(), port)) {
return crate::egress::tunnel::connect_with_timeout(*addr, host, port).await;
}
}
crate::egress::tunnel::open(st.egress.config(), host, port).await
}
async fn terminate(
st: Arc<ModelRelayState>,
ic: Arc<ca::Interceptor>,
io: TokioIo<Upgraded>,
host: String,
port: u16,
preflight: bool,
) {
let start = match tokio::time::timeout(
INTERCEPT_HANDSHAKE_TIMEOUT,
tokio_rustls::LazyConfigAcceptor::new(rustls::server::Acceptor::default(), io),
)
.await
{
Ok(Ok(start)) => start,
_ => {
tracing::debug!(host = %host, "model relay intercept closed before a ClientHello");
return;
}
};
let sni_ok = start
.client_hello()
.server_name()
.is_none_or(|s| s.eq_ignore_ascii_case(&host));
let config = if sni_ok {
ic.pinned_config(&host)
} else {
None
};
let Some(config) = config else {
tracing::debug!(
host = %host,
"model relay intercept SNI does not match the CONNECT authority"
);
return;
};
let tls =
match tokio::time::timeout(INTERCEPT_HANDSHAKE_TIMEOUT, start.into_stream(config)).await {
Ok(Ok(tls)) => {
if !preflight {
st.refusals.record_success(&host);
}
tls
}
Ok(Err(e)) => {
if !preflight {
st.refusals.record_refusal(&host, classify_refusal(&e));
}
proxy::record_pass_through_failure("intercept_handshake");
return;
}
Err(_) => {
if !preflight {
st.refusals
.record_refusal(&host, RefusalKind::Other("handshake timed out".into()));
}
proxy::record_pass_through_failure("intercept_handshake");
return;
}
};
let svc = super::router(st).layer(axum::Extension(ConnectAuthority { host, port }));
if let Err(e) = hyper::server::conn::http1::Builder::new()
.serve_connection(
TokioIo::new(tls),
hyper_util::service::TowerToHyperService::new(svc),
)
.with_upgrades()
.await
{
tracing::debug!(error = %e, "model relay intercept connection ended");
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RefusalKind {
Alert(String),
HandshakeEof,
Reset,
Other(String),
}
fn classify_refusal(e: &std::io::Error) -> RefusalKind {
if let Some(rustls::Error::AlertReceived(d)) =
e.get_ref().and_then(|i| i.downcast_ref::<rustls::Error>())
{
return RefusalKind::Alert(format!("{d:?}"));
}
match e.kind() {
std::io::ErrorKind::UnexpectedEof => RefusalKind::HandshakeEof,
std::io::ErrorKind::ConnectionReset
| std::io::ErrorKind::ConnectionAborted
| std::io::ErrorKind::BrokenPipe => RefusalKind::Reset,
_ => RefusalKind::Other(e.to_string()),
}
}
const TRIP_AFTER: u32 = 3;
const SUCCESS_WINDOW: Duration = Duration::from_secs(10 * 60);
#[derive(Default)]
struct HostRecord {
consecutive: u32,
last_success: Option<Instant>,
tripped_logged: bool,
}
#[derive(Default)]
pub struct RefusalLedger {
hosts: Mutex<BTreeMap<String, HostRecord>>,
}
impl RefusalLedger {
pub fn record_refusal(&self, host: &str, kind: RefusalKind) {
let host = host.to_ascii_lowercase();
let mut hosts = self.hosts.lock().unwrap_or_else(|e| e.into_inner());
let record = hosts.entry(host.clone()).or_default();
record.consecutive += 1;
tracing::debug!(host = %host, kind = ?kind, "model relay intercept handshake refused");
if record.consecutive >= TRIP_AFTER && !record.tripped_logged {
record.tripped_logged = true;
tracing::warn!(host = %host, kind = ?kind, "model relay intercept host tripped");
}
}
pub fn record_success(&self, host: &str) {
self.record_success_at(host, Instant::now());
}
pub fn tripped_hosts(&self) -> Vec<String> {
self.tripped_hosts_at(Instant::now())
}
pub fn clear(&self, host: &str) {
let host = host.to_ascii_lowercase();
self.hosts
.lock()
.unwrap_or_else(|e| e.into_inner())
.remove(&host);
}
fn record_success_at(&self, host: &str, at: Instant) {
let host = host.to_ascii_lowercase();
let mut hosts = self.hosts.lock().unwrap_or_else(|e| e.into_inner());
let record = hosts.entry(host).or_default();
record.consecutive = 0;
record.last_success = Some(at);
record.tripped_logged = false;
}
fn tripped_hosts_at(&self, now: Instant) -> Vec<String> {
let hosts = self.hosts.lock().unwrap_or_else(|e| e.into_inner());
hosts
.iter()
.filter(|(_, r)| {
r.consecutive >= TRIP_AFTER
&& r.last_success
.is_none_or(|t| now.saturating_duration_since(t) > SUCCESS_WINDOW)
})
.map(|(h, _)| h.clone())
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use rustls::pki_types::pem::PemObject;
#[test]
fn a_tls_alert_is_classified_by_its_description() {
let err = std::io::Error::new(
std::io::ErrorKind::InvalidData,
rustls::Error::AlertReceived(rustls::AlertDescription::UnknownCA),
);
assert_eq!(
classify_refusal(&err),
RefusalKind::Alert("UnknownCA".to_string())
);
}
#[test]
fn eof_reset_and_other_are_classified() {
let eof = std::io::Error::from(std::io::ErrorKind::UnexpectedEof);
assert_eq!(classify_refusal(&eof), RefusalKind::HandshakeEof);
for kind in [
std::io::ErrorKind::ConnectionReset,
std::io::ErrorKind::ConnectionAborted,
std::io::ErrorKind::BrokenPipe,
] {
assert_eq!(
classify_refusal(&std::io::Error::from(kind)),
RefusalKind::Reset
);
}
let other = std::io::Error::new(
std::io::ErrorKind::InvalidData,
rustls::Error::NoApplicationProtocol,
);
assert!(matches!(classify_refusal(&other), RefusalKind::Other(_)));
}
#[test]
fn three_consecutive_refusals_trip_a_host() {
let ledger = RefusalLedger::default();
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_refusal("a.test", RefusalKind::Reset);
assert!(ledger.tripped_hosts().is_empty());
ledger.record_refusal("a.test", RefusalKind::Reset);
assert_eq!(ledger.tripped_hosts(), vec!["a.test".to_string()]);
ledger.record_refusal("b.test", RefusalKind::Reset);
assert_eq!(ledger.tripped_hosts(), vec!["a.test".to_string()]);
}
#[test]
fn a_success_within_ten_minutes_holds_the_trip() {
let ledger = RefusalLedger::default();
let t0 = Instant::now();
ledger.record_success_at("a.test", t0);
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_refusal("a.test", RefusalKind::Reset);
assert!(ledger
.tripped_hosts_at(t0 + Duration::from_secs(9 * 60 + 59))
.is_empty());
assert_eq!(
ledger.tripped_hosts_at(t0 + Duration::from_secs(10 * 60 + 1)),
vec!["a.test".to_string()]
);
}
#[test]
fn a_success_resets_the_count() {
let ledger = RefusalLedger::default();
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_success("a.test");
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_refusal("a.test", RefusalKind::Reset);
assert!(ledger
.tripped_hosts_at(Instant::now() + Duration::from_secs(11 * 60))
.is_empty());
}
#[test]
fn clear_forgets_a_host() {
let ledger = RefusalLedger::default();
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_refusal("a.test", RefusalKind::Reset);
assert_eq!(ledger.tripped_hosts(), vec!["a.test".to_string()]);
ledger.clear("a.test");
assert!(ledger.tripped_hosts().is_empty());
ledger.record_refusal("a.test", RefusalKind::Reset);
assert!(ledger.tripped_hosts().is_empty());
}
#[tokio::test]
async fn a_preflight_connect_is_not_counted() {
let tmp = tempfile::tempdir().expect("tempdir");
let interceptor =
Arc::new(ca::Interceptor::new(&tmp.path().join("ca"), ["a.test"]).expect("new"));
let ledger = Arc::new(RefusalLedger::default());
let closed = crate::model_relay::mock::closed_port().await;
let mut resolve = BTreeMap::new();
resolve.insert(
("a.test".to_string(), 443u16),
std::net::SocketAddr::from(([127, 0, 0, 1], closed)),
);
let slot: super::super::InterceptSlot =
Arc::new(std::sync::RwLock::new(Some(interceptor.clone())));
let state = Arc::new(
ModelRelayState::new(
reqwest::Url::parse("http://127.0.0.1:1").unwrap(),
0,
8,
&[],
)
.with_intercept(slot)
.with_refusals(ledger.clone())
.with_resolve_override(resolve),
);
let port = super::super::serve_ephemeral(state).await;
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_refusal("a.test", RefusalKind::Reset);
ledger.record_refusal("a.test", RefusalKind::Reset);
assert_eq!(ledger.tripped_hosts(), vec!["a.test".to_string()]);
probe(port, &interceptor.ca().pem_path(), true).await;
assert_eq!(
ledger.tripped_hosts(),
vec!["a.test".to_string()],
"a preflight connect's handshake must not be recorded, success or refusal"
);
probe(port, &interceptor.ca().pem_path(), false).await;
assert!(ledger.tripped_hosts().is_empty());
}
async fn probe(port: u16, ca_pem_path: &std::path::Path, preflight: bool) {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut stream = tokio::net::TcpStream::connect(("127.0.0.1", port))
.await
.expect("dial relay");
let mut head = "CONNECT a.test:443 HTTP/1.1\r\nHost: a.test:443\r\n".to_string();
if preflight {
head.push_str(&format!(
"{}: 1\r\n",
super::super::preflight::PREFLIGHT_HEADER
));
}
head.push_str("\r\n");
stream
.write_all(head.as_bytes())
.await
.expect("write CONNECT");
let mut buf = Vec::new();
let mut byte = [0u8; 1];
while !buf.ends_with(b"\r\n\r\n") {
stream
.read_exact(&mut byte)
.await
.expect("read CONNECT response");
buf.push(byte[0]);
}
assert!(
String::from_utf8_lossy(&buf).starts_with("HTTP/1.1 200"),
"unexpected CONNECT response: {}",
String::from_utf8_lossy(&buf)
);
let ca_pem = std::fs::read_to_string(ca_pem_path).expect("read ca pem");
let ca_der =
rustls::pki_types::CertificateDer::from_pem_slice(ca_pem.as_bytes()).expect("parse");
let mut roots = rustls::RootCertStore::empty();
roots.add(ca_der).expect("trust the CA");
let client_config = rustls::ClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth();
let connector = tokio_rustls::TlsConnector::from(Arc::new(client_config));
let server_name = rustls::pki_types::ServerName::try_from("a.test").expect("server name");
let mut tls = connector
.connect(server_name, stream)
.await
.expect("tls handshake against our own leaf must succeed");
tls.write_all(b"GET / HTTP/1.1\r\nHost: a.test\r\nConnection: close\r\n\r\n")
.await
.expect("write GET");
let mut response = Vec::new();
let _ = tls.read_to_end(&mut response).await;
}
}