1use std::collections::HashMap;
18use std::net::SocketAddr;
19
20use serde::de::DeserializeOwned;
21use serde::{Deserialize, Serialize};
22use tokio::io::{AsyncReadExt, AsyncWriteExt};
23use tokio::net::{TcpListener, TcpStream};
24use tokio::runtime::Handle;
25
26use fv_streams_exchange::Link;
27
28use crate::dataflow::WorkerExchange;
29use crate::placement::{Members, WorkerId};
30
31#[derive(Debug, Serialize, Deserialize)]
33struct Register {
34 exchange_addr: String,
35 wire: u8,
37}
38
39#[derive(Debug, Serialize, Deserialize)]
42struct Assignment {
43 workers: Vec<String>,
44 fingerprint: String,
45 wire: u8,
47}
48
49fn wire_mismatch(who: &str, got: u8, mine: u8) -> std::io::Error {
51 invalid(format!(
52 "{who} speaks wire format v{got}, this worker v{mine} — every worker must run the same fv-streams version"
53 ))
54}
55
56#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
63pub enum ToWorker {
64 Barrier(u64),
65 Commit(u64),
66 Stop,
67}
68
69#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
75pub enum ToLeader {
76 Acked {
77 epoch: u64,
78 contribution: fv_streams_state::Manifest,
79 },
80 Finished,
81}
82
83pub struct LeaderCtl {
87 conns: Vec<(WorkerId, TcpStream)>,
88}
89
90impl LeaderCtl {
91 pub async fn broadcast(&mut self, msg: ToWorker) -> std::io::Result<()> {
93 for (_, s) in self.conns.iter_mut() {
94 write_msg(s, &msg).await?;
95 }
96 Ok(())
97 }
98
99 pub async fn collect(&mut self) -> std::io::Result<Vec<(WorkerId, ToLeader)>> {
102 let mut out = Vec::with_capacity(self.conns.len());
103 for (w, s) in self.conns.iter_mut() {
104 out.push((*w, read_msg(s).await?));
105 }
106 Ok(out)
107 }
108}
109
110pub struct WorkerCtl {
112 conn: TcpStream,
113}
114
115impl WorkerCtl {
116 pub async fn recv(&mut self) -> std::io::Result<ToWorker> {
118 read_msg(&mut self.conn).await
119 }
120
121 pub async fn report(&mut self, msg: ToLeader) -> std::io::Result<()> {
123 write_msg(&mut self.conn, &msg).await
124 }
125}
126
127pub enum Coord {
130 Leader(LeaderCtl),
131 Worker(WorkerCtl),
132}
133
134fn invalid<E: std::fmt::Display>(e: E) -> std::io::Error {
135 std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string())
136}
137
138async fn write_msg<T: Serialize>(s: &mut TcpStream, msg: &T) -> std::io::Result<()> {
140 let bytes = postcard::to_allocvec(msg).map_err(invalid)?;
141 s.write_all(&(bytes.len() as u32).to_le_bytes()).await?;
142 s.write_all(&bytes).await?;
143 s.flush().await
144}
145
146async fn read_msg<T: DeserializeOwned>(s: &mut TcpStream) -> std::io::Result<T> {
148 let mut len = [0u8; 4];
149 s.read_exact(&mut len).await?;
150 let n = u32::from_le_bytes(len) as usize;
151 let mut buf = vec![0u8; n];
152 s.read_exact(&mut buf).await?;
153 postcard::from_bytes(&buf).map_err(invalid)
154}
155
156pub async fn lead(
161 control: &TcpListener,
162 my_exchange_addr: &str,
163 n_workers: usize,
164 fingerprint: &str,
165) -> std::io::Result<(Members, LeaderCtl)> {
166 let mut regs: Vec<(String, TcpStream)> = Vec::new();
168 let mut addrs = vec![my_exchange_addr.to_string()];
169 let mut foreign: Option<(String, u8)> = None;
170 while addrs.len() < n_workers {
171 let (mut s, _) = control.accept().await?;
172 let reg: Register = read_msg(&mut s).await?;
173 if reg.wire != fv_streams_exchange::WIRE_VERSION {
174 foreign.get_or_insert((reg.exchange_addr.clone(), reg.wire));
175 }
176 addrs.push(reg.exchange_addr.clone());
177 regs.push((reg.exchange_addr, s));
178 }
179 addrs.sort();
180 addrs.dedup();
181 let members = Members::new(addrs.clone(), my_exchange_addr).map_err(invalid)?;
182 let assign = Assignment {
183 workers: addrs.clone(),
184 fingerprint: fingerprint.to_string(),
185 wire: fv_streams_exchange::WIRE_VERSION,
186 };
187 let mut conns = Vec::with_capacity(regs.len());
190 for (addr, mut s) in regs {
191 write_msg(&mut s, &assign).await?;
192 let wid = addrs
193 .iter()
194 .position(|a| *a == addr)
195 .expect("registered address is a member") as WorkerId;
196 conns.push((wid, s));
197 }
198 conns.sort_by_key(|(w, _)| *w);
199 if let Some((addr, got)) = foreign {
200 return Err(wire_mismatch(
201 &format!("the worker at {addr}"),
202 got,
203 fv_streams_exchange::WIRE_VERSION,
204 ));
205 }
206 Ok((members, LeaderCtl { conns }))
207}
208
209#[derive(Debug, Clone, PartialEq, Eq)]
212pub struct ClusterCfg {
213 pub join: Option<String>,
215 pub n_workers: usize,
217 pub exchange_bind: String,
220 pub control_bind: String,
222 pub flows: usize,
224 pub window: usize,
225}
226
227impl ClusterCfg {
228 pub fn parse(
232 join: &str,
233 n_workers: usize,
234 exchange_bind: &str,
235 control_bind: &str,
236 flows: usize,
237 window: usize,
238 ) -> Option<ClusterCfg> {
239 if join.is_empty() && n_workers <= 1 {
240 return None;
241 }
242 Some(ClusterCfg {
243 join: (!join.is_empty()).then(|| join.to_string()),
244 n_workers: n_workers.max(2),
245 exchange_bind: if exchange_bind.is_empty() {
246 "127.0.0.1:0".into()
247 } else {
248 exchange_bind.into()
249 },
250 control_bind: if control_bind.is_empty() {
251 "0.0.0.0:7100".into()
252 } else {
253 control_bind.into()
254 },
255 flows: flows.max(1),
256 window: window.max(1),
257 })
258 }
259}
260
261pub async fn connect(cfg: &ClusterCfg, base_fp: &str) -> std::io::Result<(Members, TcpListener, Coord)> {
267 let exchange = TcpListener::bind(&cfg.exchange_bind).await?;
268 let my_addr = exchange.local_addr()?.to_string();
269 let (members, coord) = match &cfg.join {
270 Some(leader) => {
271 eprintln!("fv-streams: joining the leader at {leader} (exchange on {my_addr})…");
272 let addr: SocketAddr = leader.parse().map_err(invalid)?;
273 let (m, leader_fp, ctl) = join(addr, &my_addr).await?;
274 if leader_fp != base_fp {
275 return Err(invalid(
276 "topology fingerprint mismatch with the leader — this worker's steps/inputs/vnodes differ; \
277 every worker must run the same pipeline.toml and STREAM_VNODES",
278 ));
279 }
280 eprintln!("fv-streams: joined — worker {} of {}", m.me(), m.len());
281 (m, Coord::Worker(ctl))
282 }
283 None => {
284 let control = TcpListener::bind(&cfg.control_bind).await?;
285 let addr = control.local_addr()?;
286 eprintln!("fv-streams: leader ready on {addr} (exchange {my_addr}).");
287 eprintln!("fv-streams: workers join with: fv-streams run <pipeline.toml> --join {addr}");
288 eprintln!("fv-streams: waiting for {} worker(s) to join…", cfg.n_workers - 1);
289 let (m, ctl) = lead(&control, &my_addr, cfg.n_workers, base_fp).await?;
290 eprintln!("fv-streams: all {} workers joined — starting.", m.len());
291 (m, Coord::Leader(ctl))
292 }
293 };
294 Ok((members, exchange, coord))
295}
296
297pub async fn join(leader_control: SocketAddr, my_exchange_addr: &str) -> std::io::Result<(Members, String, WorkerCtl)> {
301 join_as(leader_control, my_exchange_addr, fv_streams_exchange::WIRE_VERSION).await
302}
303
304async fn join_as(
306 leader_control: SocketAddr,
307 my_exchange_addr: &str,
308 wire: u8,
309) -> std::io::Result<(Members, String, WorkerCtl)> {
310 let mut s = TcpStream::connect(leader_control).await?;
311 write_msg(
312 &mut s,
313 &Register {
314 exchange_addr: my_exchange_addr.to_string(),
315 wire,
316 },
317 )
318 .await?;
319 let assign: Assignment = read_msg(&mut s).await?;
320 if assign.wire != wire {
321 return Err(wire_mismatch("the leader", assign.wire, wire));
323 }
324 let members = Members::new(assign.workers, my_exchange_addr).map_err(invalid)?;
325 Ok((members, assign.fingerprint, WorkerCtl { conn: s }))
326}
327
328pub async fn mesh(
333 members: &Members,
334 exchange: &TcpListener,
335 handle: Handle,
336 flows: usize,
337 window: usize,
338) -> std::io::Result<WorkerExchange> {
339 let me = members.me();
340 let higher: Vec<WorkerId> = members.peers().filter(|&p| p > me).collect();
341 let lower: Vec<WorkerId> = members.peers().filter(|&p| p < me).collect();
342
343 let accept_fut = Link::accept_mesh(exchange, &higher, flows, window);
344 let connect_fut = async {
345 let mut m: HashMap<WorkerId, Link> = HashMap::new();
346 for p in &lower {
347 let addr: SocketAddr = members.addr(*p).expect("peer address").parse().map_err(invalid)?;
348 m.insert(*p, Link::connect_as(addr, me, flows, window).await?);
349 }
350 Ok::<_, std::io::Error>(m)
351 };
352 let (accepted, connected) = tokio::join!(accept_fut, connect_fut);
353 let links = accepted?.into_iter().chain(connected?);
354
355 let mut senders = HashMap::new();
356 let mut receivers = Vec::new();
357 for (peer, link) in links {
358 let (s, r) = link.into_split();
359 senders.insert(peer, s);
360 receivers.push(r);
361 }
362 Ok(WorkerExchange {
363 handle,
364 senders,
365 receivers,
366 })
367}
368
369#[cfg(test)]
370mod tests {
371 use super::*;
372 use fv_streams_exchange::FrameKind;
373
374 fn rt() -> tokio::runtime::Runtime {
375 tokio::runtime::Builder::new_multi_thread()
376 .worker_threads(4)
377 .enable_all()
378 .build()
379 .unwrap()
380 }
381
382 #[test]
383 fn cluster_cfg_parse_is_solo_unless_joined_or_multi_worker() {
384 assert_eq!(ClusterCfg::parse("", 1, "", "", 4, 1 << 20), None);
386 assert_eq!(ClusterCfg::parse("", 0, "", "", 4, 1 << 20), None);
387 let w = ClusterCfg::parse("10.0.0.1:7100", 1, "", "", 4, 1 << 20).unwrap();
389 assert_eq!(w.join.as_deref(), Some("10.0.0.1:7100"));
390 assert_eq!(w.n_workers, 2, "a joiner implies at least two workers");
391 assert_eq!(w.exchange_bind, "127.0.0.1:0", "ephemeral exchange bind by default");
392 let l = ClusterCfg::parse("", 3, "1.2.3.4:9000", "0.0.0.0:7100", 8, 1 << 22).unwrap();
394 assert_eq!(l.join, None);
395 assert_eq!((l.n_workers, l.flows, l.window), (3, 8, 1 << 22));
396 assert_eq!(l.exchange_bind, "1.2.3.4:9000");
397 }
398
399 #[test]
400 fn a_worker_on_another_wire_version_is_refused_at_join_by_both_sides() {
401 let rt = rt();
405 rt.block_on(async {
406 let lx = TcpListener::bind("127.0.0.1:0").await.unwrap();
407 let leader_x = lx.local_addr().unwrap().to_string();
408 let control = TcpListener::bind("127.0.0.1:0").await.unwrap();
409 let caddr = control.local_addr().unwrap();
410 let wx = TcpListener::bind("127.0.0.1:0").await.unwrap();
411 let worker_x = wx.local_addr().unwrap().to_string();
412 let leader = tokio::spawn({
413 let leader_x = leader_x.clone();
414 async move { lead(&control, &leader_x, 2, "fp-1").await }
415 });
416 let foreign = fv_streams_exchange::WIRE_VERSION + 1;
417 let worker_err = join_as(caddr, &worker_x, foreign)
418 .await
419 .err()
420 .expect("the joiner refuses");
421 let leader_err = leader.await.unwrap().err().expect("the leader refuses");
422 for (who, e) in [("joiner", worker_err.to_string()), ("leader", leader_err.to_string())] {
423 assert!(
424 e.contains("wire format v") && e.contains("same fv-streams version"),
425 "{who}: {e}"
426 );
427 }
428 assert!(leader_err.to_string().contains(&format!("v{foreign}")), "{leader_err}");
429 });
430 }
431
432 #[test]
433 fn rendezvous_settles_the_same_membership_on_every_worker() {
434 let rt = rt();
435 rt.block_on(async {
436 let lx = TcpListener::bind("127.0.0.1:0").await.unwrap();
439 let leader_x = lx.local_addr().unwrap().to_string();
440 let control = TcpListener::bind("127.0.0.1:0").await.unwrap();
441 let caddr = control.local_addr().unwrap();
442 let wx = TcpListener::bind("127.0.0.1:0").await.unwrap();
443 let worker_x = wx.local_addr().unwrap().to_string();
444
445 let leader = tokio::spawn({
446 let leader_x = leader_x.clone();
447 async move { lead(&control, &leader_x, 2, "fp-1").await.unwrap() }
448 });
449 let (worker_members, fp, mut worker_ctl) = join(caddr, &worker_x).await.unwrap();
450 let (leader_members, mut leader_ctl) = leader.await.unwrap();
451
452 assert_eq!(leader_members.len(), 2);
454 assert_eq!(worker_members.len(), 2);
455 assert_eq!(fp, "fp-1");
456 assert_ne!(leader_members.me(), worker_members.me(), "distinct worker ids");
457 assert_eq!(leader_members.fingerprint_clause(), worker_members.fingerprint_clause());
458
459 leader_ctl.broadcast(ToWorker::Barrier(7)).await.unwrap();
462 assert_eq!(worker_ctl.recv().await.unwrap(), ToWorker::Barrier(7));
463 let contribution = fv_streams_state::Manifest {
464 v: 1,
465 epoch: 7,
466 fingerprint: "fp-1".into(),
467 sources: Default::default(),
468 objects: vec!["s0-op3".into()],
469 at_ms: 0,
470 files: Default::default(),
471 };
472 worker_ctl
473 .report(ToLeader::Acked {
474 epoch: 7,
475 contribution: contribution.clone(),
476 })
477 .await
478 .unwrap();
479 assert_eq!(
480 leader_ctl.collect().await.unwrap(),
481 vec![(worker_members.me(), ToLeader::Acked { epoch: 7, contribution })]
482 );
483 leader_ctl.broadcast(ToWorker::Commit(7)).await.unwrap();
484 assert_eq!(worker_ctl.recv().await.unwrap(), ToWorker::Commit(7));
485 leader_ctl.broadcast(ToWorker::Stop).await.unwrap();
486 assert_eq!(worker_ctl.recv().await.unwrap(), ToWorker::Stop);
487 worker_ctl.report(ToLeader::Finished).await.unwrap();
488 assert_eq!(
489 leader_ctl.collect().await.unwrap(),
490 vec![(worker_members.me(), ToLeader::Finished)]
491 );
492 });
493 }
494
495 #[test]
496 fn a_three_worker_mesh_links_every_pair_and_carries_frames() {
497 let rt = rt();
498 let handle = rt.handle().clone();
499 rt.block_on(async {
500 let listeners: Vec<TcpListener> = {
502 let mut v = Vec::new();
503 for _ in 0..3 {
504 v.push(TcpListener::bind("127.0.0.1:0").await.unwrap());
505 }
506 v
507 };
508 let addrs: Vec<String> = listeners.iter().map(|l| l.local_addr().unwrap().to_string()).collect();
509 let members: Vec<Members> = addrs.iter().map(|a| Members::new(addrs.clone(), a).unwrap()).collect();
510
511 let (flows, window) = (2usize, 1usize << 20);
513 let mut tasks = Vec::new();
514 for (m, l) in members.into_iter().zip(listeners) {
515 let h = handle.clone();
516 tasks.push(tokio::spawn(async move {
517 let ex = mesh(&m, &l, h, flows, window).await.unwrap();
518 (m.me(), ex)
519 }));
520 }
521 let mut exchanges: HashMap<WorkerId, WorkerExchange> = HashMap::new();
522 for t in tasks {
523 let (me, ex) = t.await.unwrap();
524 assert_eq!(ex.senders.len(), 2, "worker {me} has a sender to each peer");
526 assert_eq!(ex.receivers.len(), 2, "worker {me} has a receiver from each peer");
527 exchanges.insert(me, ex);
528 }
529
530 let ex0 = exchanges.get(&0).unwrap();
532 ex0.senders[&2]
533 .send(5, FrameKind::Data, bytes::Bytes::from_static(b"hi-2"))
534 .await
535 .unwrap();
536 let ex2 = exchanges.get_mut(&2).unwrap();
537 let mut saw = None;
538 for r in ex2.receivers.iter_mut() {
539 if let Ok(Some(f)) = tokio::time::timeout(std::time::Duration::from_secs(2), r.recv()).await {
540 saw = Some(f);
541 break;
542 }
543 }
544 let f = saw.expect("worker 2 received worker 0's frame");
545 assert_eq!((f.channel, &f.body[..]), (5, &b"hi-2"[..]));
546 });
547 }
548}