use std::{net::SocketAddr, time::Instant};
use mio::{Token, net::TcpStream};
use rusty_ulid::Ulid;
use super::{AlpnMatcher, Input, Output, PrereadConfig, SniPrereadCore};
use crate::{
Readiness, SessionMetrics, SessionResult,
metrics::names,
pool::Checkout,
socket::{SocketHandler, SocketResult},
sozu_command::{ready::Ready, state::ClusterId},
};
macro_rules! log_context {
($self:expr) => {{
let (open, reset, grey, gray, white) = sozu_command::logging::ansi_palette();
format!(
"{open}TCP-SNI{reset}\t{grey}Session{reset}({gray}frontend{reset}={white}{frontend}{reset})\t >>>",
open = open,
reset = reset,
grey = grey,
gray = gray,
white = white,
frontend = $self.frontend_token.0,
)
}};
}
const LOGGED_ALPN_LIMIT: usize = 8;
fn render_alpn_for_log(alpn: &[Vec<u8>]) -> String {
let shown: Vec<String> = alpn
.iter()
.take(LOGGED_ALPN_LIMIT)
.map(|p| String::from_utf8_lossy(p).into_owned())
.collect();
if alpn.len() > LOGGED_ALPN_LIMIT {
format!("{shown:?} (+{} more)", alpn.len() - LOGGED_ALPN_LIMIT)
} else {
format!("{shown:?}")
}
}
#[derive(Debug, Clone)]
pub struct RoutedOutcome {
pub cluster: ClusterId,
pub content_offset: usize,
pub proxy_source: Option<SocketAddr>,
pub sni: String,
pub alpn: Vec<Vec<u8>>,
pub matched_sni_pattern: String,
pub matched_alpn: AlpnMatcher,
}
pub struct SniPreread<Front: SocketHandler> {
pub frontend: Front,
pub frontend_token: Token,
pub frontend_readiness: Readiness,
pub backend_readiness: Readiness,
pub backend: Option<TcpStream>,
pub backend_token: Option<Token>,
pub request_id: Ulid,
pub frontend_buffer: Checkout,
effective_max_bytes: usize,
core: SniPrereadCore,
outcome: Option<RoutedOutcome>,
started_at: Instant,
}
impl<Front: SocketHandler> SniPreread<Front> {
pub fn new(
frontend: Front,
frontend_token: Token,
request_id: Ulid,
frontend_buffer: Checkout,
effective_max_bytes: usize,
) -> Self {
SniPreread {
frontend,
frontend_token,
frontend_readiness: Readiness {
interest: Ready::READABLE | Ready::HUP | Ready::ERROR,
event: Ready::EMPTY,
},
backend_readiness: Readiness {
interest: Ready::HUP | Ready::ERROR,
event: Ready::EMPTY,
},
backend: None,
backend_token: None,
request_id,
frontend_buffer,
effective_max_bytes,
core: SniPrereadCore::new(),
outcome: None,
started_at: Instant::now(),
}
}
pub fn effective_max_bytes(&self) -> usize {
self.effective_max_bytes
}
pub fn outcome(&self) -> Option<&RoutedOutcome> {
self.outcome.as_ref()
}
pub fn is_routed(&self) -> bool {
self.outcome.is_some()
}
pub fn started_at(&self) -> Instant {
self.started_at
}
pub fn has_received_bytes(&self) -> bool {
self.frontend_buffer.available_data() > 0
}
pub fn front_socket(&self) -> &TcpStream {
self.frontend.socket_ref()
}
pub fn back_socket_mut(&mut self) -> Option<&mut TcpStream> {
self.backend.as_mut()
}
pub fn set_back_socket(&mut self, socket: TcpStream) {
self.backend = Some(socket);
}
pub fn set_back_token(&mut self, token: Token) {
self.backend_token = Some(token);
}
pub fn readable(
&mut self,
metrics: &mut SessionMetrics,
cfg: &PrereadConfig<'_>,
) -> SessionResult {
if self.outcome.is_some() {
return SessionResult::Continue;
}
let cap_remaining = self
.effective_max_bytes
.saturating_sub(self.frontend_buffer.available_data());
let space_before = self.frontend_buffer.available_space();
let read_len = space_before.min(cap_remaining);
let (sz, socket_result) = self
.frontend
.socket_read(&mut self.frontend_buffer.space()[..read_len]);
debug_assert!(
sz <= read_len,
"socket_read cannot return more bytes than the capped space it was handed"
);
if sz > 0 {
let data_before = self.frontend_buffer.available_data();
self.frontend_buffer.fill(sz);
debug_assert_eq!(
self.frontend_buffer.available_data(),
data_before + sz,
"fill must expose exactly the bytes just read"
);
count!(names::backend::BYTES_IN, sz as i64);
metrics.bin += sz;
if socket_result == SocketResult::Error {
return self.front_gone(cfg);
}
if socket_result == SocketResult::WouldBlock {
self.frontend_readiness.event.remove(Ready::READABLE);
}
let output = self.core.handle_input(
cfg,
Input::Bytes {
buf: self.frontend_buffer.data(),
now: Instant::now(),
},
);
return self.handle_output(output);
}
match socket_result {
SocketResult::Error | SocketResult::Closed => self.front_gone(cfg),
SocketResult::WouldBlock => {
self.frontend_readiness.event.remove(Ready::READABLE);
SessionResult::Continue
}
SocketResult::Continue => SessionResult::Continue,
}
}
pub fn on_timeout(&mut self, cfg: &PrereadConfig<'_>) {
let output = self.core.handle_input(
cfg,
Input::Timeout {
now: Instant::now(),
},
);
let _ = self.handle_output(output);
}
pub fn on_front_closed(&mut self, cfg: &PrereadConfig<'_>) {
let _ = self.front_gone(cfg);
}
fn front_gone(&mut self, cfg: &PrereadConfig<'_>) -> SessionResult {
if self.has_received_bytes() {
let output = self.core.handle_input(cfg, Input::FrontClosed);
self.handle_output(output)
} else {
trace!(
"{} front socket closed with 0 bytes during SNI preread",
log_context!(self)
);
self.frontend_readiness.reset();
SessionResult::Close
}
}
fn handle_output(&mut self, output: Output) -> SessionResult {
match output {
Output::NeedMore { .. } => SessionResult::Continue,
Output::Routed {
cluster,
content_offset,
proxy_source,
sni,
alpn,
matched_sni_pattern,
matched_alpn,
} => {
debug_assert!(
self.outcome.is_none(),
"a route must be captured at most once per SniPreread lifetime"
);
info!(
"{} SNI preread routed to cluster {} (sni={:?}, alpn={})",
log_context!(self),
cluster,
sni,
render_alpn_for_log(&alpn)
);
incr!(names::tcp::sni_preread::ROUTED);
self.backend_readiness.interest.insert(Ready::WRITABLE);
self.outcome = Some(RoutedOutcome {
cluster,
content_offset,
proxy_source,
sni,
alpn,
matched_sni_pattern,
matched_alpn,
});
SessionResult::Continue
}
Output::Reject(reason) => {
debug!("{} SNI preread rejected: {:?}", log_context!(self), reason);
incr!(names::tcp::sni_preread::rejected_name(reason));
self.frontend_readiness.reset();
self.backend_readiness.reset();
SessionResult::Close
}
}
}
pub fn back_writable(&self) -> SessionResult {
debug_assert!(
self.outcome.is_some(),
"back_writable on SniPreread must only fire after a route decision"
);
SessionResult::Upgrade
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
use crate::protocol::tcp_preread::{AlpnMatcher, RejectReason};
use crate::router::pattern_trie::TrieNode;
#[test]
fn every_reject_reason_has_a_distinct_metric_name() {
let reasons = [
RejectReason::NotTls,
RejectReason::MalformedRecord,
RejectReason::MalformedHandshake,
RejectReason::Fragmented,
RejectReason::TooLarge,
RejectReason::NoSni,
RejectReason::EchOuterAbsent,
RejectReason::SniUnmatched,
RejectReason::AlpnUnmatched,
RejectReason::ProxyHeaderInvalid,
RejectReason::FrontClosed,
];
let mut names_seen = std::collections::HashSet::new();
for reason in reasons {
let name = names::tcp::sni_preread::rejected_name(reason);
assert!(
name.starts_with("tcp.sni_preread.rejected."),
"unexpected metric name shape for {reason:?}: {name}"
);
assert!(
names_seen.insert(name),
"two RejectReason variants mapped to the same metric name: {name}"
);
}
}
fn cfg(routes: &TrieNode<Vec<(AlpnMatcher, ClusterId)>>) -> PrereadConfig<'_> {
PrereadConfig {
routes,
inbound_proxy: false,
max_bytes: 16 * 1024,
timeout: Duration::from_secs(3),
accept_wildcard: true,
}
}
#[test]
fn render_alpn_for_log_shows_everything_within_the_limit() {
let alpn: Vec<Vec<u8>> = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
assert_eq!(render_alpn_for_log(&alpn), r#"["h2", "http/1.1"]"#);
}
#[test]
fn render_alpn_for_log_truncates_past_the_limit_with_a_count_suffix() {
let alpn: Vec<Vec<u8>> = (0..(LOGGED_ALPN_LIMIT + 5))
.map(|i| i.to_string().into_bytes())
.collect();
let rendered = render_alpn_for_log(&alpn);
assert!(
rendered.ends_with("(+5 more)"),
"expected a truncation suffix for 5 entries past the limit, got {rendered:?}"
);
let shown_count = alpn
.iter()
.take(LOGGED_ALPN_LIMIT)
.filter(|p| {
let s = String::from_utf8_lossy(p).into_owned();
rendered.contains(&format!("{s:?}"))
})
.count();
assert_eq!(shown_count, LOGGED_ALPN_LIMIT);
}
#[test]
fn preread_config_helper_shape_is_sane() {
let routes = TrieNode::root();
let cfg = cfg(&routes);
assert!(!cfg.inbound_proxy);
assert_eq!(cfg.max_bytes, 16 * 1024);
assert!(cfg.accept_wildcard);
}
#[test]
fn readable_caps_the_accumulator_at_effective_max_bytes() {
use std::io::Write as _;
use std::net::{TcpListener as StdTcpListener, TcpStream as StdTcpStream};
use mio::net::TcpStream as MioTcpStream;
use crate::pool::Pool;
let listener = StdTcpListener::bind("127.0.0.1:0").expect("bind test listener");
let addr = listener.local_addr().expect("listener local addr");
let mut client = StdTcpStream::connect(addr).expect("connect test client");
let (server, _) = listener.accept().expect("accept test server");
server.set_nonblocking(true).expect("server nonblocking");
let flood = vec![0x16u8; 4096];
client.write_all(&flood).expect("write flood");
client.flush().ok();
let mut pool = Pool::with_capacity(1, 1, 16 * 1024);
let frontend_buffer = pool.checkout().expect("frontend buffer");
assert!(
frontend_buffer.available_space() > 4096,
"buffer must dwarf the cap for this test to be meaningful"
);
let effective_max_bytes = 16usize;
let mut preread = SniPreread::new(
MioTcpStream::from_std(server),
Token(0),
Ulid::generate(),
frontend_buffer,
effective_max_bytes,
);
let routes = TrieNode::root();
let cfg = PrereadConfig {
routes: &routes,
inbound_proxy: false,
max_bytes: effective_max_bytes,
timeout: Duration::from_secs(3),
accept_wildcard: true,
};
let mut metrics = SessionMetrics::new(Some(Duration::ZERO));
let result = preread.readable(&mut metrics, &cfg);
assert!(
preread.frontend_buffer.available_data() <= effective_max_bytes,
"readable accumulated {} bytes, past the {}-byte cap",
preread.frontend_buffer.available_data(),
effective_max_bytes
);
assert_eq!(
result,
SessionResult::Close,
"an over-cap preread window must reject-and-close"
);
drop(client);
}
}