use std::net::SocketAddr;
use std::sync::{Arc, Mutex};
use dashmap::DashMap;
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;
fn generate_self_signed() -> (Vec<rustls::Certificate>, rustls::PrivateKey) {
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)) {
return (vec![rustls::Certificate(cert_der)], rustls::PrivateKey(key_der));
}
}
let cert = rcgen::generate_simple_self_signed(vec!["localhost".to_string()])
.expect("self-signed cert generation failed");
let cert_der = cert.serialize_der().expect("cert serialize failed");
let key_der = cert.serialize_private_key_der();
let _ = std::fs::create_dir_all(&dir);
let _ = std::fs::write(&cert_path, &cert_der);
let _ = std::fs::write(&key_path, &key_der);
return (vec![rustls::Certificate(cert_der)], rustls::PrivateKey(key_der));
}
let cert = rcgen::generate_simple_self_signed(vec!["localhost".to_string()])
.expect("self-signed cert generation failed");
let cert_der = cert.serialize_der().expect("cert serialize failed");
let key_der = cert.serialize_private_key_der();
(vec![rustls::Certificate(cert_der)], rustls::PrivateKey(key_der))
}
fn make_server_config(
certs: Vec<rustls::Certificate>,
key: rustls::PrivateKey,
) -> quinn::ServerConfig {
let mut sc = rustls::ServerConfig::builder()
.with_safe_defaults()
.with_no_client_auth()
.with_single_cert(certs, key)
.expect("bad server TLS config");
sc.alpn_protocols = vec![b"zk-quic".to_vec()];
quinn::ServerConfig::with_crypto(Arc::new(sc))
}
fn make_client_config() -> quinn::ClientConfig {
let mut cc = rustls::ClientConfig::builder()
.with_safe_defaults()
.with_custom_certificate_verifier(Arc::new(SkipVerify))
.with_no_client_auth();
cc.alpn_protocols = vec![b"zk-quic".to_vec()];
quinn::ClientConfig::new(Arc::new(cc))
}
struct SkipVerify;
impl rustls::client::ServerCertVerifier for SkipVerify {
fn verify_server_cert(
&self,
_end_entity: &rustls::Certificate,
_intermediates: &[rustls::Certificate],
_server_name: &rustls::ServerName,
_scts: &mut dyn Iterator<Item = &[u8]>,
_ocsp_response: &[u8],
_now: std::time::SystemTime,
) -> Result<rustls::client::ServerCertVerified, rustls::Error> {
Ok(rustls::client::ServerCertVerified::assertion())
}
}
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().await.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(())
}
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 = u32::from_be_bytes(len_buf) as usize;
if len > 64 * 1024 * 1024 {
return Err("message too large".into());
}
let mut buf = vec![0u8; len];
recv.read_exact(&mut buf)
.await
.map_err(|e| e.to_string())?;
Ok(buf)
}
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>,
}
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 (certs, key) = generate_self_signed();
let server_cfg = make_server_config(certs, key);
let client_cfg = make_client_config();
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(),
});
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 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());
}
}
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();
for (_, offer_tx) in &subs {
let _ = offer_tx.send((offer.clone(), reply_tx.clone()));
}
drop(reply_tx); let n = subs.len();
for _ in 0..n {
if let Ok((url, resp)) = reply_rx.recv() {
if let Ok(result) = resp {
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, mut 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(),
})
.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: Arc<dyn Fn(&BrokerState, &TaskOffer) -> Result<TaskResult, TaskReject> + Send + Sync>,
) {
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 format!("{}:{}", connect_host, quic_port).parse() {
Ok(a) => a,
Err(_) => 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: Arc<dyn Fn(&BrokerState, &TaskOffer) -> Result<TaskResult, TaskReject> + Send + Sync>,
) {
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,
};
if let Ok(sub) = serde_json::from_slice::<SubscribeMessage>(&msg) {
if sub.action == "subscribe" && !sub.peer_url.is_empty() {
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;
}
};
handle_stream(send, recv, offer, state).await;
}
fn subscriber_loop(
offer_rx: mpsc::Receiver<(TaskOffer, mpsc::Sender<(String, Result<TaskResult, TaskReject>)>)>,
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,
mut 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.uri.clone();
let pph = worker.pricing.price_per_hour;
let owner = state.config.owner_user_id.clone().unwrap_or_default();
let exec = tokio::task::spawn_blocking(move || {
let agent = ureq::AgentBuilder::new()
.timeout(std::time::Duration::from_secs_f64(timeout + 5.0))
.build();
let start = std::time::Instant::now();
let resp = agent
.post(&worker_uri)
.set("Content-Type", "application/octet-stream")
.set("X-Zakuro-Request-Id", &task_id)
.set("X-Zakuro-Timeout-Secs", &format!("{:.1}", timeout))
.send_bytes(&payload);
let duration_ms = start.elapsed().as_secs_f64() * 1000.0;
match resp {
Ok(r) => {
let worker_pid = r.header("X-Zakuro-Pid").map(|s| s.to_string());
let worker_ip = r.header("X-Zakuro-IP").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_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,
worker_ip,
})
}
Err(e) => Err(format!("worker: {}", e)),
}
})
.await;
let resp_bytes = match exec {
Ok(Ok(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;
}