Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
140 changes: 126 additions & 14 deletions crates/tw-gateway/src/listen.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,18 @@
//! 不接新连接:用户刚按下保存,所有客户端就连不上了。所以先绑新的,绑上了
//! 再让旧的退场;**绑不上就什么都不动**,旧的照常服务,并且说清为什么。
//!
//! # 新旧重叠时,旧的先让出来
//!
//! 「新的先起来」有一种情况做不到:新旧两套地址在同一个端口上互相覆盖。
//! Linux 不让 `127.0.0.1:8788` 和已经在听的 `0.0.0.0:8788` 并存(反过来也
//! 一样),于是从 `bind: all` 换到一张网卡,新的回环那一个永远绑不上 ——
//! 而挡着它的正是我们自己要退场的那一个。
//!
//! 这时退一步:先让要退场的那几个停止接新连接、**把端口真正放掉**(等 axum
//! 丢下监听器,见 [`Tracked`]),再绑新的。在跑的请求不受影响 —— 它们的连接
//! 和监听器是两回事;中间只有几毫秒不接新连接,客户端的重试盖得住。新的还是
//! 绑不上(端口真被别的程序占了),就把退场的那几个原样绑回去,照旧说清为什么。
//!
//! # 为什么要记下「此刻在听哪儿」
//!
//! 界面上的地址以前是启动时记的一次。换了端口,网关已经在新端口上服务,
Expand Down Expand Up @@ -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<SocketAddr> = 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())
})
}

/// 一个地址上的服务。
Expand All @@ -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<Output = (Self::Io, Self::Addr)> + Send {
axum::serve::Listener::accept(&mut self.inner)
}

fn local_addr(&self) -> std::io::Result<SocketAddr> {
self.inner.local_addr()
}
}

/// 绑这几个地址。**一个绑不上就全部作废** —— 半套监听器比一套都没换更难说清。
Expand All @@ -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(
Expand All @@ -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<Msg>) -> Listening {
Expand Down Expand Up @@ -186,21 +262,57 @@ pub async fn serve_at(state: AppState, want: Vec<SocketAddr>, follow: bool) -> s
// **还要的地址原样留着,只绑新增的** —— 从「仅本机」换到「局域网」时,
// 回环那一个一直在听,正在上面跑的请求感觉不到任何变化
let fresh: Vec<SocketAddr> = want.iter().filter(|a| !have.contains(a)).copied().collect();
let (keep, gone): (Vec<Bound>, Vec<Bound>) =
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<SocketAddr> = 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<Bound>, Vec<Bound>) =
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)?);
Expand Down
168 changes: 168 additions & 0 deletions crates/tw-gateway/tests/hotreload.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}
Loading