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
113 changes: 102 additions & 11 deletions crates/gateway/tests/claims_reliability.rs
Original file line number Diff line number Diff line change
Expand Up @@ -980,14 +980,21 @@ async fn a_request_specific_403_neither_walks_nor_cools_the_pool_key() {
/// claim: REL-16, BIL-21
#[tokio::test]
async fn sigterm_drains_an_in_flight_request_and_bills_it() {
let up = ReplyUpstream::start(|_, _| {
Reply::Delayed(Duration::from_millis(1500), Box::new(Reply::ok()))
})
.await;
// Held, not delayed: the answer is released only once the gateway has acted on SIGTERM, so the
// request is in flight through the drain on any runner. A fixed sleep before the signal (400 ms
// against a 1.5 s reply) failed every run on a starved host: the request had not reached the
// provider yet when the test looked.
let release = Arc::new(tokio::sync::Notify::new());
let held = Arc::clone(&release);
let up =
ReplyUpstream::start(move |_, _| Reply::Held(Arc::clone(&held), Box::new(Reply::ok())))
.await;
let (pubkey, sk) = test_keypair(161);
let gw = Gateway::builder(unused_nats_port(), &up.authority(), &b64(&pubkey))
.providers(&["openai"])
.config_line("shutdown_grace_period_secs = 10")
// Pingora's own "SIGTERM received" line is how the test knows the signal has landed.
.env("AI_LOG", "warn,ai.usage=info,pingora_core::server=info")
.start()
.await;
let key = billing_vkey(&sk, 1601);
Expand All @@ -1006,13 +1013,18 @@ async fn sigterm_drains_an_in_flight_request_and_bills_it() {
resp.text().await.unwrap_or_default(),
)
});
tokio::time::sleep(Duration::from_millis(400)).await;
let deadline = std::time::Instant::now() + CONDITION_BUDGET;
while up.hits() == 0 && std::time::Instant::now() < deadline {
tokio::time::sleep(Duration::from_millis(25)).await;
}
assert_eq!(
up.hits(),
1,
"the request reached the provider before SIGTERM"
);
gw.sigterm();
gw.wait_for_log_line(&["SIGTERM received"]).await;
release.notify_one();
let (status, text) = inflight.await.unwrap();
assert_eq!(status, 200, "{text}");
assert!(text.contains("chatcmpl"), "{text}");
Expand Down Expand Up @@ -1419,8 +1431,77 @@ async fn concurrent_requests_share_an_upstream_h2_connection() {
);
}

/// A server-side socket that closes the way a real HTTP/2 server's does: once the connection is
/// finished, it keeps reading (and discarding) what the client still sends for a while before the
/// descriptor is closed.
///
/// Closing a socket with unread bytes in its receive queue makes the kernel answer with RST instead
/// of FIN, and an RST can destroy the GOAWAY still on its way to the client: the client's next write
/// fails with a broken pipe before it reads the frame that says its stream was never processed. A
/// gateway that never saw the GOAWAY cannot tell an unprocessed stream from a processed one, so it
/// rightly does not resend the POST, and answers 502 "upstream failed after receiving the request".
/// hyper's `graceful_shutdown` drops the socket the moment its last stream ends, so
/// [`goaway_upstream`] hit exactly that under CPU contention (2 of 288 runs, up to 7 of 16 requests;
/// once in a mutation run's unmutated baseline). Real servers linger for this reason: Go's
/// `x/net/http2` waits `goAwayTimeout` (1 s) after its final GOAWAY before closing, and nginx does a
/// lingering close.
struct Lingering(Option<tokio::net::TcpStream>);

impl Lingering {
fn io(&mut self) -> std::pin::Pin<&mut tokio::net::TcpStream> {
std::pin::Pin::new(self.0.as_mut().expect("present until drop"))
}
}

impl tokio::io::AsyncRead for Lingering {
fn poll_read(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
self.get_mut().io().poll_read(cx, buf)
}
}

impl tokio::io::AsyncWrite for Lingering {
fn poll_write(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &[u8],
) -> std::task::Poll<std::io::Result<usize>> {
self.get_mut().io().poll_write(cx, buf)
}
fn poll_flush(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<std::io::Result<()>> {
self.get_mut().io().poll_flush(cx)
}
fn poll_shutdown(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<std::io::Result<()>> {
self.get_mut().io().poll_shutdown(cx)
}
}

impl Drop for Lingering {
fn drop(&mut self) {
if let (Some(mut sock), Ok(rt)) = (self.0.take(), tokio::runtime::Handle::try_current()) {
rt.spawn(async move {
let mut sink = [0u8; 4096];
let _ = tokio::time::timeout(Duration::from_secs(1), async {
while matches!(sock.read(&mut sink).await, Ok(n) if n > 0) {}
})
.await;
});
}
}
}

/// A TLS H2 upstream that sends GOAWAY (graceful) on every connection right after its first
/// request — what a provider's load balancer does when it drains a node.
/// request — what a provider's load balancer does when it drains a node. Its sockets close like a
/// real server's ([`Lingering`]).
async fn goaway_upstream() -> (u16, Arc<AtomicUsize>, tokio::task::JoinHandle<()>) {
use http_body_util::{BodyExt, Full};
use hyper::service::service_fn;
Expand All @@ -1442,7 +1523,7 @@ async fn goaway_upstream() -> (u16, Arc<AtomicUsize>, tokio::task::JoinHandle<()
while let Ok((s, _)) = listener.accept().await {
let (acceptor, counter) = (acceptor.clone(), counter.clone());
tokio::spawn(async move {
let Ok(tls) = acceptor.accept(s).await else {
let Ok(tls) = acceptor.accept(Lingering(Some(s))).await else {
return;
};
let first = Arc::new(tokio::sync::Notify::new());
Expand Down Expand Up @@ -1873,18 +1954,28 @@ async fn upstream_goaway_is_handled() {
.body(CHAT)
.send()
.await
.map(|r| r.status().as_u16())
.unwrap_or(0)
}));
}
let mut statuses = Vec::new();
let mut failures = Vec::new();
for t in tasks {
statuses.push(t.await.unwrap());
match t.await.unwrap() {
Ok(r) if r.status() == 200 => statuses.push(200),
Ok(r) => {
statuses.push(r.status().as_u16());
failures.push(r.text().await.unwrap_or_default());
}
Err(e) => {
statuses.push(0);
failures.push(e.to_string());
}
}
}
task.abort();
assert!(
statuses.iter().all(|s| *s == 200),
"concurrent requests across GOAWAYs: {statuses:?}"
"concurrent requests across GOAWAYs: {statuses:?}\nfailures: {failures:#?}\nlog:\n{}",
log_tail(&gw.log())
);
assert!(served.load(Ordering::SeqCst) >= 26);
}
Expand Down
Loading
Loading