Skip to content
12 changes: 12 additions & 0 deletions dstack/gateway/src/models.rs
Original file line number Diff line number Diff line change
Expand Up @@ -451,6 +451,18 @@ impl<C: Counting> EnteredCounter<C> {
Self(connections)
}
}
impl EnteredCounter {
/// Enter only if `others + counter < max`, checking and incrementing in one
/// atomic step.
pub fn try_enter(counter: &Arc<AtomicU64>, others: u64, max: u64) -> Option<Self> {
counter
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |n| {
(n + others < max).then_some(n + 1)
})
.ok()?;
Some(Self(counter.clone()))
}
}
impl<C: Counting> Drop for EnteredCounter<C> {
fn drop(&mut self) {
self.0.dec();
Expand Down
62 changes: 62 additions & 0 deletions dstack/gateway/src/proxy.rs
Original file line number Diff line number Diff line change
Expand Up @@ -88,10 +88,20 @@ pub(crate) mod stats;
mod tls_passthough;
mod tls_terminate;

/// One TLS record. `extract_sni` does not reassemble records, so a larger
/// ClientHello cannot be parsed anyway.
const MAX_CLIENT_HELLO_BYTES: usize = 5 + (1 << 14);

async fn take_sni(stream: &mut TcpStream) -> Result<(Option<String>, Vec<u8>)> {
let mut buffer = vec![0u8; 4096];
let mut data_len = 0;
loop {
if data_len == buffer.len() {
if buffer.len() >= MAX_CLIENT_HELLO_BYTES {
bail!("no sni in the first {MAX_CLIENT_HELLO_BYTES} bytes of the client hello");
}
buffer.resize((buffer.len() * 2).min(MAX_CLIENT_HELLO_BYTES), 0);
}
// read data from stream
let n = stream
.read(&mut buffer[data_len..])
Expand Down Expand Up @@ -646,6 +656,58 @@ mod tests {
assert_eq!(config.ktls.is_some(), probe_ktls().is_ok());
}

/// A ClientHello whose `server_name` sits behind `filler` bytes of key share.
fn client_hello(sni: &str, filler: usize) -> Vec<u8> {
fn u16_prefixed(body: &[u8]) -> Vec<u8> {
let mut out = (body.len() as u16).to_be_bytes().to_vec();
out.extend_from_slice(body);
out
}

let mut server_name = vec![0u8]; // name_type: host_name
server_name.extend_from_slice(&u16_prefixed(sni.as_bytes()));
let mut extensions = vec![0x00, 0x33]; // key_share
extensions.extend_from_slice(&u16_prefixed(&vec![0u8; filler]));
extensions.extend_from_slice(&[0x00, 0x00]); // server_name
extensions.extend_from_slice(&u16_prefixed(&u16_prefixed(&server_name)));

let mut body = vec![0x03, 0x03]; // legacy_version
body.extend_from_slice(&[0u8; 32]); // random
body.push(0); // session_id
body.extend_from_slice(&u16_prefixed(&[0x13, 0x01])); // cipher_suites
body.extend_from_slice(&[1, 0]); // compression_methods
body.extend_from_slice(&u16_prefixed(&extensions));

let mut handshake = vec![0x01]; // client_hello
handshake.extend_from_slice(&(body.len() as u32).to_be_bytes()[1..]);
handshake.extend_from_slice(&body);

let mut record = vec![0x16, 0x03, 0x01]; // handshake record
record.extend_from_slice(&u16_prefixed(&handshake));
record
}

/// Feed `hello` to `take_sni`, keeping the client open so a short read is
/// not mistaken for a hangup.
async fn sniff(hello: Vec<u8>) -> (Result<(Option<String>, Vec<u8>)>, TcpStream) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let mut client = TcpStream::connect(listener.local_addr().unwrap())
.await
.unwrap();
let (mut server, _) = listener.accept().await.unwrap();
tokio::io::AsyncWriteExt::write_all(&mut client, &hello)
.await
.unwrap();
(take_sni(&mut server).await, client)
}

#[tokio::test]
async fn an_sni_past_the_first_read_is_still_found() {
let (sniffed, _client) = sniff(client_hello("app.example.com", 6000)).await;
let (sni, _buffer) = sniffed.unwrap();
assert_eq!(sni.as_deref(), Some("app.example.com"));
}

#[test]
fn test_parse_destination() {
// Test basic app_id only
Expand Down
95 changes: 61 additions & 34 deletions dstack/gateway/src/proxy/adaptive_ktls.rs
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,34 @@ where
};

let phase = loop {
// A peer that stops reading parks the write, so the watchdog must be
// polled during it. `write_all` is not cancel-safe; single `write`s are.
macro_rules! write_watched {
($w:expr, $buf:expr, $n:expr, $ctx:expr) => {{
let n = $n;
let mut written = 0usize;
while written < n {
let count = tokio::select! {
() = async { match watchdog.as_mut() {
Some(w) => w.tick().await,
None => std::future::pending().await,
} } => {
if watchdog.as_mut().is_some_and(|w| w.stalled(moved)) {
bail!("idle timeout");
}
continue;
}
r = $w.write(&$buf[written..n]) => r.context($ctx)?,
};
if count == 0 {
bail!("write accepted no bytes");
}
written += count;
moved += count as u64;
}
}};
}

// One side closing is not the end of the connection: a client that ends
// its request with close_notify still expects the response. Propagate
// the EOF to that direction's peer, then drain the other direction
Expand All @@ -93,36 +121,7 @@ where
if n == 0 {
break;
}
// Not `write_all`: it is not cancel-safe, so it cannot sit
// in a `select!`, and leaving it outside meant a client
// that stopped reading blocked the drain in its write with
// the watchdog unpolled -- the mirror of the silent-backend
// stall, and just as good for holding a connection to
// `timeouts.total`. Single `write` calls are cancel-safe
// (nothing is written when the other branch wins), so the
// partial-write loop is ours to drive.
let mut written = 0usize;
while written < n {
let count = tokio::select! {
() = async { match watchdog.as_mut() {
Some(w) => w.tick().await,
None => std::future::pending().await,
} } => {
if watchdog.as_mut().is_some_and(|w| w.stalled(moved)) {
bail!("idle timeout");
}
continue;
}
r = $w.write(&$buf[written..n]) => r.context("write error")?,
};
if count == 0 {
bail!("write accepted no bytes");
}
written += count;
// Per partial write, so a peer draining slowly still
// counts as progress and is not reaped for being slow.
moved += count as u64;
}
write_watched!($w, $buf, n, "write error");
}
$w.shutdown().await.ok();
break Phase::Eof;
Expand All @@ -142,14 +141,12 @@ where
r = tr.read(&mut down) => {
let n = r.context("read from client failed")?;
if n == 0 { finish_one!(uw, ur, tw, up); }
uw.write_all(&down[..n]).await.context("write to app failed")?;
moved += n as u64;
write_watched!(uw, down, n, "write to app failed");
}
r = ur.read(&mut up) => {
let n = r.context("read from app failed")?;
if n == 0 { finish_one!(tw, tr, uw, down); }
tw.write_all(&up[..n]).await.context("write to client failed")?;
moved += n as u64;
write_watched!(tw, up, n, "write to client failed");
}
}
if gate.reached(moved, start) {
Expand Down Expand Up @@ -242,6 +239,36 @@ mod tests {
}
}

/// Shrink the socket buffers so a few KiB of unread data parks a write.
fn throttle(stream: &TcpStream) {
let sock = socket2::SockRef::from(stream);
sock.set_recv_buffer_size(4096).unwrap();
sock.set_send_buffer_size(4096).unwrap();
}

// Real timers: a paused clock fires the watchdog before the write blocks.
#[tokio::test]
async fn a_client_that_stops_reading_is_reaped_by_the_idle_timeout() {
let (client, mut tls_side) = connected_pair().await;
let (mut upstream, mut backend) = connected_pair().await;
for s in [&client, &tls_side, &upstream, &backend] {
throttle(s);
}
let _app = tokio::spawn(async move {
backend.write_all(&vec![0u8; 1 << 20]).await.ok();
backend
});

let idle = Duration::from_secs(1);
let relayed = tokio::time::timeout(
idle * 10,
relay_until(&mut tls_side, &mut upstream, &ungated(), Some(idle)),
)
.await
.expect("the relay outlived its idle window");
assert!(matches!(relayed, Err(e) if e.to_string().contains("idle timeout")));
}

#[tokio::test]
async fn response_survives_a_client_half_close_before_the_gate() {
let (mut client, mut tls_side) = connected_pair().await;
Expand Down
78 changes: 77 additions & 1 deletion dstack/gateway/src/proxy/io_bridge.rs
Original file line number Diff line number Diff line change
Expand Up @@ -56,10 +56,16 @@ where
Ok(false)
}
NextStep::Write => {
self.writer
let n = self
.writer
.write_buf(&mut self.buf)
.await
.context("write error")?;
// Otherwise the loop re-enters at once and the empty write
// counts as progress, spinning until `timeouts.total`.
if n == 0 && !self.buf.is_empty() {
bail!("write accepted no bytes");
}
self.progress += 1;
if self.buf.is_empty() {
self.next_step = NextStep::Flush;
Expand Down Expand Up @@ -245,3 +251,73 @@ where
}
}
}

#[cfg(test)]
mod tests {
use super::*;
use std::pin::Pin;
use std::task::{Context as TaskContext, Poll};
use std::time::Duration;

/// Reports success while accepting nothing.
struct AcceptsNothing;

impl AsyncWrite for AcceptsNothing {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut TaskContext<'_>,
_buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Poll::Ready(Ok(0))
}

fn poll_flush(
self: Pin<&mut Self>,
_cx: &mut TaskContext<'_>,
) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}

fn poll_shutdown(
self: Pin<&mut Self>,
_cx: &mut TaskContext<'_>,
) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
}

// Own thread and runtime: the pre-fix spin never yields, so an in-runtime
// timeout would never fire.
#[test]
fn a_writer_that_accepts_no_bytes_is_an_error_not_a_spin() {
let (tx, rx) = std::sync::mpsc::channel();
std::thread::spawn(move || {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
let outcome = rt.block_on(async {
let config: ProxyConfig = crate::config::load_config_figment(None)
.focus("core.proxy")
.extract()
.unwrap();
let mut client_rx: &[u8] = b"request";
relay(
&mut client_rx,
&mut tokio::io::sink(),
&mut tokio::io::empty(),
&mut AcceptsNothing,
&config,
)
.await
});
tx.send(outcome.map_err(|e| format!("{e:#}"))).ok();
});

let err = rx
.recv_timeout(Duration::from_secs(5))
.expect("the bridge spun on a zero-length write")
.unwrap_err();
assert!(err.contains("write accepted no bytes"), "{err}");
}
}
15 changes: 15 additions & 0 deletions dstack/gateway/src/proxy/sni.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,9 @@
//
// SPDX-License-Identifier: Apache-2.0

//! ClientHello SNI sniffing. Records are not reassembled, so a ClientHello
//! fragmented across records yields no SNI.

use parcelona::parser_combinators::{Msg, PErr};
use parcelona::u8::*;
use tracing::trace;
Expand All @@ -28,6 +31,7 @@ fn take_record_u8_len(b: &[u8]) -> Result<(&[u8], &[u8]), PErr<'_, u8>> {
}

fn extract_sni_inner(b: &[u8]) -> Result<(usize, &[u8]), PErr<'_, u8>> {
const CONTENT_TYPE_HANDSHAKE: u8 = 0x16;
const HANDSHAKE_TYPE_CLIENT_HELLO: usize = 1;
const EXTENSION_TYPE_SNI: usize = 0;
const NAME_TYPE_HOST_NAME: usize = 0;
Expand All @@ -37,6 +41,10 @@ fn extract_sni_inner(b: &[u8]) -> Result<(usize, &[u8]), PErr<'_, u8>> {
return Err(PErr::new(b));
}

if b[0] != CONTENT_TYPE_HANDSHAKE {
let err = PErr::new(b).user_msg_push(Msg::Str("not a handshake record"));
return Err(err);
}
let b = &b[5..];
// Handshake message type.
let (b, c) = take_len_be_u8(b)?;
Expand Down Expand Up @@ -141,6 +149,13 @@ mod tests {
assert_eq!(extract_sni(&hello), Some(&b"app.example.com"[..]));
}

#[test]
fn a_record_that_is_not_a_handshake_has_no_sni() {
let mut hello = client_hello(b"app.example.com");
hello[0] = 0x17; // application_data
assert_eq!(extract_sni(&hello), None);
}

#[test]
fn a_truncated_client_hello_yields_no_name_rather_than_a_wrong_one() {
let hello = client_hello(b"app.example.com");
Expand Down
Loading
Loading