diff --git a/crates/tw-gateway/src/listen.rs b/crates/tw-gateway/src/listen.rs index b1fce9ba..deb7f9a3 100644 --- a/crates/tw-gateway/src/listen.rs +++ b/crates/tw-gateway/src/listen.rs @@ -7,6 +7,18 @@ //! 不接新连接:用户刚按下保存,所有客户端就连不上了。所以先绑新的,绑上了 //! 再让旧的退场;**绑不上就什么都不动**,旧的照常服务,并且说清为什么。 //! +//! # 新旧重叠时,旧的先让出来 +//! +//! 「新的先起来」有一种情况做不到:新旧两套地址在同一个端口上互相覆盖。 +//! Linux 不让 `127.0.0.1:8788` 和已经在听的 `0.0.0.0:8788` 并存(反过来也 +//! 一样),于是从 `bind: all` 换到一张网卡,新的回环那一个永远绑不上 —— +//! 而挡着它的正是我们自己要退场的那一个。 +//! +//! 这时退一步:先让要退场的那几个停止接新连接、**把端口真正放掉**(等 axum +//! 丢下监听器,见 [`Tracked`]),再绑新的。在跑的请求不受影响 —— 它们的连接 +//! 和监听器是两回事;中间只有几毫秒不接新连接,客户端的重试盖得住。新的还是 +//! 绑不上(端口真被别的程序占了),就把退场的那几个原样绑回去,照旧说清为什么。 +//! //! # 为什么要记下「此刻在听哪儿」 //! //! 界面上的地址以前是启动时记的一次。换了端口,网关已经在新端口上服务, @@ -80,13 +92,30 @@ pub fn unresolved(e: &tw_config::BindError) -> Msg { /// /// 控制面在写配置之前问这一句:绑不上的配置不该落盘,否则网关守着旧地址, /// 配置文件却说着新地址,两边从那一刻起各说各的。 +/// +/// **和我们自己在听的地址同端口时,「被占」答不了。**`0.0.0.0:8788` 在听, +/// `127.0.0.1:8788` 就一定报被占 —— 占着它的是我们,换监听时会先让出来。 +/// 这种就不在这里判,交给换监听那一步:真绑不上,它守住旧的并说清原因。 pub async fn check(state: &AppState, want: &[SocketAddr]) -> Result<(), Msg> { let held = state.listening().addrs; - let fresh: Vec = want.iter().filter(|a| !held.contains(a)).copied().collect(); - bind_all(&fresh) - .await - .map(drop) - .map_err(|(a, e)| bind_failure(a, &e)) + for a in want.iter().filter(|a| !held.contains(a)) { + match TcpListener::bind(*a).await { + Ok(l) => drop(l), + Err(e) if e.kind() == std::io::ErrorKind::AddrInUse && overlaps_ours(*a, &held) => {} + Err(e) => return Err(bind_failure(*a, &e)), + } + } + Ok(()) +} + +/// 这个地址会不会撞上我们自己在听的某一个:同端口,而且有一边是通配地址 +/// (或者干脆是同一个地址)。 +fn overlaps_ours(a: SocketAddr, held: &[SocketAddr]) -> bool { + held.iter().any(|h| { + h.port() == a.port() + && h.is_ipv4() == a.is_ipv4() + && (h.ip() == a.ip() || h.ip().is_unspecified() || a.ip().is_unspecified()) + }) } /// 一个地址上的服务。 @@ -97,6 +126,39 @@ struct Bound { actual: SocketAddr, /// 发一下(或者丢掉)就停止接新连接,已经在跑的请求自己跑完 stop: tokio::sync::oneshot::Sender<()>, + /// 端口真正放掉的那一刻:axum 丢下了监听器 + released: tokio::sync::oneshot::Receiver<()>, +} + +impl Bound { + /// 停止接新连接,等端口放掉。在跑的请求照常跑完,不用等它们。 + async fn release(self) { + let _ = self.stop.send(()); + // axum 收到信号后下一轮调度就丢下监听器;等不到也不能卡死换监听 + let _ = tokio::time::timeout(std::time::Duration::from_secs(5), self.released).await; + } +} + +/// 交给 axum 的监听器,**被丢掉时说一声**。 +/// +/// 停止信号发出去之后,axum 什么时候真的关掉监听的 socket,外面看不见; +/// 而「旧的先让出来」那一步要等的正是这一刻 —— 在它之前去绑,照样撞上。 +struct Tracked { + inner: TcpListener, + _released: tokio::sync::oneshot::Sender<()>, +} + +impl axum::serve::Listener for Tracked { + type Io = tokio::net::TcpStream; + type Addr = SocketAddr; + + fn accept(&mut self) -> impl std::future::Future + Send { + axum::serve::Listener::accept(&mut self.inner) + } + + fn local_addr(&self) -> std::io::Result { + self.inner.local_addr() + } } /// 绑这几个地址。**一个绑不上就全部作废** —— 半套监听器比一套都没换更难说清。 @@ -117,8 +179,17 @@ fn start(state: &AppState, want: SocketAddr, listener: TcpListener) -> std::io:: let actual = listener.local_addr()?; tracing::info!(%actual, "the gateway is listening"); let (stop, stopped) = tokio::sync::oneshot::channel::<()>(); + let (release, released) = tokio::sync::oneshot::channel::<()>(); let st = state.clone(); tokio::spawn(async move { + use axum::serve::ListenerExt; + let listener = Tracked { + inner: listener, + _released: release, + } + // 自己的监听器要经 `tap_io` 才拿得到对端地址(axum 只给 TcpListener + // 和 TapIo 实现了 ConnectInfo) + .tap_io(|_| {}); // `into_make_service_with_connect_info` 是拿到对端地址的唯一办法 —— // 少了它,来源白名单收到的永远是 unwrap 出来的默认值。 let r = axum::serve( @@ -134,7 +205,12 @@ fn start(state: &AppState, want: SocketAddr, listener: TcpListener) -> std::io:: Err(e) => tracing::error!(%actual, %e, "the listener ended"), } }); - Ok(Bound { want, actual, stop }) + Ok(Bound { + want, + actual, + stop, + released, + }) } fn snapshot(bound: &[Bound], error: Option) -> Listening { @@ -186,21 +262,57 @@ pub async fn serve_at(state: AppState, want: Vec, follow: bool) -> s // **还要的地址原样留着,只绑新增的** —— 从「仅本机」换到「局域网」时, // 回环那一个一直在听,正在上面跑的请求感觉不到任何变化 let fresh: Vec = want.iter().filter(|a| !have.contains(a)).copied().collect(); + let (keep, gone): (Vec, Vec) = + bound.into_iter().partition(|b| want.contains(&b.want)); let listeners = match bind_all(&fresh).await { - Ok(l) => l, + // 新的都绑上了,旧的才退场:停止接新连接,在跑的请求自己跑完 + Ok(l) => { + for b in gone { + tracing::info!(addr = %b.actual, "no longer wanted; draining"); + let _ = b.stop.send(()); + } + l + } + // 挡路的可能是我们自己要退场的那几个(见模块文档):先让它们放掉 + // 端口再试。**只在同端口时这么做** —— 别的端口被占,让出来也没用 + Err((a, e)) + if e.kind() == std::io::ErrorKind::AddrInUse + && gone.iter().any(|b| b.want.port() == a.port()) => + { + let back: Vec = gone.iter().map(|b| b.want).collect(); + for b in gone { + tracing::info!(addr = %b.actual, "no longer wanted and in the way; releasing it first"); + b.release().await; + } + match bind_all(&fresh).await { + Ok(l) => l, + Err((a, e)) => { + tracing::error!(%a, %e, "the new listen address cannot be bound; going back to the current one"); + bound = keep; + // 原样绑回去。这一步也失败的话(刚放掉就被别人抢了), + // 能听的照听,`addrs` 说的就是此刻真在听的 + for w in back { + match TcpListener::bind(w).await { + Ok(l) => bound.push(start(&state, w, l)?), + Err(e) => { + tracing::error!(%w, %e, "the previous listen address could not be taken back") + } + } + } + bound.sort_by_key(|b| have.iter().position(|w| *w == b.want)); + state.set_listening(snapshot(&bound, Some(bind_failure(a, &e)))); + continue; + } + } + } Err((a, e)) => { tracing::error!(%a, %e, "the new listen address cannot be bound; keeping the current one"); + bound = keep.into_iter().chain(gone).collect(); + bound.sort_by_key(|b| have.iter().position(|w| *w == b.want)); state.set_listening(snapshot(&bound, Some(bind_failure(a, &e)))); continue; } }; - // 新的都绑上了,旧的才退场:停止接新连接,在跑的请求自己跑完 - let (keep, gone): (Vec, Vec) = - bound.into_iter().partition(|b| want.contains(&b.want)); - for b in gone { - tracing::info!(addr = %b.actual, "no longer wanted; draining"); - let _ = b.stop.send(()); - } bound = keep; for (a, l) in listeners { bound.push(start(&state, a, l)?); diff --git a/crates/tw-gateway/tests/hotreload.rs b/crates/tw-gateway/tests/hotreload.rs index 64499fd5..1cca2897 100644 --- a/crates/tw-gateway/tests/hotreload.rs +++ b/crates/tw-gateway/tests/hotreload.rs @@ -666,3 +666,171 @@ async fn a_request_in_flight_survives_the_listener_being_rebuilt() { assert_eq!(r.status(), 200); assert!(r.text().await.unwrap().contains("slow")); } + +fn free_port() -> u16 { + std::net::TcpListener::bind("127.0.0.1:0") + .unwrap() + .local_addr() + .unwrap() + .port() +} + +fn bind(raw: &str) -> tw_config::Bind { + serde_yaml_ng::from_str(raw).unwrap() +} + +/// 换一次监听,等它换完(看状态里的地址,不靠睡多久)。 +async fn switch(state: &tw_gateway::AppState, next: Config) { + let want = next.listen.gateway.addrs().unwrap(); + state.reload(next).unwrap(); + for _ in 0..100 { + let now = state.listening(); + if now.addrs == want || now.error.is_some() { + return; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!("the listener never switched: {:?}", state.listening()); +} + +#[tokio::test] +async fn switching_between_all_interfaces_and_one_address_on_the_same_port_takes_effect() { + // **同一个端口上,通配地址和具体地址不能并存**(Linux):`bind: all` 换成 + // 一张网卡时,新的 127.0.0.1 撞上的是我们自己还没退场的 0.0.0.0。以前这一步 + // 报「端口被占」然后守着旧的,要重启才生效 + let (up, _) = counting_upstream("a").await; + let mut c = cfg(vec![provider("a", up)], vec![]); + let port = free_port(); + c.listen.gateway.port = port; + c.listen.gateway.bind = bind("all"); + let state = tw_gateway::AppState::new(c.clone()).unwrap(); + following(&state, &c).await; + let gw: SocketAddr = format!("127.0.0.1:{port}").parse().unwrap(); + assert!(ask(gw).await.contains("\"a\"")); + + for to in ["127.0.0.1", "all", "127.0.0.1"] { + let mut next = c.clone(); + next.listen.gateway.bind = bind(to); + switch(&state, next.clone()).await; + let now = state.listening(); + assert_eq!(now.error, None, "switching to {to}: {now:?}"); + assert_eq!( + now.addrs, + next.listen.gateway.addrs().unwrap(), + "switching to {to}" + ); + assert!( + ask(gw).await.contains("\"a\""), + "nothing answers after switching to {to}" + ); + } +} + +#[tokio::test] +async fn the_pre_save_check_does_not_call_our_own_port_taken() { + // 控制面在写配置之前问「绑不绑得上」。挡路的是我们自己在听的那一个时, + // 不该回答「被别的程序占了」—— 那样这份设置根本存不进去 + let (up, _) = counting_upstream("a").await; + let mut c = cfg(vec![provider("a", up)], vec![]); + c.listen.gateway.port = free_port(); + c.listen.gateway.bind = bind("all"); + let state = tw_gateway::AppState::new(c.clone()).unwrap(); + following(&state, &c).await; + let mut next = c.clone(); + next.listen.gateway.bind = bind("127.0.0.1"); + let want = next.listen.gateway.addrs().unwrap(); + assert_eq!(tw_gateway::listen::check(&state, &want).await, Ok(())); + + // 真被别人占着的,照样说 + let squatter = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); + let taken = squatter.local_addr().unwrap(); + let e = tw_gateway::listen::check(&state, &[taken]) + .await + .unwrap_err(); + assert_eq!(e.code, "gw.listen.port_taken"); +} + +#[tokio::test] +async fn a_request_in_flight_survives_the_old_listener_giving_way() { + // 旧的先让出端口时,**让的只是监听**:已经连上、跑到一半的请求照常跑完 + let slow = { + let app = Router::new().fallback(any(|| async { + tokio::time::sleep(Duration::from_millis(600)).await; + axum::response::Response::builder() + .header("content-type", "application/json") + .body(axum::body::Body::from(r#"{"by":"slow"}"#)) + .unwrap() + })); + let l = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let a = l.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(l, app).await.unwrap() }); + a + }; + let mut c = cfg(vec![provider("slow", slow)], vec![]); + let port = free_port(); + c.listen.gateway.port = port; + c.listen.gateway.bind = bind("all"); + let state = tw_gateway::AppState::new(c.clone()).unwrap(); + following(&state, &c).await; + + let inflight = tokio::spawn(async move { + reqwest::Client::new() + .post(format!("http://127.0.0.1:{port}/v1/messages")) + .header("x-api-key", "tw-k") + .timeout(Duration::from_secs(10)) + .body(r#"{"model":"m","messages":[]}"#) + .send() + .await + }); + tokio::time::sleep(Duration::from_millis(120)).await; + let mut next = c.clone(); + next.listen.gateway.bind = bind("127.0.0.1"); + switch(&state, next).await; + assert_eq!(state.listening().error, None); + assert!(!inflight.is_finished(), "测试的前提:旧请求还在跑"); + assert!( + tokio::net::TcpStream::connect(("127.0.0.1", port)) + .await + .is_ok(), + "旧请求还没跑完,新的监听就该能连了" + ); + let r = inflight + .await + .unwrap() + .expect("让出端口把跑到一半的请求掐了"); + assert_eq!(r.status(), 200); +} + +#[tokio::test] +async fn when_the_new_address_is_really_taken_the_old_one_is_taken_back() { + // 让出来之后新的还是绑不上(这回真是别的程序):旧的原样绑回去,并说清原因。 + // + // 场景:0.0.0.0:p 在听,换到 [::1]:p,而 [::1]:p 被别人占着。同端口,所以 + // 走「先让出来」那条路;让了也没用,于是回到 0.0.0.0:p + let port = free_port(); + let Ok(squatter) = std::net::TcpListener::bind(("::1", port)) else { + eprintln!("no IPv6 loopback here; skipping"); + return; + }; + let (up, _) = counting_upstream("a").await; + let mut c = cfg(vec![provider("a", up)], vec![]); + c.listen.gateway.port = port; + c.listen.gateway.bind = bind("all"); + let state = tw_gateway::AppState::new(c.clone()).unwrap(); + following(&state, &c).await; + let before = state.listening().addrs; + + let mut next = c.clone(); + next.listen.gateway.bind = bind("::1"); + switch(&state, next).await; + let now = state.listening(); + assert_eq!( + now.error.as_ref().map(|m| m.code.as_str()), + Some("gw.listen.port_taken"), + "{now:?}" + ); + assert_eq!(now.addrs, before, "旧的该原样回来"); + let gw: SocketAddr = format!("127.0.0.1:{port}").parse().unwrap(); + assert!(ask(gw).await.contains("\"a\""), "旧的地址不该停"); + drop(squatter); +}