diff --git a/tests/common/git_server.rs b/tests/common/git_server.rs index edc986e..7ca5a4d 100644 --- a/tests/common/git_server.rs +++ b/tests/common/git_server.rs @@ -150,6 +150,7 @@ impl SimpleGitServer { let repo_path = Arc::new(bare_repo_path); let handle = tokio::spawn(async move { + let mut connections = tokio::task::JoinSet::new(); println!("[SmartGitServer] Server loop started on port {}", port); eprintln!("[SmartGitServer] Server loop started on port {}", port); loop { @@ -161,7 +162,7 @@ impl SimpleGitServer { let repo_path = Arc::clone(&repo_path); let io = TokioIo::new(stream); - tokio::spawn(async move { + connections.spawn(async move { let service = service_fn(move |req| { let repo_path = Arc::clone(&repo_path); async move { handle_request(req, &repo_path).await } @@ -183,18 +184,20 @@ impl SimpleGitServer { } } } + _ = connections.join_next(), if !connections.is_empty() => {}, _ = &mut shutdown_rx => { // Shutdown signal received break; } } } + connections.shutdown().await; }); let url = format!("http://127.0.0.1:{}", port); // 7. Wait for server to be ready - wait_for_server_ready(port).await; + // The reserved listener is already bound; the runtime drives accepts. Self { shutdown_tx: Some(shutdown_tx), @@ -236,8 +239,10 @@ impl Drop for SimpleGitServer { if let Some(tx) = self.shutdown_tx.take() { let _ = tx.send(()); } - // Note: We can't await the handle in drop, but the temp_dir cleanup - // will happen automatically when _temp_dir is dropped + // Dropping the owner cancels every connection through its JoinSet. + if let Some(handle) = self.handle.take() { + handle.abort(); + } } } @@ -304,31 +309,6 @@ fn guess_content_type(path: &Path) -> &'static str { } } -/// Wait for the server to be ready to accept connections. -async fn wait_for_server_ready(port: u16) { - let max_attempts = 50; // 5 seconds total - let delay = std::time::Duration::from_millis(100); - - for attempt in 0..max_attempts { - match tokio::net::TcpStream::connect(format!("127.0.0.1:{}", port)).await { - Ok(_) => { - // Connection successful, server is ready - tokio::time::sleep(std::time::Duration::from_millis(50)).await; - return; - } - Err(_) => { - if attempt == max_attempts - 1 { - panic!( - "SimpleGitServer failed to start after {} attempts", - max_attempts - ); - } - tokio::time::sleep(delay).await; - } - } - } -} - #[cfg(test)] mod tests { use super::*; @@ -597,6 +577,7 @@ impl SmartGitServer { let repo_path = Arc::new(bare_repo_path); let handle = tokio::spawn(async move { + let mut connections = tokio::task::JoinSet::new(); loop { tokio::select! { accept_result = listener.accept() => { @@ -605,7 +586,7 @@ impl SmartGitServer { let repo_path = Arc::clone(&repo_path); let io = TokioIo::new(stream); - tokio::spawn(async move { + connections.spawn(async move { let service = service_fn(move |req| { let repo_path = Arc::clone(&repo_path); async move { handle_smart_request(req, &repo_path).await } @@ -627,18 +608,20 @@ impl SmartGitServer { } } } + _ = connections.join_next(), if !connections.is_empty() => {}, _ = &mut shutdown_rx => { // Shutdown signal received break; } } } + connections.shutdown().await; }); let url = format!("http://127.0.0.1:{}", port); // 6. Wait for server to be ready - wait_for_server_ready(port).await; + // The reserved listener is already bound; the runtime drives accepts. Self { shutdown_tx: Some(shutdown_tx), @@ -680,8 +663,10 @@ impl Drop for SmartGitServer { if let Some(tx) = self.shutdown_tx.take() { let _ = tx.send(()); } - // Note: We can't await the handle in drop, but the temp_dir cleanup - // will happen automatically when _temp_dir is dropped + // Dropping the owner cancels every connection through its JoinSet. + if let Some(handle) = self.handle.take() { + handle.abort(); + } } } @@ -755,11 +740,11 @@ async fn handle_info_refs_upload_pack( git_protocol_version: Option<&str>, ) -> Result>, hyper::Error> { use std::process::Stdio; - use tokio::io::AsyncReadExt; use tokio::process::Command as TokioCommand; // Spawn git upload-pack --advertise-refs let mut cmd = TokioCommand::from(grasp_audit::git_command()); + cmd.kill_on_drop(true); cmd.arg("-c") .arg("uploadpack.allowReachableSHA1InWant=true") .arg("-c") @@ -780,7 +765,7 @@ async fn handle_info_refs_upload_pack( .stdout(Stdio::piped()) .stderr(Stdio::piped()); - let mut child = match cmd.spawn() { + let child = match cmd.spawn() { Ok(child) => child, Err(e) => { eprintln!("Failed to spawn git upload-pack: {}", e); @@ -791,20 +776,15 @@ async fn handle_info_refs_upload_pack( } }; - // Read stdout - let mut output = Vec::new(); - if let Some(mut stdout) = child.stdout.take() { - if let Err(e) = stdout.read_to_end(&mut output).await { - eprintln!("Failed to read git output: {}", e); - } - } - - // Wait for process - let status = child.wait().await; - if let Ok(s) = &status { - if !s.success() { - eprintln!("git upload-pack --advertise-refs failed"); - } + let output = match child.wait_with_output().await { + Ok(output) => output, + Err(error) => return Ok(git_io_failure(error)), + }; + if !output.status.success() { + eprintln!( + "git upload-pack --advertise-refs failed: {}", + String::from_utf8_lossy(&output.stderr) + ); } // Build response with pkt-line header @@ -821,7 +801,7 @@ async fn handle_info_refs_upload_pack( response_body.extend_from_slice(b"0000"); // Then the git output - response_body.extend_from_slice(&output); + response_body.extend_from_slice(&output.stdout); Ok(Response::builder() .status(StatusCode::OK) @@ -844,7 +824,6 @@ async fn handle_upload_pack( ) -> Result>, hyper::Error> { use http_body_util::BodyExt; use std::process::Stdio; - use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::process::Command as TokioCommand; // Read request body @@ -852,6 +831,7 @@ async fn handle_upload_pack( // Spawn git upload-pack let mut cmd = TokioCommand::from(grasp_audit::git_command()); + cmd.kill_on_drop(true); cmd.arg("-c") .arg("uploadpack.allowReachableSHA1InWant=true") .arg("-c") @@ -871,7 +851,7 @@ async fn handle_upload_pack( .stdout(Stdio::piped()) .stderr(Stdio::piped()); - let mut child = match cmd.spawn() { + let child = match cmd.spawn() { Ok(child) => child, Err(e) => { eprintln!("Failed to spawn git upload-pack: {}", e); @@ -882,51 +862,119 @@ async fn handle_upload_pack( } }; - // Write request body to stdin - if let Some(mut stdin) = child.stdin.take() { - if let Err(e) = stdin.write_all(&body_bytes).await { - eprintln!("Failed to write to git stdin: {}", e); - } - // Close stdin to signal end of input - drop(stdin); - } - - // Read stdout - let mut output = Vec::new(); - if let Some(mut stdout) = child.stdout.take() { - if let Err(e) = stdout.read_to_end(&mut output).await { - eprintln!("Failed to read git output: {}", e); - } - } - - // Read stderr for debugging - let mut stderr_output = Vec::new(); - if let Some(mut stderr) = child.stderr.take() { - let _ = stderr.read_to_end(&mut stderr_output).await; - } - - // Wait for process - let status = child.wait().await; - if let Ok(s) = &status { - if !s.success() { - let stderr_str = String::from_utf8_lossy(&stderr_output); - eprintln!("git upload-pack failed: {}", stderr_str); - } + let output = match collect_git_output(child, &body_bytes).await { + Ok(output) => output, + Err(error) => return Ok(git_io_failure(error)), + }; + if !output.status.success() { + eprintln!( + "git upload-pack failed: {}", + String::from_utf8_lossy(&output.stderr) + ); } Ok(Response::builder() .status(StatusCode::OK) .header("Content-Type", "application/x-git-upload-pack-result") .header("Cache-Control", "no-cache") - .body(Full::new(Bytes::from(output))) + .body(Full::new(Bytes::from(output.stdout))) .unwrap()) } +/// Write the request while draining both output pipes to prevent backpressure +/// deadlocks. Cancelling the future kills the command's child process. +async fn collect_git_output( + mut child: tokio::process::Child, + input: &[u8], +) -> std::io::Result { + use tokio::io::AsyncWriteExt; + let stdin = child.stdin.take(); + let write = async { + if let Some(mut stdin) = stdin { + stdin.write_all(input).await?; + } + Ok::<_, std::io::Error>(()) + }; + let (_, output) = tokio::try_join!(write, child.wait_with_output())?; + Ok(output) +} + +fn git_io_failure(error: std::io::Error) -> Response> { + eprintln!("Failed to collect git output: {error}"); + Response::builder() + .status(StatusCode::INTERNAL_SERVER_ERROR) + .body(Full::new(Bytes::from("Failed to collect git output"))) + .unwrap() +} + #[cfg(test)] mod smart_git_server_tests { use super::*; use crate::common::purgatory_helpers::{create_test_repo_with_commit, CommitVariant}; + async fn assert_stop_closes_connection( + url: String, + stop: impl std::future::Future, + ) { + use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader}; + tokio::time::timeout(std::time::Duration::from_secs(10), async { + let address = url.strip_prefix("http://").unwrap(); + let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); + stream + .write_all(b"GET /missing HTTP/1.1\r\nHost: localhost\r\n\r\n") + .await + .unwrap(); + let mut reader = BufReader::new(stream); + let mut status = String::new(); + reader.read_line(&mut status).await.unwrap(); + assert!(status.starts_with("HTTP/1.1 404")); + stop.await; + let mut remaining = Vec::new(); + reader.read_to_end(&mut remaining).await.unwrap(); + }) + .await + .expect("stopped server must close keep-alive connections"); + } + + #[tokio::test] + async fn both_git_servers_close_owned_connections_on_stop() { + let directory = tempfile::tempdir().unwrap(); + create_test_repo_with_commit(directory.path(), CommitVariant::StateTest).unwrap(); + let simple = SimpleGitServer::start(directory.path()).await; + assert_stop_closes_connection(simple.url().to_string(), simple.stop()).await; + let smart = SmartGitServer::start(directory.path()).await; + assert_stop_closes_connection(smart.url().to_string(), smart.stop()).await; + } + + #[cfg(unix)] + #[tokio::test] + async fn git_io_drains_output_while_writing_input() { + use std::process::Stdio; + let directory = tempfile::tempdir().unwrap(); + let payload = vec![b'x'; 1024 * 1024]; + let path = directory.path().join("payload"); + std::fs::write(&path, &payload).unwrap(); + let child = tokio::process::Command::new("sh") + .args(["-c", "cat \"$1\"; cat \"$1\" >&2; cat", "test-io"]) + .arg(&path) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .kill_on_drop(true) + .spawn() + .unwrap(); + let output = tokio::time::timeout( + std::time::Duration::from_secs(10), + collect_git_output(child, &payload), + ) + .await + .expect("all three pipes must make progress") + .unwrap(); + assert!(output.status.success()); + assert_eq!(output.stdout, payload.repeat(2)); + assert_eq!(output.stderr, payload); + } + #[tokio::test] async fn test_smart_git_server_starts_and_stops() { // Create a test repo