1use std::collections::HashMap;
29use std::io::{self, Read, Write};
30use std::net::{Shutdown, SocketAddr, TcpListener, TcpStream};
31use std::os::fd::AsFd;
32use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
33use std::sync::{Arc, Mutex, MutexGuard};
34use std::thread::JoinHandle;
35use std::time::{Duration, Instant};
36
37use crate::error::{Error, Result};
38
39pub const MAX_CONNS_PER_ROUTE: usize = 4096;
41pub const CONNECT_TIMEOUT: Duration = Duration::from_secs(3);
43pub const BACKOFF_MIN: Duration = Duration::from_secs(2);
45pub const BACKOFF_MAX: Duration = Duration::from_secs(30);
47const ACCEPT_POLL: Duration = Duration::from_millis(100);
49const KEEPALIVE_IDLE: Duration = Duration::from_secs(60);
52const COPY_BUF: usize = 32 * 1024;
53const CONN_STACK: usize = 128 * 1024;
54
55#[derive(Debug, Clone, PartialEq, Eq)]
57pub struct RouteStatus {
58 pub key: String,
59 pub listen: SocketAddr,
61 pub backends: Vec<BackendStatus>,
62 pub draining: Vec<BackendStatus>,
64 pub accepted: u64,
66 pub failures: u64,
68 pub rejected: u64,
70}
71
72#[derive(Debug, Clone, PartialEq, Eq)]
73pub struct BackendStatus {
74 pub addr: SocketAddr,
75 pub active: usize,
77 pub down: bool,
79 pub connect_failures: u64,
81}
82
83#[derive(Clone, Default)]
87pub struct Balancer {
88 inner: Arc<Inner>,
89}
90
91#[derive(Default)]
92struct Inner {
93 routes: Mutex<HashMap<String, Route>>,
94 draining: Mutex<Vec<(String, Arc<Backend>)>>,
96}
97
98struct Route {
99 requested: SocketAddr,
101 shared: Arc<RouteShared>,
102 acceptor: Acceptor,
103}
104
105struct Acceptor {
106 bound: SocketAddr,
107 stop: Arc<AtomicBool>,
108 thread: Option<JoinHandle<()>>,
109}
110
111impl Acceptor {
112 fn stop(mut self) {
115 self.stop.store(true, Ordering::SeqCst);
116 if let Some(t) = self.thread.take() {
117 let _ = t.join();
118 }
119 }
120}
121
122struct RouteShared {
123 key: String,
124 backends: Mutex<Vec<Arc<Backend>>>,
125 rr: AtomicUsize,
126 conns: AtomicUsize,
127 accepted: AtomicU64,
128 failures: AtomicU64,
129 rejected: AtomicU64,
130}
131
132struct Backend {
133 addr: SocketAddr,
134 active: AtomicUsize,
135 failures: AtomicU64,
136 health: Mutex<Health>,
137}
138
139#[derive(Default)]
140struct Health {
141 down_until: Option<Instant>,
142 backoff: Duration,
143}
144
145impl Backend {
146 fn new(addr: SocketAddr) -> Arc<Backend> {
147 Arc::new(Backend {
148 addr,
149 active: AtomicUsize::new(0),
150 failures: AtomicU64::new(0),
151 health: Mutex::new(Health::default()),
152 })
153 }
154
155 fn is_down(&self, now: Instant) -> bool {
156 lock(&self.health).down_until.is_some_and(|t| t > now)
157 }
158
159 fn status(&self) -> BackendStatus {
160 BackendStatus {
161 addr: self.addr,
162 active: self.active.load(Ordering::SeqCst),
163 down: self.is_down(Instant::now()),
164 connect_failures: self.failures.load(Ordering::SeqCst),
165 }
166 }
167}
168
169fn lock<T>(m: &Mutex<T>) -> MutexGuard<'_, T> {
172 m.lock().unwrap_or_else(|e| e.into_inner())
173}
174
175impl Balancer {
176 pub fn new() -> Balancer {
177 Balancer::default()
178 }
179
180 pub fn set_route(
188 &self,
189 key: &str,
190 listen: SocketAddr,
191 backends: Vec<SocketAddr>,
192 ) -> Result<SocketAddr> {
193 let mut routes = lock(&self.inner.routes);
194 let existing = routes.get(key).filter(|r| r.requested == listen);
195 let Some(route) = existing else {
196 let listener = TcpListener::bind(listen).map_err(|e| {
198 let msg = format!("balance: route {key}: cannot listen on {listen}: {e}");
199 Error::Io(io::Error::new(e.kind(), msg))
200 })?;
201 let shared = match routes.get(key) {
202 Some(r) => r.shared.clone(),
203 None => Arc::new(RouteShared {
204 key: key.to_string(),
205 backends: Mutex::new(Vec::new()),
206 rr: AtomicUsize::new(0),
207 conns: AtomicUsize::new(0),
208 accepted: AtomicU64::new(0),
209 failures: AtomicU64::new(0),
210 rejected: AtomicU64::new(0),
211 }),
212 };
213 let acceptor = spawn_acceptor(listener, shared.clone())?;
214 let bound = acceptor.bound;
215 let old = routes.insert(
216 key.to_string(),
217 Route {
218 requested: listen,
219 shared: shared.clone(),
220 acceptor,
221 },
222 );
223 match &old {
224 Some(o) => eprintln!(
225 "balance: route {key}: moved from {} to {bound}",
226 o.acceptor.bound
227 ),
228 None => eprintln!("balance: route {key}: listening on {bound}"),
229 }
230 self.update_backends(&shared, backends);
231 drop(routes);
232 if let Some(o) = old {
233 o.acceptor.stop();
234 }
235 return Ok(bound);
236 };
237 let bound = route.acceptor.bound;
238 let shared = route.shared.clone();
239 drop(routes);
240 self.update_backends(&shared, backends);
241 Ok(bound)
242 }
243
244 fn update_backends(&self, shared: &RouteShared, wanted: Vec<SocketAddr>) {
245 let key = &shared.key;
246 let mut current = lock(&shared.backends);
247 let mut draining = lock(&self.inner.draining);
248 let before: Vec<SocketAddr> = current.iter().map(|b| b.addr).collect();
249 let mut next: Vec<Arc<Backend>> = Vec::with_capacity(wanted.len());
250 for addr in wanted {
251 if next.iter().any(|b| b.addr == addr) {
252 continue;
253 }
254 let b = if let Some(i) = current.iter().position(|b| b.addr == addr) {
257 current.swap_remove(i)
258 } else if let Some(i) = draining
259 .iter()
260 .position(|(k, b)| k == key && b.addr == addr)
261 {
262 draining.swap_remove(i).1
263 } else {
264 Backend::new(addr)
265 };
266 next.push(b);
267 }
268 let removed = std::mem::replace(&mut *current, next);
270 for b in removed {
271 eprintln!(
272 "balance: route {key}: backend {} removed ({} open, draining)",
273 b.addr,
274 b.active.load(Ordering::SeqCst)
275 );
276 draining.push((key.clone(), b));
277 }
278 draining.retain(|(_, b)| b.active.load(Ordering::SeqCst) > 0);
279 let after: Vec<SocketAddr> = current.iter().map(|b| b.addr).collect();
280 if before != after {
281 let list: Vec<String> = after.iter().map(|a| a.to_string()).collect();
282 eprintln!("balance: route {key}: backends [{}]", list.join(", "));
283 }
284 }
285
286 pub fn remove_route(&self, key: &str) {
289 let Some(route) = lock(&self.inner.routes).remove(key) else {
290 return;
291 };
292 let bound = route.acceptor.bound;
293 route.acceptor.stop();
294 let backends = std::mem::take(&mut *lock(&route.shared.backends));
295 lock(&self.inner.draining).extend(
296 backends
297 .into_iter()
298 .filter(|b| b.active.load(Ordering::SeqCst) > 0)
299 .map(|b| (key.to_string(), b)),
300 );
301 eprintln!("balance: route {key}: removed (was {bound})");
302 }
303
304 pub fn clear(&self) {
306 let keys: Vec<String> = lock(&self.inner.routes).keys().cloned().collect();
307 for k in keys {
308 self.remove_route(&k);
309 }
310 }
311
312 pub fn routes(&self) -> Vec<RouteStatus> {
313 let snapshot: Vec<(String, SocketAddr, Arc<RouteShared>)> = lock(&self.inner.routes)
316 .iter()
317 .map(|(k, r)| (k.clone(), r.acceptor.bound, r.shared.clone()))
318 .collect();
319 let mut out: Vec<RouteStatus> = snapshot
320 .into_iter()
321 .map(|(key, listen, shared)| {
322 let backends = lock(&shared.backends).iter().map(|b| b.status()).collect();
323 let draining = lock(&self.inner.draining)
324 .iter()
325 .filter(|(k, b)| *k == key && b.active.load(Ordering::SeqCst) > 0)
326 .map(|(_, b)| b.status())
327 .collect();
328 RouteStatus {
329 key,
330 listen,
331 backends,
332 draining,
333 accepted: shared.accepted.load(Ordering::SeqCst),
334 failures: shared.failures.load(Ordering::SeqCst),
335 rejected: shared.rejected.load(Ordering::SeqCst),
336 }
337 })
338 .collect();
339 out.sort_by(|a, b| a.key.cmp(&b.key));
340 out
341 }
342
343 pub fn active(&self, key: &str, backend: SocketAddr) -> usize {
346 let mut n = 0;
347 if let Some(r) = lock(&self.inner.routes).get(key) {
348 n += lock(&r.shared.backends)
349 .iter()
350 .filter(|b| b.addr == backend)
351 .map(|b| b.active.load(Ordering::SeqCst))
352 .sum::<usize>();
353 }
354 let mut draining = lock(&self.inner.draining);
355 draining.retain(|(_, b)| b.active.load(Ordering::SeqCst) > 0);
356 n + draining
357 .iter()
358 .filter(|(k, b)| k == key && b.addr == backend)
359 .map(|(_, b)| b.active.load(Ordering::SeqCst))
360 .sum::<usize>()
361 }
362
363 pub fn wait_drained(&self, key: &str, backend: SocketAddr, timeout: Duration) -> bool {
368 let deadline = Instant::now() + timeout;
369 loop {
370 if self.active(key, backend) == 0 {
371 return true;
372 }
373 let now = Instant::now();
374 if now >= deadline {
375 return false;
376 }
377 std::thread::sleep((deadline - now).min(Duration::from_millis(20)));
378 }
379 }
380}
381
382fn spawn_acceptor(listener: TcpListener, shared: Arc<RouteShared>) -> Result<Acceptor> {
383 let bound = listener.local_addr()?;
384 listener.set_nonblocking(true)?;
387 let stop = Arc::new(AtomicBool::new(false));
388 let thread = std::thread::Builder::new()
389 .name(format!("isb-lb-{}", shared.key))
390 .spawn({
391 let stop = stop.clone();
392 move || accept_loop(listener, shared, stop)
393 })?;
394 Ok(Acceptor {
395 bound,
396 stop,
397 thread: Some(thread),
398 })
399}
400
401fn accept_loop(listener: TcpListener, shared: Arc<RouteShared>, stop: Arc<AtomicBool>) {
402 let timeout = rustix::event::Timespec {
403 tv_sec: 0,
404 tv_nsec: ACCEPT_POLL.as_nanos() as _,
405 };
406 let mut pause = Duration::ZERO;
407 let mut last_log: Option<(io::ErrorKind, Option<i32>, Instant)> = None;
408 while !stop.load(Ordering::SeqCst) {
409 let mut fds = [rustix::event::PollFd::new(
410 &listener,
411 rustix::event::PollFlags::IN,
412 )];
413 match rustix::event::poll(&mut fds, Some(&timeout)) {
414 Ok(0) | Err(rustix::io::Errno::INTR) => continue,
415 Ok(_) => {}
416 Err(e) => {
417 eprintln!("balance: route {}: poll: {e}", shared.key);
419 std::thread::sleep(ACCEPT_POLL);
420 continue;
421 }
422 }
423 match listener.accept() {
424 Ok((client, _)) => {
425 pause = Duration::ZERO;
426 handle(client, &shared);
427 }
428 Err(e) if e.kind() == io::ErrorKind::WouldBlock => {}
429 Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
430 Err(e) => {
431 let now = Instant::now();
435 let same = last_log.as_ref().is_some_and(|(k, raw, t)| {
436 *k == e.kind()
437 && *raw == e.raw_os_error()
438 && now.duration_since(*t) < Duration::from_secs(60)
439 });
440 if !same {
441 eprintln!("balance: route {}: accept: {e} (retrying)", shared.key);
442 last_log = Some((e.kind(), e.raw_os_error(), now));
443 }
444 pause = (pause * 2).clamp(Duration::from_millis(5), Duration::from_secs(1));
445 std::thread::sleep(pause);
446 }
447 }
448 }
449}
450
451struct ConnSlot(Arc<RouteShared>);
453
454impl Drop for ConnSlot {
455 fn drop(&mut self) {
456 self.0.conns.fetch_sub(1, Ordering::SeqCst);
457 }
458}
459
460struct BackendSlot(Arc<Backend>);
462
463impl Drop for BackendSlot {
464 fn drop(&mut self) {
465 self.0.active.fetch_sub(1, Ordering::SeqCst);
466 }
467}
468
469fn handle(client: TcpStream, shared: &Arc<RouteShared>) {
470 shared.accepted.fetch_add(1, Ordering::SeqCst);
471 if shared.conns.fetch_add(1, Ordering::SeqCst) >= MAX_CONNS_PER_ROUTE {
472 shared.conns.fetch_sub(1, Ordering::SeqCst);
473 shared.rejected.fetch_add(1, Ordering::SeqCst);
474 return;
475 }
476 let slot = ConnSlot(shared.clone());
477 let spawned = std::thread::Builder::new()
479 .name("isb-lb-conn".into())
480 .stack_size(CONN_STACK)
481 .spawn(move || serve(client, slot));
482 if let Err(e) = spawned {
483 shared.rejected.fetch_add(1, Ordering::SeqCst);
486 eprintln!("balance: route {}: cannot spawn: {e}", shared.key);
487 }
488}
489
490fn serve(client: TcpStream, slot: ConnSlot) {
491 let shared = slot.0.clone();
492 let Some((upstream, backend)) = connect_any(&shared) else {
493 shared.failures.fetch_add(1, Ordering::SeqCst);
494 return;
495 };
496 let _ = client.set_nonblocking(false);
498 for s in [&client, &upstream] {
499 let _ = s.set_nodelay(true);
500 let _ = rustix::net::sockopt::set_socket_keepalive(s.as_fd(), true);
501 let _ = rustix::net::sockopt::set_tcp_keepidle(s.as_fd(), KEEPALIVE_IDLE);
502 }
503 let (Ok(client2), Ok(upstream2)) = (client.try_clone(), upstream.try_clone()) else {
504 return;
505 };
506 let other = std::thread::Builder::new()
507 .name("isb-lb-conn".into())
508 .stack_size(CONN_STACK)
509 .spawn(move || pipe(upstream2, client2));
510 let Ok(other) = other else {
511 return;
512 };
513 pipe(client, upstream);
514 let _ = other.join();
515 drop(backend);
516 drop(slot);
517}
518
519fn pipe(mut from: TcpStream, mut to: TcpStream) {
522 let mut buf = vec![0u8; COPY_BUF];
523 loop {
524 match from.read(&mut buf) {
525 Ok(0) => {
526 let _ = to.shutdown(Shutdown::Write);
527 return;
528 }
529 Ok(n) => {
530 if to.write_all(&buf[..n]).is_err() {
531 break;
532 }
533 }
534 Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
535 Err(_) => break,
536 }
537 }
538 let _ = from.shutdown(Shutdown::Both);
539 let _ = to.shutdown(Shutdown::Both);
540}
541
542fn connect_any(shared: &RouteShared) -> Option<(TcpStream, BackendSlot)> {
546 let mut tried: Vec<SocketAddr> = Vec::new();
547 loop {
548 let backend = pick(shared, &tried)?;
549 tried.push(backend.0.addr);
550 match TcpStream::connect_timeout(&backend.0.addr, CONNECT_TIMEOUT) {
551 Ok(s) => {
552 mark_up(&shared.key, &backend.0);
553 return Some((s, backend));
554 }
555 Err(e) => mark_down(&shared.key, &backend.0, &e),
556 }
557 }
558}
559
560fn pick(shared: &RouteShared, tried: &[SocketAddr]) -> Option<BackendSlot> {
564 let backends = lock(&shared.backends);
565 let n = backends.len();
566 if n == 0 {
567 return None;
568 }
569 let now = Instant::now();
570 let start = shared.rr.fetch_add(1, Ordering::SeqCst) % n;
571 let order = || (0..n).map(|i| &backends[(start + i) % n]);
572 let untried = |b: &&Arc<Backend>| !tried.contains(&b.addr);
573 let best = order()
574 .filter(untried)
575 .filter(|b| !b.is_down(now))
576 .min_by_key(|b| b.active.load(Ordering::SeqCst))
577 .or_else(|| {
578 order()
579 .filter(untried)
580 .min_by_key(|b| b.active.load(Ordering::SeqCst))
581 })?;
582 best.active.fetch_add(1, Ordering::SeqCst);
583 Some(BackendSlot(best.clone()))
584}
585
586fn mark_down(key: &str, b: &Backend, err: &io::Error) {
587 b.failures.fetch_add(1, Ordering::SeqCst);
588 let mut h = lock(&b.health);
589 let now = Instant::now();
590 let was_up = h.down_until.is_none_or(|t| t <= now);
591 h.backoff = if h.backoff.is_zero() {
592 BACKOFF_MIN
593 } else {
594 (h.backoff * 2).min(BACKOFF_MAX)
595 };
596 h.down_until = Some(now + h.backoff);
597 if was_up {
598 eprintln!(
599 "balance: route {key}: backend {} down: {err} (retry in {:?})",
600 b.addr, h.backoff
601 );
602 }
603}
604
605fn mark_up(key: &str, b: &Backend) {
606 let mut h = lock(&b.health);
607 if !h.backoff.is_zero() {
608 *h = Health::default();
609 eprintln!("balance: route {key}: backend {} up", b.addr);
610 }
611}
612
613#[cfg(test)]
614mod tests {
615 use super::*;
616
617 fn any() -> SocketAddr {
618 "127.0.0.1:0".parse().unwrap()
619 }
620
621 fn echo(id: &'static str) -> SocketAddr {
624 let l = TcpListener::bind(any()).unwrap();
625 let addr = l.local_addr().unwrap();
626 std::thread::spawn(move || {
627 for s in l.incoming() {
628 let Ok(mut s) = s else { continue };
629 std::thread::spawn(move || {
630 s.write_all(format!("{id}\n").as_bytes()).unwrap();
631 let mut r = s.try_clone().unwrap();
632 let _ = io::copy(&mut r, &mut s);
633 let _ = s.shutdown(Shutdown::Write);
634 });
635 }
636 });
637 addr
638 }
639
640 fn dead() -> SocketAddr {
649 let ip = if cfg!(target_os = "macos") {
650 [127, 0, 0, 1]
651 } else {
652 [127, 0, 0, 3]
653 };
654 SocketAddr::from((ip, free().port()))
655 }
656
657 fn free() -> SocketAddr {
660 TcpListener::bind(any()).unwrap().local_addr().unwrap()
661 }
662
663 fn read_line(s: &mut TcpStream) -> io::Result<String> {
664 let mut out = Vec::new();
665 let mut b = [0u8; 1];
666 loop {
667 match s.read(&mut b)? {
668 0 => break,
669 _ if b[0] == b'\n' => break,
670 _ => out.push(b[0]),
671 }
672 }
673 Ok(String::from_utf8(out).unwrap())
674 }
675
676 fn open(addr: SocketAddr) -> (TcpStream, String) {
678 let mut s = TcpStream::connect(addr).unwrap();
679 s.set_read_timeout(Some(Duration::from_secs(10))).unwrap();
680 let tag = read_line(&mut s).unwrap();
681 (s, tag)
682 }
683
684 fn roundtrip(s: &mut TcpStream, msg: &str) -> String {
685 s.write_all(format!("{msg}\n").as_bytes()).unwrap();
686 read_line(s).unwrap()
687 }
688
689 fn status(lb: &Balancer, key: &str) -> RouteStatus {
690 lb.routes().into_iter().find(|r| r.key == key).unwrap()
691 }
692
693 #[test]
694 fn least_connections_spreads_and_refills() {
695 let (a, b, c) = (echo("a"), echo("b"), echo("c"));
696 let lb = Balancer::new();
697 let at = lb.set_route("web", any(), vec![a, b, c]).unwrap();
698 let mut conns: Vec<(TcpStream, String)> = (0..6).map(|_| open(at)).collect();
699 for id in ["a", "b", "c"] {
700 assert_eq!(conns.iter().filter(|(_, t)| t == id).count(), 2, "{id}");
701 }
702 let st = status(&lb, "web");
703 assert!(st.backends.iter().all(|b| b.active == 2), "{st:?}");
704 assert_eq!(st.accepted, 6);
705
706 conns.retain(|(_, t)| t != "a");
708 assert!(lb.wait_drained("web", a, Duration::from_secs(5)));
709 let next: Vec<_> = (0..2).map(|_| open(at)).collect();
710 assert!(next.iter().all(|(_, t)| t == "a"), "least-conn ignored");
711 lb.clear();
712 }
713
714 #[test]
715 fn dead_backend_is_skipped_and_retried_elsewhere() {
716 let (gone, live) = (dead(), echo("live"));
717 let lb = Balancer::new();
718 let at = lb.set_route("web", any(), vec![gone, live]).unwrap();
719 for _ in 0..5 {
720 let (mut s, tag) = open(at);
721 assert_eq!(tag, "live");
722 assert_eq!(roundtrip(&mut s, "hi"), "hi");
723 }
724 let st = status(&lb, "web");
725 assert_eq!(st.failures, 0);
726 let g = st.backends.iter().find(|b| b.addr == gone).unwrap();
727 assert!(g.down, "{g:?}");
728 assert_eq!(g.connect_failures, 1);
730 lb.clear();
731 }
732
733 #[test]
734 fn no_reachable_backend_closes_client() {
735 let lb = Balancer::new();
736 let empty = lb.set_route("empty", any(), vec![]).unwrap();
737 let only_dead = lb.set_route("dead", any(), vec![dead()]).unwrap();
738 for at in [empty, only_dead] {
739 let (_, tag) = open(at);
740 assert_eq!(tag, "", "expected EOF");
741 }
742 assert_eq!(status(&lb, "empty").failures, 1);
743 assert_eq!(status(&lb, "dead").failures, 1);
744 lb.clear();
745 }
746
747 #[test]
748 fn removed_backend_drains() {
749 let (old, new) = (echo("old"), echo("new"));
750 let lb = Balancer::new();
751 let at = lb.set_route("web", any(), vec![old]).unwrap();
752 let (mut s, tag) = open(at);
753 assert_eq!(tag, "old");
754
755 lb.set_route("web", any(), vec![new]).unwrap();
756 assert_eq!(roundtrip(&mut s, "still here"), "still here");
757 assert_eq!(open(at).1, "new");
758 let st = status(&lb, "web");
759 assert_eq!(st.draining.len(), 1);
760 assert_eq!(st.draining[0].addr, old);
761 assert!(!lb.wait_drained("web", old, Duration::from_millis(200)));
762
763 drop(s);
764 assert!(lb.wait_drained("web", old, Duration::from_secs(5)));
765 assert!(status(&lb, "web").draining.is_empty());
766 lb.clear();
767 }
768
769 #[test]
770 fn remove_route_stops_listening_and_keeps_connections() {
771 let a = echo("a");
772 let lb = Balancer::new();
773 let own = if cfg!(target_os = "macos") {
778 "127.0.0.1:0"
779 } else {
780 "127.0.0.2:0"
781 };
782 let at = lb.set_route("web", own.parse().unwrap(), vec![a]).unwrap();
783 let (mut s, _) = open(at);
784 lb.remove_route("web");
785 let err = TcpStream::connect(at).unwrap_err();
786 assert_eq!(err.kind(), io::ErrorKind::ConnectionRefused);
787 assert!(lb.routes().is_empty());
788 assert_eq!(roundtrip(&mut s, "after"), "after");
789 assert_eq!(lb.active("web", a), 1);
790 drop(s);
791 assert!(lb.wait_drained("web", a, Duration::from_secs(5)));
792 }
793
794 #[test]
795 fn updating_backends_does_not_rebind() {
796 let (a, b) = (echo("a"), echo("b"));
797 let lb = Balancer::new();
798 let at = lb.set_route("web", any(), vec![a]).unwrap();
799 let (mut s, _) = open(at);
800 let again = lb.set_route("web", any(), vec![a, b]).unwrap();
801 assert_eq!(at, again);
802 assert_eq!(status(&lb, "web").listen, at);
803 assert_eq!(roundtrip(&mut s, "x"), "x");
804 assert_eq!(open(at).1, "b");
805 lb.clear();
806 }
807
808 #[test]
809 fn half_close_is_forwarded() {
810 let lb = Balancer::new();
811 let at = lb.set_route("web", any(), vec![echo("a")]).unwrap();
812 let mut s = TcpStream::connect(at).unwrap();
813 s.set_read_timeout(Some(Duration::from_secs(10))).unwrap();
814 let payload = "x".repeat(200_000);
815 let mut r = s.try_clone().unwrap();
818 let reader = std::thread::spawn(move || {
819 let mut got = String::new();
820 r.read_to_string(&mut got).map(|_| got)
821 });
822 s.write_all(payload.as_bytes()).unwrap();
823 s.shutdown(Shutdown::Write).unwrap();
824 assert_eq!(reader.join().unwrap().unwrap(), format!("a\n{payload}"));
825 lb.clear();
826 }
827
828 #[test]
829 fn bind_conflict_is_an_error() {
830 let taken = TcpListener::bind(any()).unwrap();
831 let lb = Balancer::new();
832 let err = lb
833 .set_route("web", taken.local_addr().unwrap(), vec![echo("a")])
834 .unwrap_err();
835 assert!(err.to_string().contains("cannot listen"), "{err}");
836 assert!(matches!(&err, Error::Io(e) if e.kind() == io::ErrorKind::AddrInUse));
837 assert!(lb.routes().is_empty());
838 }
839
840 #[test]
841 fn changing_listen_moves_the_route() {
842 let a = echo("a");
843 let lb = Balancer::new();
844 let first = lb.set_route("web", any(), vec![a]).unwrap();
845 let (mut s, _) = open(first);
846
847 let taken = TcpListener::bind(any()).unwrap();
849 let conflict = lb.set_route("web", taken.local_addr().unwrap(), vec![a]);
850 assert!(conflict.is_err());
851 assert_eq!(status(&lb, "web").listen, first);
852 assert_eq!(open(first).1, "a");
853
854 let (target, moved) = (0..20)
855 .find_map(|_| {
856 let t = free();
857 lb.set_route("web", t, vec![a]).ok().map(|m| (t, m))
858 })
859 .expect("no free port to move to");
860 assert_eq!(moved, target);
861 assert_eq!(status(&lb, "web").listen, target);
862 assert_eq!(open(moved).1, "a");
863 match TcpStream::connect(first) {
867 Err(e) => assert_eq!(e.kind(), io::ErrorKind::ConnectionRefused),
868 Ok(_) => eprintln!("old port {first} was reused by another listener"),
869 }
870 assert_eq!(roundtrip(&mut s, "y"), "y");
872 assert_eq!(status(&lb, "web").accepted, 3);
873 lb.clear();
874 }
875}