use std::net::{SocketAddr, ToSocketAddrs};
use std::sync::{Arc, Mutex};
use dashmap::DashMap;
use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, ServerName, UnixTime};
use tokio::runtime::Runtime;
use std::collections::HashSet;
use std::sync::mpsc;
use std::thread;
use std::time::Duration;
use super::task_board::{SubscribeMessage, TaskOffer, TaskReject, TaskResult};
use super::BrokerState;
type ExecuteFn =
Arc<dyn Fn(&BrokerState, &TaskOffer) -> Result<TaskResult, TaskReject> + Send + Sync>;
type OfferReceiver = mpsc::Receiver<(
TaskOffer,
mpsc::Sender<(String, Result<TaskResult, TaskReject>)>,
)>;
fn parse_persisted_key(key_der: Vec<u8>) -> Option<PrivateKeyDer<'static>> {
PrivateKeyDer::try_from(key_der).ok()
}
fn quic_seed(peer_key: &str) -> [u8; 32] {
let mut input = Vec::with_capacity(12 + peer_key.len());
input.extend_from_slice(b"zc-quic-pin\0");
input.extend_from_slice(peer_key.as_bytes());
let digest = ring::digest::digest(&ring::digest::SHA256, &input);
let mut out = [0u8; 32];
out.copy_from_slice(digest.as_ref());
out
}
fn ed25519_pkcs8(seed: &[u8; 32]) -> Vec<u8> {
const PREFIX: [u8; 16] = [
0x30, 0x2e, 0x02, 0x01, 0x00, 0x30, 0x05, 0x06, 0x03, 0x2b, 0x65, 0x70, 0x04, 0x22, 0x04,
0x20,
];
let mut der = Vec::with_capacity(48);
der.extend_from_slice(&PREFIX);
der.extend_from_slice(seed);
der
}
fn derive_quic_identity(
peer_key: &str,
) -> Result<(CertificateDer<'static>, PrivateKeyDer<'static>), String> {
let seed = quic_seed(peer_key);
let pkcs8 = ed25519_pkcs8(&seed);
let pkcs8_der = PrivatePkcs8KeyDer::from(pkcs8.clone());
let key_pair = rcgen::KeyPair::from_pkcs8_der_and_sign_algo(&pkcs8_der, &rcgen::PKCS_ED25519)
.map_err(|e| format!("derive keypair: {}", e))?;
let mut params = rcgen::CertificateParams::new(vec!["zakuro-broker".to_string()])
.map_err(|e| format!("cert params: {}", e))?;
params.serial_number = Some(rcgen::SerialNumber::from_slice(&[0x01]));
params.not_before = rcgen::date_time_ymd(2020, 1, 1);
params.not_after = rcgen::date_time_ymd(2100, 1, 1);
params.distinguished_name = {
let mut dn = rcgen::DistinguishedName::new();
dn.push(rcgen::DnType::CommonName, "zakuro-broker");
dn
};
let cert = params
.self_signed(&key_pair)
.map_err(|e| format!("self-sign: {}", e))?;
let cert_der = CertificateDer::from(cert.der().to_vec());
let key = PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(pkcs8));
Ok((cert_der, key))
}
fn cert_sha256(cert: &CertificateDer) -> [u8; 32] {
let digest = ring::digest::digest(&ring::digest::SHA256, cert.as_ref());
let mut out = [0u8; 32];
out.copy_from_slice(digest.as_ref());
out
}
fn generate_self_signed() -> (Vec<CertificateDer<'static>>, PrivateKeyDer<'static>) {
if let Some(home) = std::env::var("HOME")
.ok()
.or_else(|| std::env::var("USERPROFILE").ok())
{
let dir = std::path::Path::new(&home).join(".zakuro");
let cert_path = dir.join("quic_cert.der");
let key_path = dir.join("quic_key.der");
if cert_path.exists() && key_path.exists() {
if let (Ok(cert_der), Ok(key_der)) =
(std::fs::read(&cert_path), std::fs::read(&key_path))
{
if let Some(key) = parse_persisted_key(key_der) {
return (vec![CertificateDer::from(cert_der)], key);
}
eprintln!("[quic] persisted QUIC key is invalid; regenerating self-signed cert");
}
}
let cert = rcgen::generate_simple_self_signed(vec!["localhost".to_string()])
.expect("self-signed cert generation failed");
let cert_der: Vec<u8> = cert.cert.der().to_vec();
let key_der: Vec<u8> = cert.key_pair.serialize_der();
let _ = std::fs::create_dir_all(&dir);
let _ = std::fs::write(&cert_path, &cert_der);
let _ = std::fs::write(&key_path, &key_der);
if let Some(key) = parse_persisted_key(key_der) {
return (vec![CertificateDer::from(cert_der)], key);
}
}
let cert = rcgen::generate_simple_self_signed(vec!["localhost".to_string()])
.expect("self-signed cert generation failed");
let cert_der: Vec<u8> = cert.cert.der().to_vec();
let key_der: Vec<u8> = cert.key_pair.serialize_der();
(
vec![CertificateDer::from(cert_der)],
PrivateKeyDer::try_from(key_der).expect("rcgen-generated QUIC key was not parseable"),
)
}
fn make_server_config(
certs: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
) -> quinn::ServerConfig {
let mut sc = rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(certs, key)
.expect("bad server TLS config");
sc.alpn_protocols = vec![b"zk-quic".to_vec()];
let quic_sc = quinn::crypto::rustls::QuicServerConfig::try_from(sc)
.expect("QUIC requires a rustls config supporting TLS 1.3");
let mut cfg = quinn::ServerConfig::with_crypto(Arc::new(quic_sc));
cfg.transport_config(Arc::new(task_feed_transport_config()));
cfg
}
fn task_feed_transport_config() -> quinn::TransportConfig {
let mut tc = quinn::TransportConfig::default();
tc.keep_alive_interval(Some(Duration::from_secs(5)));
tc.max_idle_timeout(Some(
Duration::from_secs(30)
.try_into()
.expect("30s is a valid QUIC idle timeout"),
));
tc
}
fn resolve_peer_addr(host: &str, port: u16) -> Option<SocketAddr> {
let candidates: Vec<SocketAddr> = (host, port).to_socket_addrs().ok()?.collect();
candidates
.iter()
.find(|a| a.is_ipv4())
.or_else(|| candidates.first())
.copied()
}
fn make_client_config(
verifier: Arc<dyn rustls::client::danger::ServerCertVerifier>,
) -> quinn::ClientConfig {
let mut cc = rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(verifier)
.with_no_client_auth();
cc.alpn_protocols = vec![b"zk-quic".to_vec()];
let quic_cc = quinn::crypto::rustls::QuicClientConfig::try_from(cc)
.expect("QUIC requires a rustls config supporting TLS 1.3");
let mut cfg = quinn::ClientConfig::new(Arc::new(quic_cc));
cfg.transport_config(Arc::new(task_feed_transport_config()));
cfg
}
fn build_tls_configs(peer_key: &str) -> (quinn::ServerConfig, quinn::ClientConfig) {
if !peer_key.is_empty() {
match derive_quic_identity(peer_key) {
Ok((cert, key)) => {
let pin = cert_sha256(&cert);
let server_cfg = make_server_config(vec![cert], key);
let client_cfg = make_client_config(Arc::new(PinnedVerify { expected: pin }));
return (server_cfg, client_cfg);
}
Err(e) => {
eprintln!(
"[quic] peer-key cert derivation failed ({e}); falling back to unpinned TLS"
);
}
}
}
let (certs, key) = generate_self_signed();
(
make_server_config(certs, key),
make_client_config(Arc::new(SkipVerify)),
)
}
#[derive(Debug)]
struct SkipVerify;
impl rustls::client::danger::ServerCertVerifier for SkipVerify {
fn verify_server_cert(
&self,
_end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp_response: &[u8],
_now: UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &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()
}
}
#[derive(Debug)]
struct PinnedVerify {
expected: [u8; 32],
}
impl rustls::client::danger::ServerCertVerifier for PinnedVerify {
fn verify_server_cert(
&self,
end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp_response: &[u8],
_now: UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
if cert_sha256(end_entity) == self.expected {
Ok(rustls::client::danger::ServerCertVerified::assertion())
} else {
Err(rustls::Error::InvalidCertificate(
rustls::CertificateError::ApplicationVerificationFailure,
))
}
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &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()
}
}
async fn write_msg(send: &mut quinn::SendStream, payload: &[u8]) -> Result<(), String> {
let len = (payload.len() as u32).to_be_bytes();
send.write_all(&len).await.map_err(|e| e.to_string())?;
send.write_all(payload).await.map_err(|e| e.to_string())?;
send.finish().map_err(|e| e.to_string())?;
Ok(())
}
async fn write_msg_keep_open(send: &mut quinn::SendStream, payload: &[u8]) -> Result<(), String> {
let len = (payload.len() as u32).to_be_bytes();
send.write_all(&len).await.map_err(|e| e.to_string())?;
send.write_all(payload).await.map_err(|e| e.to_string())?;
Ok(())
}
const MAX_QUIC_MSG_BYTES: usize = 8 * 1024 * 1024;
fn cap_msg_len(len_bytes: [u8; 4]) -> Result<usize, String> {
let len = u32::from_be_bytes(len_bytes) as usize;
if len > MAX_QUIC_MSG_BYTES {
return Err(format!(
"message too large: {} bytes (max {})",
len, MAX_QUIC_MSG_BYTES
));
}
Ok(len)
}
async fn read_msg(recv: &mut quinn::RecvStream) -> Result<Vec<u8>, String> {
let mut len_buf = [0u8; 4];
recv.read_exact(&mut len_buf)
.await
.map_err(|e| e.to_string())?;
let len = cap_msg_len(len_buf)?;
const CHUNK: usize = 64 * 1024;
let mut buf: Vec<u8> = Vec::with_capacity(len.min(CHUNK));
let mut remaining = len;
let mut chunk = [0u8; CHUNK];
while remaining > 0 {
let want = remaining.min(CHUNK);
recv.read_exact(&mut chunk[..want])
.await
.map_err(|e| e.to_string())?;
buf.extend_from_slice(&chunk[..want]);
remaining -= want;
}
Ok(buf)
}
fn peer_key_ok(expected: &str, provided: &str) -> bool {
expected.is_empty() || expected == provided
}
type SubscriberOfferTx = mpsc::Sender<(
TaskOffer,
mpsc::Sender<(String, Result<TaskResult, TaskReject>)>,
)>;
pub struct QuicTransport {
runtime: Runtime,
endpoint: quinn::Endpoint,
connections: DashMap<SocketAddr, quinn::Connection>,
subscriber_txs: DashMap<String, SubscriberOfferTx>,
peer_key: String,
}
impl QuicTransport {
pub fn new(bind_addr: SocketAddr, state: Arc<BrokerState>) -> Result<Arc<Self>, String> {
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(2)
.enable_all()
.build()
.map_err(|e| format!("tokio runtime: {}", e))?;
let peer_key = state.peer_manager.peer_key().to_string();
let (server_cfg, client_cfg) = build_tls_configs(&peer_key);
let endpoint = runtime.block_on(async {
let mut ep = quinn::Endpoint::server(server_cfg, bind_addr)
.map_err(|e| format!("QUIC bind {}: {}", bind_addr, e))?;
ep.set_default_client_config(client_cfg);
Ok::<_, String>(ep)
})?;
let transport = Arc::new(Self {
runtime,
endpoint,
connections: DashMap::new(),
subscriber_txs: DashMap::new(),
peer_key,
});
let t = transport.clone();
transport.runtime.spawn(async move {
accept_loop(t, state).await;
});
Ok(transport)
}
pub fn offer_task(
&self,
peer_addr: SocketAddr,
offer: &TaskOffer,
) -> Result<TaskResult, String> {
self.runtime
.block_on(self.offer_task_inner(peer_addr, offer))
}
async fn offer_task_inner(
&self,
peer_addr: SocketAddr,
offer: &TaskOffer,
) -> Result<TaskResult, String> {
let conn = self.get_or_connect(peer_addr).await?;
let (mut send, mut recv) = conn
.open_bi()
.await
.map_err(|e| format!("open stream: {}", e))?;
let mut offer = offer.clone();
offer.peer_key = self.peer_key.clone();
let payload = serde_json::to_vec(&offer).map_err(|e| e.to_string())?;
write_msg(&mut send, &payload).await?;
let resp = read_msg(&mut recv).await?;
if let Ok(result) = serde_json::from_slice::<TaskResult>(&resp) {
return Ok(result);
}
if let Ok(reject) = serde_json::from_slice::<TaskReject>(&resp) {
return Err(format!("peer rejected: {}", reject.reason));
}
Err("invalid peer response".into())
}
async fn get_or_connect(&self, addr: SocketAddr) -> Result<quinn::Connection, String> {
if let Some(c) = self.connections.get(&addr) {
if c.close_reason().is_none() {
return Ok(c.clone());
}
}
self.connections.remove(&addr);
let conn = self
.endpoint
.connect(addr, "localhost")
.map_err(|e| format!("connect {}: {}", addr, e))?
.await
.map_err(|e| format!("handshake {}: {}", addr, e))?;
self.connections.insert(addr, conn.clone());
Ok(conn)
}
pub fn local_addr(&self) -> Option<SocketAddr> {
self.endpoint.local_addr().ok()
}
pub fn broadcast_offer_to_subscribers(
&self,
offer: &TaskOffer,
) -> Option<(String, TaskResult)> {
let subs: Vec<(String, SubscriberOfferTx)> = self
.subscriber_txs
.iter()
.map(|e| (e.key().clone(), e.value().clone()))
.collect();
if subs.is_empty() {
return None;
}
let (reply_tx, reply_rx) = mpsc::channel();
let mut dead: Vec<String> = Vec::new();
for (url, offer_tx) in &subs {
if offer_tx.send((offer.clone(), reply_tx.clone())).is_err() {
dead.push(url.clone());
}
}
for url in &dead {
self.subscriber_txs.remove(url);
eprintln!(" [QUIC] pruned disconnected subscriber {}", url);
}
drop(reply_tx); let n = subs.len() - dead.len();
for _ in 0..n {
if let Ok((url, Ok(result))) = reply_rx.recv() {
return Some((url, result));
}
}
None
}
pub fn subscriber_urls(&self) -> Vec<String> {
self.subscriber_txs
.iter()
.map(|e| e.key().clone())
.collect()
}
pub fn connect_and_subscribe(
&self,
peer_addr: SocketAddr,
our_url: &str,
) -> Result<(quinn::SendStream, quinn::RecvStream), String> {
self.runtime
.block_on(self.connect_and_subscribe_async(peer_addr, our_url))
}
async fn connect_and_subscribe_async(
&self,
peer_addr: SocketAddr,
our_url: &str,
) -> Result<(quinn::SendStream, quinn::RecvStream), String> {
let conn = self
.endpoint
.connect(peer_addr, "localhost")
.map_err(|e| format!("connect: {}", e))?
.await
.map_err(|e| format!("handshake: {}", e))?;
let (mut send, recv) = conn
.open_bi()
.await
.map_err(|e| format!("open stream: {}", e))?;
let msg = serde_json::to_vec(&SubscribeMessage {
action: "subscribe".to_string(),
peer_url: our_url.to_string(),
peer_key: self.peer_key.clone(),
})
.map_err(|e| e.to_string())?;
write_msg_keep_open(&mut send, &msg).await?;
Ok((send, recv))
}
}
pub fn run_subscription_client(state: Arc<BrokerState>, our_url: String, execute_fn: ExecuteFn) {
let spawned = Arc::new(Mutex::new(HashSet::<String>::new()));
thread::spawn(move || {
let mut interval = 0u32;
loop {
thread::sleep(Duration::from_millis(if interval < 20 {
2000
} else {
10000
}));
interval = interval.saturating_add(1);
let transport = match state.quic.get().cloned() {
Some(t) => t,
None => continue,
};
let peer_urls: Vec<String> = state.peer_manager.peer_urls();
for peer_url in peer_urls {
if peer_url == our_url {
continue;
}
let quic_port = state.peer_manager.get_quic_port(&peer_url);
if quic_port == 0 {
continue;
}
let url_trimmed = peer_url
.trim_start_matches("http://")
.trim_start_matches("https://");
let host = url_trimmed.split(':').next().unwrap_or("127.0.0.1");
let connect_host = if host.starts_with("127.") {
"127.0.0.1"
} else {
host
};
let addr: SocketAddr = match resolve_peer_addr(connect_host, quic_port) {
Some(a) => a,
None => {
eprintln!(
" [QUIC subscribe] {}:{}: could not resolve — skipping",
connect_host, quic_port
);
continue;
}
};
{
let mut g = spawned.lock().unwrap();
if g.contains(&peer_url) {
continue;
}
g.insert(peer_url.clone());
}
let transport_c = transport.clone();
let state_c = state.clone();
let our_url_c = our_url.clone();
let execute_fn_c = execute_fn.clone();
let peer_url_c = peer_url.clone();
thread::spawn(move || {
run_subscription_to_peer(
transport_c,
state_c,
addr,
&peer_url_c,
&our_url_c,
execute_fn_c,
);
});
}
}
});
}
fn run_subscription_to_peer(
transport: Arc<QuicTransport>,
state: Arc<BrokerState>,
peer_addr: SocketAddr,
_peer_url: &str,
our_url: &str,
execute_fn: ExecuteFn,
) {
let runtime = transport.runtime.handle().clone();
loop {
let (mut send, mut recv) = match transport.connect_and_subscribe(peer_addr, our_url) {
Ok(pair) => pair,
Err(e) => {
eprintln!(" [QUIC subscribe] {}: {}", peer_addr, e);
thread::sleep(Duration::from_secs(5));
continue;
}
};
eprintln!(" [QUIC subscribe] Subscribed to {} (task feed)", peer_addr);
while let Ok(msg) = runtime.block_on(read_msg(&mut recv)) {
let offer: TaskOffer = match serde_json::from_slice(&msg) {
Ok(o) => o,
Err(_) => continue,
};
let resp = match execute_fn(&state, &offer) {
Ok(r) => serde_json::to_vec(&r).unwrap_or_default(),
Err(rej) => serde_json::to_vec(&rej).unwrap_or_default(),
};
if runtime
.block_on(write_msg_keep_open(&mut send, &resp))
.is_err()
{
break;
}
}
thread::sleep(Duration::from_secs(2));
}
}
impl Drop for QuicTransport {
fn drop(&mut self) {
self.endpoint.close(0u32.into(), b"shutdown");
}
}
async fn accept_loop(transport: Arc<QuicTransport>, state: Arc<BrokerState>) {
loop {
let incoming = match transport.endpoint.accept().await {
Some(i) => i,
None => break,
};
let transport_clone = transport.clone();
let state_clone = state.clone();
tokio::spawn(async move {
let conn = match incoming.await {
Ok(c) => c,
Err(e) => {
eprintln!(" [QUIC] accept error: {}", e);
return;
}
};
loop {
let stream = match conn.accept_bi().await {
Ok(s) => s,
Err(quinn::ConnectionError::ApplicationClosed(_)) => break,
Err(_) => break,
};
let t = transport_clone.clone();
let st = state_clone.clone();
tokio::spawn(async move {
handle_incoming_stream(stream, t, st).await;
});
}
});
}
}
async fn handle_incoming_stream(
(mut send, mut recv): (quinn::SendStream, quinn::RecvStream),
transport: Arc<QuicTransport>,
state: Arc<BrokerState>,
) {
let msg = match read_msg(&mut recv).await {
Ok(m) => m,
Err(_) => return,
};
let expected_key = state.peer_manager.peer_key();
if let Ok(sub) = serde_json::from_slice::<SubscribeMessage>(&msg) {
if sub.action == "subscribe" && !sub.peer_url.is_empty() {
if !peer_key_ok(expected_key, &sub.peer_key) {
eprintln!(
" [QUIC] rejected subscribe from {}: bad peer_key",
sub.peer_url
);
return;
}
let peer_url = sub.peer_url.clone();
let (offer_tx, offer_rx) = mpsc::channel();
let runtime_handle = transport.runtime.handle().clone();
thread::spawn(move || {
subscriber_loop(offer_rx, send, recv, peer_url, runtime_handle);
});
transport.subscriber_txs.insert(sub.peer_url, offer_tx);
return;
}
}
let offer: TaskOffer = match serde_json::from_slice(&msg) {
Ok(o) => o,
Err(e) => {
let r = TaskReject {
task_id: String::new(),
reason: format!("bad payload: {}", e),
};
let _ = write_msg(&mut send, &serde_json::to_vec(&r).unwrap()).await;
return;
}
};
if !peer_key_ok(expected_key, &offer.peer_key) {
let r = TaskReject {
task_id: offer.task_id.clone(),
reason: "unauthorized peer".into(),
};
let _ = write_msg(&mut send, &serde_json::to_vec(&r).unwrap()).await;
return;
}
handle_stream(send, recv, offer, state).await;
}
fn subscriber_loop(
offer_rx: OfferReceiver,
mut send: quinn::SendStream,
mut recv: quinn::RecvStream,
peer_url: String,
runtime_handle: tokio::runtime::Handle,
) {
while let Ok((offer, reply_tx)) = offer_rx.recv() {
let payload = match serde_json::to_vec(&offer) {
Ok(p) => p,
Err(_) => continue,
};
let write_read = async {
write_msg_keep_open(&mut send, &payload).await?;
read_msg(&mut recv).await
};
let resp_bytes = match runtime_handle.block_on(write_read) {
Ok(b) => b,
Err(_) => break,
};
let result = if let Ok(r) = serde_json::from_slice::<TaskResult>(&resp_bytes) {
Ok(r)
} else if let Ok(rej) = serde_json::from_slice::<TaskReject>(&resp_bytes) {
Err(rej)
} else {
continue;
};
let _ = reply_tx.send((peer_url.clone(), result));
}
}
async fn handle_stream(
mut send: quinn::SendStream,
_recv: quinn::RecvStream,
offer: TaskOffer,
state: Arc<BrokerState>,
) {
let local_workers: Vec<_> = state
.workers
.healthy()
.into_iter()
.filter(|w| state.is_local_worker(&w.uri))
.filter(|w| w.pricing.price_per_hour <= offer.max_price_per_hour)
.filter(|w| {
offer
.worker_type
.as_ref()
.map(|wt| w.worker_type.as_str() == wt.as_str())
.unwrap_or(true)
})
.collect();
if local_workers.is_empty() {
let r = TaskReject {
task_id: offer.task_id,
reason: "no matching local worker".into(),
};
let _ = write_msg(&mut send, &serde_json::to_vec(&r).unwrap()).await;
return;
}
let idx = {
use std::time::SystemTime;
SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap()
.subsec_nanos() as usize
% local_workers.len()
};
let worker = local_workers[idx].clone();
let payload = match base64::Engine::decode(
&base64::engine::general_purpose::STANDARD,
&offer.payload_b64,
) {
Ok(p) => p,
Err(e) => {
let r = TaskReject {
task_id: offer.task_id,
reason: format!("bad base64: {}", e),
};
let _ = write_msg(&mut send, &serde_json::to_vec(&r).unwrap()).await;
return;
}
};
let worker_uri = format!("{}/execute", worker.uri.trim_end_matches('/'));
let timeout = if offer.timeout_secs > 0.0 {
offer.timeout_secs
} else {
300.0
};
let task_id = offer.task_id.clone();
let wname = worker.name.clone();
let wuri = worker.zc_uri();
let pph = worker.pricing.price_per_hour;
let owner = state.config.owner_user_id.clone().unwrap_or_default();
let node_uri = state.node_key.node_uri();
let exec = tokio::task::spawn_blocking(move || {
let agent = ureq::Agent::new_with_config(
ureq::Agent::config_builder()
.timeout_global(Some(std::time::Duration::from_secs_f64(timeout + 5.0)))
.build(),
);
let start = std::time::Instant::now();
let resp = agent
.post(&worker_uri)
.header("Content-Type", "application/octet-stream")
.header("X-Zakuro-Request-Id", &task_id)
.header("X-Zakuro-Timeout-Secs", &format!("{:.1}", timeout))
.send(&payload);
let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
match resp {
Ok(r) => {
let worker_pid = r
.headers()
.get("X-Zakuro-Pid")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
let actual_cost = (pph / 3600.0 * duration_ms / 1000.0).max(0.001);
let mut body = Vec::new();
let _ = std::io::Read::read_to_end(&mut r.into_body().into_reader(), &mut body);
let b64 = base64::Engine::encode(&base64::engine::general_purpose::STANDARD, &body);
Ok(TaskResult {
task_id: task_id.to_string(),
payload_b64: b64,
duration_ms,
actual_cost,
worker_name: wname,
worker_uri: wuri,
price_per_hour: pph,
executor_owner: owner,
worker_pid,
executor_node_uri: node_uri,
executor_node_id: None,
executor_sig: None,
})
}
Err(e) => Err(format!("worker: {}", e)),
}
})
.await;
let resp_bytes = match exec {
Ok(Ok(mut res)) => {
res.executor_node_id = Some(state.node_key.public_b64());
res.executor_sig = Some(
state
.node_key
.sign(&crate::broker::task_board::receipt_bytes(&res)),
);
serde_json::to_vec(&res).unwrap()
}
Ok(Err(e)) => serde_json::to_vec(&TaskReject {
task_id: offer.task_id,
reason: e,
})
.unwrap(),
Err(e) => serde_json::to_vec(&TaskReject {
task_id: offer.task_id,
reason: format!("internal: {}", e),
})
.unwrap(),
};
let _ = write_msg(&mut send, &resp_bytes).await;
}
#[cfg(test)]
mod panic_fix_tests {
use super::{cap_msg_len, parse_persisted_key, peer_key_ok};
use crate::broker::task_board::{SubscribeMessage, TaskOffer};
#[test]
fn parse_persisted_key_rejects_garbage_without_panicking() {
assert!(parse_persisted_key(vec![0u8, 1, 2, 3]).is_none());
assert!(parse_persisted_key(Vec::new()).is_none());
}
#[test]
fn cap_msg_len_bounds() {
assert_eq!(cap_msg_len(1024u32.to_be_bytes()).unwrap(), 1024usize);
assert_eq!(cap_msg_len(0u32.to_be_bytes()).unwrap(), 0usize);
let over = (super::MAX_QUIC_MSG_BYTES as u32) + 1;
assert!(cap_msg_len(over.to_be_bytes()).is_err());
}
#[test]
fn empty_expected_allows_any_provided() {
assert!(peer_key_ok("", ""));
assert!(peer_key_ok("", "whatever"));
}
#[test]
fn matching_key_is_accepted() {
assert!(peer_key_ok("s3cret", "s3cret"));
}
#[test]
fn mismatched_or_missing_key_is_rejected() {
assert!(!peer_key_ok("s3cret", "wrong"));
assert!(!peer_key_ok("s3cret", ""));
}
#[test]
fn task_offer_peer_key_defaults_when_absent() {
let json = r#"{
"task_id":"t1","payload_b64":"","max_price_per_hour":1.0,
"estimated_duration_secs":1.0,"timeout_secs":10.0,
"requester_user_id":"u","source_broker":"n"
}"#;
let offer: TaskOffer = serde_json::from_str(json).unwrap();
assert_eq!(offer.peer_key, "");
}
#[test]
fn task_offer_peer_key_round_trips() {
let json = r#"{
"task_id":"t1","payload_b64":"","max_price_per_hour":1.0,
"estimated_duration_secs":1.0,"timeout_secs":10.0,
"requester_user_id":"u","source_broker":"n","peer_key":"abc"
}"#;
let offer: TaskOffer = serde_json::from_str(json).unwrap();
assert_eq!(offer.peer_key, "abc");
let back = serde_json::to_string(&offer).unwrap();
assert!(back.contains("\"peer_key\":\"abc\""));
}
#[test]
fn subscribe_message_peer_key_defaults_and_round_trips() {
let absent: SubscribeMessage =
serde_json::from_str(r#"{"action":"subscribe","peer_url":"http://x"}"#).unwrap();
assert_eq!(absent.peer_key, "");
let present: SubscribeMessage =
serde_json::from_str(r#"{"action":"subscribe","peer_url":"http://x","peer_key":"k"}"#)
.unwrap();
assert_eq!(present.peer_key, "k");
}
#[test]
fn dead_subscribers_are_pruned_on_send_failure() {
use dashmap::DashMap;
use std::sync::mpsc;
let subs: DashMap<String, mpsc::Sender<u8>> = DashMap::new();
let (live_tx, _live_rx) = mpsc::channel::<u8>();
let (dead_tx, dead_rx) = mpsc::channel::<u8>();
subs.insert("live".into(), live_tx);
subs.insert("dead".into(), dead_tx);
drop(dead_rx);
let snapshot: Vec<(String, mpsc::Sender<u8>)> = subs
.iter()
.map(|e| (e.key().clone(), e.value().clone()))
.collect();
let mut dead = Vec::new();
for (url, tx) in &snapshot {
if tx.send(1).is_err() {
dead.push(url.clone());
}
}
for url in &dead {
subs.remove(url);
}
assert!(subs.contains_key("live"));
assert!(!subs.contains_key("dead"), "dead subscriber must be pruned");
}
#[test]
fn quic_seed_is_stable_and_key_separated() {
assert_eq!(super::quic_seed("x"), super::quic_seed("x"));
assert_ne!(super::quic_seed("x"), super::quic_seed("y"));
}
#[test]
fn derive_quic_identity_is_deterministic() {
let (c1, _k1) = super::derive_quic_identity("shared-secret").expect("derive");
let (c2, _k2) = super::derive_quic_identity("shared-secret").expect("derive");
assert_eq!(
c1.as_ref(),
c2.as_ref(),
"same peer_key must yield byte-identical cert"
);
assert_eq!(super::cert_sha256(&c1), super::cert_sha256(&c2));
}
#[test]
fn derive_quic_identity_differs_by_key() {
let (a, _) = super::derive_quic_identity("key-a").expect("derive a");
let (b, _) = super::derive_quic_identity("key-b").expect("derive b");
assert_ne!(
super::cert_sha256(&a),
super::cert_sha256(&b),
"different peer_key must yield a different cert"
);
}
#[test]
fn derived_key_is_a_valid_private_key() {
let (_c, k) = super::derive_quic_identity("shared-secret").expect("derive");
assert!(matches!(k, rustls::pki_types::PrivateKeyDer::Pkcs8(_)));
}
#[test]
fn pinned_verify_accepts_matching_cert() {
use rustls::client::danger::ServerCertVerifier;
use rustls::pki_types::{ServerName, UnixTime};
let (cert, _k) = super::derive_quic_identity("shared-secret").expect("derive");
let v = super::PinnedVerify {
expected: super::cert_sha256(&cert),
};
let name = ServerName::try_from("localhost").unwrap();
let r = v.verify_server_cert(&cert, &[], &name, &[], UnixTime::now());
assert!(r.is_ok(), "matching cert must verify");
}
#[test]
fn pinned_verify_rejects_mismatched_cert() {
use rustls::client::danger::ServerCertVerifier;
use rustls::pki_types::{ServerName, UnixTime};
let (good, _) = super::derive_quic_identity("shared-secret").expect("derive");
let (attacker, _) = super::derive_quic_identity("attacker-key").expect("derive");
let v = super::PinnedVerify {
expected: super::cert_sha256(&good),
};
let name = ServerName::try_from("localhost").unwrap();
let r = v.verify_server_cert(&attacker, &[], &name, &[], UnixTime::now());
assert!(r.is_err(), "a cert not matching the pin must be rejected");
}
#[test]
fn build_tls_configs_keyed_and_keyless() {
let _ = super::build_tls_configs("shared-secret");
let _ = super::build_tls_configs("");
}
#[test]
fn task_feed_config_keeps_idle_connections_alive() {
let rendered = format!("{:?}", super::task_feed_transport_config());
assert!(
rendered.contains("keep_alive_interval: Some("),
"task feed needs a keepalive or idle subscriptions die: {rendered}"
);
assert!(
rendered.contains("max_idle_timeout: Some("),
"task feed needs an explicit idle timeout so it does not rely on a \
default that may be shorter than the keepalive: {rendered}"
);
}
#[test]
fn resolve_peer_addr_accepts_literal_ips() {
let addr = super::resolve_peer_addr("10.13.13.15", 9001)
.expect("a literal IP must resolve without a DNS lookup");
assert_eq!(addr.to_string(), "10.13.13.15:9001");
}
#[test]
fn resolve_peer_addr_accepts_dns_names() {
let addr = super::resolve_peer_addr("localhost", 9001)
.expect("a DNS name must resolve — parse::<SocketAddr>() cannot do this");
assert_eq!(addr.port(), 9001);
assert!(
addr.ip().is_loopback(),
"localhost should be loopback: {addr}"
);
}
#[test]
fn resolve_peer_addr_prefers_ipv4() {
let addr = super::resolve_peer_addr("localhost", 9001).expect("localhost resolves");
assert!(addr.is_ipv4(), "must prefer the v4 address, got {addr}");
}
#[test]
fn resolve_peer_addr_rejects_unresolvable_names() {
assert!(super::resolve_peer_addr("no-such-peer.invalid", 9001).is_none());
}
}