#![allow(dead_code)]
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use dashmap::DashMap;
use quinn::{ClientConfig, Connection, Endpoint, VarInt};
use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
use tokio::runtime::{Builder, Runtime};
use tokio::sync::Mutex;
pub const OP_EXECUTE: u8 = 1;
pub const OP_INFO: u8 = 2;
pub const OP_HEALTH: u8 = 3;
pub const STAT_OK: u8 = 0;
pub const STAT_USER_ERROR: u8 = 1;
pub const STAT_PROTOCOL_ERROR: u8 = 2;
pub const ALPN: &[u8] = b"zk-worker";
pub const DEFAULT_PORT: u16 = 4433;
#[derive(Debug)]
pub enum WorkerQuicError {
BadUri(String),
Dns(String),
Connect(String),
Stream(String),
Frame(String),
Protocol(String),
UserException(Vec<u8>),
Timeout(f64),
}
impl std::fmt::Display for WorkerQuicError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::BadUri(s) => write!(f, "bad uri: {}", s),
Self::Dns(s) => write!(f, "dns: {}", s),
Self::Connect(s) => write!(f, "connect: {}", s),
Self::Stream(s) => write!(f, "stream: {}", s),
Self::Frame(s) => write!(f, "malformed frame: {}", s),
Self::Protocol(s) => write!(f, "protocol error from worker: {}", s),
Self::UserException(_) => write!(f, "user exception from worker"),
Self::Timeout(s) => write!(f, "timeout after {}s", s),
}
}
}
impl std::error::Error for WorkerQuicError {}
#[derive(Debug)]
struct SkipVerify;
impl rustls::client::danger::ServerCertVerifier for SkipVerify {
fn verify_server_cert(
&self,
_: &CertificateDer<'_>,
_: &[CertificateDer<'_>],
_: &ServerName<'_>,
_: &[u8],
_: UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_: &[u8],
_: &CertificateDer<'_>,
_: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_: &[u8],
_: &CertificateDer<'_>,
_: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
rustls::crypto::ring::default_provider()
.signature_verification_algorithms
.supported_schemes()
}
}
fn make_client_config() -> ClientConfig {
let mut cc = rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(SkipVerify))
.with_no_client_auth();
cc.alpn_protocols = vec![ALPN.to_vec()];
let quic_cc = quinn::crypto::rustls::QuicClientConfig::try_from(cc)
.expect("QUIC requires a rustls config supporting TLS 1.3");
ClientConfig::new(Arc::new(quic_cc))
}
struct Pool {
conns: DashMap<String, Arc<Mutex<Option<Connection>>>>,
endpoint: Mutex<Option<Endpoint>>,
}
impl Pool {
fn new() -> Self {
Self {
conns: DashMap::new(),
endpoint: Mutex::new(None),
}
}
async fn endpoint(&self) -> Result<Endpoint, WorkerQuicError> {
let mut slot = self.endpoint.lock().await;
if let Some(ep) = slot.as_ref() {
return Ok(ep.clone());
}
let addr: SocketAddr = "0.0.0.0:0"
.parse()
.map_err(|e: std::net::AddrParseError| WorkerQuicError::Connect(e.to_string()))?;
let mut endpoint =
Endpoint::client(addr).map_err(|e| WorkerQuicError::Connect(e.to_string()))?;
endpoint.set_default_client_config(make_client_config());
*slot = Some(endpoint.clone());
Ok(endpoint)
}
async fn get(&self, host: &str, port: u16) -> Result<Connection, WorkerQuicError> {
let key = format!("{}:{}", host, port);
let slot = self
.conns
.entry(key.clone())
.or_insert_with(|| Arc::new(Mutex::new(None)))
.clone();
let mut guard = slot.lock().await;
if let Some(conn) = guard.as_ref() {
if conn.close_reason().is_none() {
return Ok(conn.clone());
}
}
let endpoint = self.endpoint().await?;
let addr = tokio::net::lookup_host((host, port))
.await
.map_err(|e| WorkerQuicError::Dns(e.to_string()))?
.next()
.ok_or_else(|| WorkerQuicError::Dns(format!("no addr for {}", host)))?;
let connecting = endpoint
.connect(addr, "localhost")
.map_err(|e| WorkerQuicError::Connect(e.to_string()))?;
let connection = connecting
.await
.map_err(|e| WorkerQuicError::Connect(e.to_string()))?;
*guard = Some(connection.clone());
Ok(connection)
}
async fn invalidate(&self, host: &str, port: u16) {
let key = format!("{}:{}", host, port);
if let Some((_, slot)) = self.conns.remove(&key) {
let mut guard = slot.lock().await;
if let Some(conn) = guard.take() {
conn.close(VarInt::from_u32(0), b"stale");
}
}
}
}
pub struct WorkerQuicClient {
runtime: Runtime,
pool: Pool,
}
impl WorkerQuicClient {
pub fn new() -> Result<Self, WorkerQuicError> {
let runtime = Builder::new_multi_thread()
.worker_threads(2)
.thread_name("zc-worker-quic")
.enable_all()
.build()
.map_err(|e| WorkerQuicError::Connect(format!("runtime: {}", e)))?;
Ok(Self {
runtime,
pool: Pool::new(),
})
}
pub fn forward(
&self,
worker_uri: &str,
body: &[u8],
request_id: &str,
effective_timeout_secs: f64,
) -> Result<Vec<u8>, WorkerQuicError> {
self.runtime.block_on(async {
forward_async(
&self.pool,
worker_uri,
body,
request_id,
effective_timeout_secs,
)
.await
})
}
pub fn health(&self, worker_uri: &str) -> bool {
self.runtime
.block_on(async { health_async(&self.pool, worker_uri).await })
}
pub fn info(&self, worker_uri: &str) -> Result<String, WorkerQuicError> {
self.runtime
.block_on(async { info_async(&self.pool, worker_uri).await })
}
}
fn parse_uri(uri: &str) -> Result<(String, u16), WorkerQuicError> {
let stripped = uri
.strip_prefix("quic://")
.ok_or_else(|| WorkerQuicError::BadUri(format!("expected quic://, got {}", uri)))?;
let no_path = stripped.split('/').next().unwrap_or(stripped);
if let Some((h, p)) = no_path.rsplit_once(':') {
let port: u16 = p
.parse()
.map_err(|_| WorkerQuicError::BadUri(format!("bad port in {}", uri)))?;
Ok((h.to_string(), port))
} else {
Ok((no_path.to_string(), DEFAULT_PORT))
}
}
async fn write_frame(
send: &mut quinn::SendStream,
op: u8,
payload: &[u8],
) -> Result<(), WorkerQuicError> {
let len = u32::try_from(payload.len())
.map_err(|_| WorkerQuicError::Frame("payload > u32::MAX".into()))?;
let mut header = [0u8; 5];
header[0] = op;
header[1..5].copy_from_slice(&len.to_be_bytes());
send.write_all(&header)
.await
.map_err(|e| WorkerQuicError::Stream(e.to_string()))?;
if !payload.is_empty() {
send.write_all(payload)
.await
.map_err(|e| WorkerQuicError::Stream(e.to_string()))?;
}
send.finish()
.map_err(|e| WorkerQuicError::Stream(e.to_string()))?;
Ok(())
}
async fn read_frame(recv: &mut quinn::RecvStream) -> Result<(u8, Vec<u8>), WorkerQuicError> {
let mut header = [0u8; 5];
recv.read_exact(&mut header)
.await
.map_err(|e| WorkerQuicError::Frame(format!("header: {}", e)))?;
let status = header[0];
let len = u32::from_be_bytes([header[1], header[2], header[3], header[4]]) as usize;
let mut body = vec![0u8; len];
if len > 0 {
recv.read_exact(&mut body)
.await
.map_err(|e| WorkerQuicError::Frame(format!("body: {}", e)))?;
}
Ok((status, body))
}
async fn forward_async(
pool: &Pool,
worker_uri: &str,
body: &[u8],
_request_id: &str,
effective_timeout_secs: f64,
) -> Result<Vec<u8>, WorkerQuicError> {
let (host, port) = parse_uri(worker_uri)?;
let work = async {
let conn = pool.get(&host, port).await?;
let (mut send, mut recv) = match conn.open_bi().await {
Ok(s) => s,
Err(e) => {
pool.invalidate(&host, port).await;
return Err(WorkerQuicError::Stream(e.to_string()));
}
};
write_frame(&mut send, OP_EXECUTE, body).await?;
let (status, payload) = read_frame(&mut recv).await?;
match status {
STAT_OK => Ok(payload),
STAT_USER_ERROR => Err(WorkerQuicError::UserException(payload)),
STAT_PROTOCOL_ERROR => Err(WorkerQuicError::Protocol(
String::from_utf8_lossy(&payload).into_owned(),
)),
other => Err(WorkerQuicError::Frame(format!(
"unknown status byte {}",
other
))),
}
};
if effective_timeout_secs > 0.0 {
match tokio::time::timeout(Duration::from_secs_f64(effective_timeout_secs + 5.0), work)
.await
{
Ok(r) => r,
Err(_) => Err(WorkerQuicError::Timeout(effective_timeout_secs)),
}
} else {
work.await
}
}
async fn health_async(pool: &Pool, worker_uri: &str) -> bool {
async fn probe(pool: &Pool, worker_uri: &str) -> Result<(), WorkerQuicError> {
let (host, port) = parse_uri(worker_uri)?;
let conn = pool.get(&host, port).await?;
let (mut send, mut recv) = conn
.open_bi()
.await
.map_err(|e| WorkerQuicError::Stream(e.to_string()))?;
write_frame(&mut send, OP_HEALTH, &[]).await?;
let (status, _) = read_frame(&mut recv).await?;
if status == STAT_OK {
Ok(())
} else {
Err(WorkerQuicError::Protocol(format!("status={}", status)))
}
}
tokio::time::timeout(Duration::from_secs(2), probe(pool, worker_uri))
.await
.ok()
.and_then(|r| r.ok())
.is_some()
}
async fn info_async(pool: &Pool, worker_uri: &str) -> Result<String, WorkerQuicError> {
let (host, port) = parse_uri(worker_uri)?;
let conn = pool.get(&host, port).await?;
let (mut send, mut recv) = conn
.open_bi()
.await
.map_err(|e| WorkerQuicError::Stream(e.to_string()))?;
write_frame(&mut send, OP_INFO, &[]).await?;
let (status, payload) = read_frame(&mut recv).await?;
if status != STAT_OK {
return Err(WorkerQuicError::Protocol(
String::from_utf8_lossy(&payload).into_owned(),
));
}
String::from_utf8(payload).map_err(|e| WorkerQuicError::Frame(format!("info utf8: {}", e)))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_uri_ok() {
let (h, p) = parse_uri("quic://worker.local:4433").unwrap();
assert_eq!(h, "worker.local");
assert_eq!(p, 4433);
}
#[test]
fn parse_uri_default_port() {
let (h, p) = parse_uri("quic://worker.local").unwrap();
assert_eq!(h, "worker.local");
assert_eq!(p, DEFAULT_PORT);
}
#[test]
fn parse_uri_strips_trailing_path() {
let (h, p) = parse_uri("quic://worker.local:4433/execute").unwrap();
assert_eq!(h, "worker.local");
assert_eq!(p, 4433);
}
#[test]
fn parse_uri_rejects_http() {
assert!(parse_uri("http://worker.local").is_err());
}
#[test]
#[ignore]
fn network_roundtrip() {
let uri = std::env::var("ZAKURO_QUIC_TEST_URI")
.expect("set ZAKURO_QUIC_TEST_URI to a running QUIC worker");
let client = WorkerQuicClient::new().expect("client");
assert!(client.health(&uri), "health failed");
let info = client.info(&uri).expect("info");
assert!(
info.contains("\"transport\":\"quic\""),
"bad info: {}",
info
);
}
}