michal/tit

Diff

3733f1bebee95d154249d7c7

tests/git_http.rs

Mode 100644100644; object cdad45f9dddcf02331498bf1

@@ -1,16 +1,20 @@
 use crate::git::{http, packetline, transport, upload_pack};
 
 use std::fs;
-use std::io::{Read, Write};
 use std::net::{Ipv4Addr, SocketAddr};
 use std::os::unix::fs::PermissionsExt;
 use std::path::{Path, PathBuf};
 use std::process::Command;
+use std::time::Duration;
 
 use http::RunningGitHttpServer;
 use tempfile::TempDir;
+use tokio::io::{AsyncReadExt, AsyncWriteExt};
+use tokio::net::TcpStream;
 use transport::GitRepositories;
 use upload_pack::{ProtocolVersion, UploadPack, UploadPackError};
+
+const RESPONSE_TIMEOUT: Duration = Duration::from_secs(10);
 
 #[tokio::test(flavor = "multi_thread", worker_threads = 4)]
 async fn stock_git_clones_and_fetches_both_hash_formats_over_smart_http() {
@@ -111,13 +115,15 @@
         "application/x-git-upload-pack-request",
         Some("version=2"),
         b"zzzz",
-    );
+    )
+    .await;
     assert!(malformed.starts_with(b"HTTP/1.1 400"));
 
     let wrong_service = raw_http_get(
         server.address(),
         "/alice/example/info/refs?service=git-receive-pack",
-    );
+    )
+    .await;
     assert!(wrong_service.starts_with(b"HTTP/1.1 400"));
 
     let wrong_content_type = raw_http_request(
@@ -126,7 +132,8 @@
         "text/plain",
         Some("version=2"),
         b"0000",
-    );
+    )
+    .await;
     assert!(wrong_content_type.starts_with(b"HTTP/1.1 415"));
 
     let wrong_version = raw_http_request(
@@ -135,7 +142,8 @@
         "application/x-git-upload-pack-request",
         Some("version=2:extra"),
         b"0000",
-    );
+    )
+    .await;
     assert!(wrong_version.starts_with(b"HTTP/1.1 400"));
 
     let oversized = raw_http_request(
@@ -144,7 +152,8 @@
         "application/x-git-upload-pack-request",
         Some("version=2"),
         &vec![b'0'; packetline::MAX_REQUEST_BYTES + 1],
-    );
+    )
+    .await;
     assert!(oversized.starts_with(b"HTTP/1.1 413"));
     server.shutdown().await.expect("stop the Git HTTP server");
 }
@@ -310,43 +319,63 @@
         .to_owned()
 }
 
-fn raw_http_request(
+async fn raw_http_request(
     address: SocketAddr,
     path: &str,
     content_type: &str,
     git_protocol: Option<&str>,
     body: &[u8],
 ) -> Vec<u8> {
-    let mut stream = std::net::TcpStream::connect(address).expect("connect to the Git HTTP server");
-    let protocol_header = git_protocol
-        .map(|value| format!("Git-Protocol: {value}\r\n"))
-        .unwrap_or_default();
-    write!(
-        stream,
-        "POST {path} HTTP/1.1\r\nHost: {address}\r\nContent-Type: {content_type}\r\n{protocol_header}Content-Length: {}\r\nConnection: close\r\n\r\n",
-        body.len()
-    )
-    .expect("write HTTP request headers");
-    stream.write_all(body).expect("write HTTP request body");
-    let mut response = Vec::new();
-    stream
-        .read_to_end(&mut response)
-        .expect("read the Git HTTP response");
-    response
+    tokio::time::timeout(RESPONSE_TIMEOUT, async {
+        let mut stream = TcpStream::connect(address)
+            .await
+            .expect("connect to the Git HTTP server");
+        let protocol_header = git_protocol
+            .map(|value| format!("Git-Protocol: {value}\r\n"))
+            .unwrap_or_default();
+        let head = format!(
+            "POST {path} HTTP/1.1\r\nHost: {address}\r\nContent-Type: {content_type}\r\n{protocol_header}Content-Length: {}\r\nConnection: close\r\n\r\n",
+            body.len()
+        );
+        stream
+            .write_all(head.as_bytes())
+            .await
+            .expect("write HTTP request headers");
+        stream
+            .write_all(body)
+            .await
+            .expect("write HTTP request body");
+        let mut response = Vec::new();
+        stream
+            .read_to_end(&mut response)
+            .await
+            .expect("read the Git HTTP response");
+        response
+    })
+    .await
+    .expect("receive a Git HTTP response before the deadline")
 }
 
-fn raw_http_get(address: SocketAddr, path: &str) -> Vec<u8> {
-    let mut stream = std::net::TcpStream::connect(address).expect("connect to the Git HTTP server");
-    write!(
-        stream,
-        "GET {path} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\n\r\n"
-    )
-    .expect("write HTTP request");
-    let mut response = Vec::new();
-    stream
-        .read_to_end(&mut response)
-        .expect("read the Git HTTP response");
-    response
+async fn raw_http_get(address: SocketAddr, path: &str) -> Vec<u8> {
+    tokio::time::timeout(RESPONSE_TIMEOUT, async {
+        let mut stream = TcpStream::connect(address)
+            .await
+            .expect("connect to the Git HTTP server");
+        let request =
+            format!("GET {path} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\n\r\n");
+        stream
+            .write_all(request.as_bytes())
+            .await
+            .expect("write HTTP request");
+        let mut response = Vec::new();
+        stream
+            .read_to_end(&mut response)
+            .await
+            .expect("read the Git HTTP response");
+        response
+    })
+    .await
+    .expect("receive a Git HTTP response before the deadline")
 }
 
 fn create_fixture(worktree: &Path, bare: &Path, format: &str) {

tests/public_routes.rs

Mode 100644100644; object c30f949702d1bd30ef1fffcd

@@ -2,8 +2,7 @@
 
 use std::collections::BTreeMap;
 use std::fs;
-use std::io::{Read, Write};
-use std::net::{Ipv4Addr, SocketAddr, TcpStream};
+use std::net::{Ipv4Addr, SocketAddr};
 use std::path::Path;
 use std::process::Command;
 use std::time::{Duration, SystemTime, UNIX_EPOCH};
@@ -12,6 +11,8 @@
 use sha2::{Digest, Sha256};
 use store::{InitialAdministrator, NewRepository, RepositoryOrigin, Store};
 use tempfile::TempDir;
+use tokio::io::{AsyncReadExt, AsyncWriteExt};
+use tokio::net::TcpStream;
 
 const RESPONSE_TIMEOUT: Duration = Duration::from_secs(10);
 
@@ -33,7 +34,7 @@
     .await
     .expect("start the public Web server");
 
-    let summary = request(server.address(), "GET", "/alice/example", &[], &[]);
+    let summary = request(server.address(), "GET", "/alice/example", &[], &[]).await;
     assert_eq!(summary.status, 200);
     assert_html_policy(&summary);
     assert!(
@@ -45,7 +46,8 @@
     assert!(!summary.text().to_ascii_lowercase().contains("<script"));
     assert!(summary.text().contains("/alice/example/issues"));
 
-    let anonymous_issues = request(server.address(), "GET", "/alice/example/issues", &[], &[]);
+    let anonymous_issues =
+        request(server.address(), "GET", "/alice/example/issues", &[], &[]).await;
     assert_eq!(anonymous_issues.status, 200);
     assert!(
         anonymous_issues
@@ -76,6 +78,7 @@
             &headers,
             rejected.as_bytes(),
         )
+        .await
         .status,
         403
     );
@@ -91,7 +94,8 @@
         "/alice/example/issues",
         &headers,
         issue.as_bytes(),
-    );
+    )
+    .await;
     assert_eq!(created.status, 303);
     assert_eq!(created.header("location"), "/alice/example/issues/1");
 
@@ -101,7 +105,8 @@
         "/alice/example/issues/1",
         &[("Cookie", cookie.as_str())],
         &[],
-    );
+    )
+    .await;
     assert_eq!(detail.status, 200);
     assert!(detail.text().contains("#1 No JavaScript workflow"));
     assert!(detail.text().contains("<strong>safe</strong>"));
@@ -116,7 +121,9 @@
         )
         .expect("make the repository private");
     assert_eq!(
-        request(server.address(), "GET", "/alice/example", &[], &[]).status,
+        request(server.address(), "GET", "/alice/example", &[], &[])
+            .await
+            .status,
         404
     );
     assert_eq!(
@@ -127,6 +134,7 @@
             &[("Cookie", cookie.as_str())],
             &[],
         )
+        .await
         .status,
         200
     );
@@ -251,38 +259,43 @@
     );
 }
 
-fn request(
+async fn request(
     address: SocketAddr,
     method: &str,
     path: &str,
     headers: &[(&str, &str)],
     body: &[u8],
 ) -> HttpResponse {
-    let mut stream = TcpStream::connect(address).expect("connect to the public Web server");
-    stream
-        .set_read_timeout(Some(RESPONSE_TIMEOUT))
-        .expect("set the response timeout");
-    stream
-        .set_write_timeout(Some(RESPONSE_TIMEOUT))
-        .expect("set the request timeout");
-    let mut head = format!(
-        "{method} {path} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nContent-Length: {}\r\n",
-        body.len()
-    );
-    for (name, value) in headers {
-        head.push_str(&format!("{name}: {value}\r\n"));
-    }
-    head.push_str("\r\n");
-    stream
-        .write_all(head.as_bytes())
-        .expect("write HTTP request headers");
-    stream.write_all(body).expect("write the HTTP request");
-    let mut response = Vec::new();
-    if let Err(error) = stream.read_to_end(&mut response)
-        && error.kind() != std::io::ErrorKind::ConnectionReset
-    {
-        panic!("read an HTTP response: {error}");
-    }
+    let response = tokio::time::timeout(RESPONSE_TIMEOUT, async {
+        let mut stream = TcpStream::connect(address)
+            .await
+            .expect("connect to the public Web server");
+        let mut head = format!(
+            "{method} {path} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nContent-Length: {}\r\n",
+            body.len()
+        );
+        for (name, value) in headers {
+            head.push_str(&format!("{name}: {value}\r\n"));
+        }
+        head.push_str("\r\n");
+        stream
+            .write_all(head.as_bytes())
+            .await
+            .expect("write HTTP request headers");
+        stream
+            .write_all(body)
+            .await
+            .expect("write the HTTP request");
+        let mut response = Vec::new();
+        if let Err(error) = stream.read_to_end(&mut response).await
+            && error.kind() != std::io::ErrorKind::ConnectionReset
+        {
+            panic!("read an HTTP response: {error}");
+        }
+        response
+    })
+    .await
+    .expect("receive an HTTP response before the deadline");
     HttpResponse::parse(&response)
 }
 

tests/web_shell.rs

Mode 100644100644; object cbc1fb38ada92bdafb662ad2

@@ -1,18 +1,20 @@
 use crate::http;
 
 use std::collections::BTreeMap;
-use std::io::{Read, Write};
-use std::net::{Ipv4Addr, SocketAddr, TcpStream};
+use std::net::{Ipv4Addr, SocketAddr};
 use std::time::Duration;
 
 use http::RunningWebServer;
 use tokio::io::{AsyncReadExt, AsyncWriteExt};
+use tokio::net::TcpStream;
+
+const RESPONSE_TIMEOUT: Duration = Duration::from_secs(10);
 
 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
 async fn serves_the_semantic_shell_without_javascript() {
     let server = start().await;
 
-    let home = request(server.address(), "GET", "/", &[]);
+    let home = request(server.address(), "GET", "/", &[]).await;
     assert_eq!(home.status, 200);
     assert_eq!(home.header("content-type"), "text/html; charset=utf-8");
     assert_eq!(home.header("cache-control"), "no-store");
@@ -35,17 +37,18 @@
         "GET",
         "/go?owner=alice&repository=example",
         &[],
-    );
+    )
+    .await;
     assert_eq!(removed_repository_form.status, 404);
     assert_security_policy(&removed_repository_form);
 
-    let head = request(server.address(), "HEAD", "/", &[]);
+    let head = request(server.address(), "HEAD", "/", &[]).await;
     assert_eq!(head.status, 200);
     assert!(head.body.is_empty());
     assert_eq!(head.header("content-length"), home.body.len().to_string());
     assert_security_policy(&head);
 
-    let css = request(server.address(), "GET", "/assets/style.css", &[]);
+    let css = request(server.address(), "GET", "/assets/style.css", &[]).await;
     assert_eq!(css.status, 200);
     assert_eq!(css.header("content-type"), "text/css; charset=utf-8");
     assert_eq!(css.header("cache-control"), "no-cache");
@@ -54,7 +57,7 @@
     assert!(css.body.contains(".two-column"));
     assert_security_policy(&css);
 
-    let css_head = request(server.address(), "HEAD", "/assets/style.css", &[]);
+    let css_head = request(server.address(), "HEAD", "/assets/style.css", &[]).await;
     assert_eq!(css_head.status, 200);
     assert!(css_head.body.is_empty());
     assert_eq!(
@@ -62,7 +65,7 @@
         css.body.len().to_string()
     );
 
-    let signup = request(server.address(), "GET", "/signup", &[]);
+    let signup = request(server.address(), "GET", "/signup", &[]).await;
     assert_eq!(signup.status, 200);
     assert!(
         signup
@@ -70,7 +73,7 @@
             .contains("<form action=\"/signup\" method=\"post\">")
     );
     assert!(signup.body.contains("name=\"invitation\""));
-    let recovery = request(server.address(), "GET", "/recover", &[]);
+    let recovery = request(server.address(), "GET", "/recover", &[]).await;
     assert_eq!(recovery.status, 200);
     assert!(
         recovery
@@ -79,7 +82,7 @@
     );
     assert!(recovery.body.contains("name=\"recovery\""));
 
-    let wrong_signup_method = request(server.address(), "PUT", "/signup", &[]);
+    let wrong_signup_method = request(server.address(), "PUT", "/signup", &[]).await;
     assert_eq!(wrong_signup_method.status, 405);
     assert_eq!(wrong_signup_method.header("allow"), "GET, HEAD, POST");
 
@@ -90,14 +93,14 @@
 async fn serves_useful_errors_and_owns_request_ids() {
     let server = start().await;
 
-    let missing = request(server.address(), "GET", "/missing", &[]);
+    let missing = request(server.address(), "GET", "/missing", &[]).await;
     assert_eq!(missing.status, 404);
     assert!(missing.body.contains("<h1>Page not found</h1>"));
     assert!(missing.body.contains("The requested page does not exist."));
     assert_security_policy(&missing);
     assert_snapshot(&missing, include_str!("snapshots/web/not-found.html"));
 
-    let missing_head = request(server.address(), "HEAD", "/missing", &[]);
+    let missing_head = request(server.address(), "HEAD", "/missing", &[]).await;
     assert_eq!(missing_head.status, 404);
     assert!(missing_head.body.is_empty());
     assert_eq!(
@@ -105,7 +108,7 @@
         missing.body.len().to_string()
     );
 
-    let method = request(server.address(), "POST", "/", &[]);
+    let method = request(server.address(), "POST", "/", &[]).await;
     assert_eq!(method.status, 405);
     assert_eq!(method.header("allow"), "GET, HEAD");
     assert!(method.body.contains("<h1>Method not allowed</h1>"));
@@ -125,8 +128,9 @@
         "GET",
         "/",
         &[("X-Request-ID", "attacker-controlled")],
-    );
-    let second = request(server.address(), "GET", "/", &[]);
+    )
+    .await;
+    let second = request(server.address(), "GET", "/", &[]).await;
     assert_request_id(first.header("x-request-id"));
     assert_request_id(second.header("x-request-id"));
     assert_ne!(first.header("x-request-id"), "attacker-controlled");
@@ -151,15 +155,16 @@
                     ("X-Forwarded-For", "198.51.100.1")
                 ]
             )
+            .await
             .status,
             400
         );
     }
-    let limited = request(server.address(), "POST", "/login", &[]);
+    let limited = request(server.address(), "POST", "/login", &[]).await;
     assert_eq!(limited.status, 429);
     assert_eq!(limited.body, "Login attempt limit exceeded.\n");
 
-    let oversized = request_with_declared_length(server.address(), "/", 1024 * 1024 + 1);
+    let oversized = request_with_declared_length(server.address(), "/", 1024 * 1024 + 1).await;
     assert_eq!(oversized.status, 413);
 
     server.shutdown().await.expect("stop the Web server");
@@ -170,9 +175,12 @@
     for path in ["/signup", "/recover"] {
         let server = start().await;
         for _ in 0..10 {
-            assert_eq!(request(server.address(), "POST", path, &[]).status, 400);
+            assert_eq!(
+                request(server.address(), "POST", path, &[]).await.status,
+                400
+            );
         }
-        let limited = request(server.address(), "POST", path, &[]);
+        let limited = request(server.address(), "POST", path, &[]).await;
         assert_eq!(limited.status, 429);
         assert_eq!(limited.body, "Account attempt limit exceeded.\n");
         server.shutdown().await.expect("stop the Web server");
@@ -216,44 +224,63 @@
         .expect("start the Web server")
 }
 
-fn request(
+async fn request(
     address: SocketAddr,
     method: &str,
     path: &str,
     headers: &[(&str, &str)],
 ) -> HttpResponse {
-    let mut stream = TcpStream::connect(address).expect("connect to the Web server");
-    let mut request =
-        format!("{method} {path} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\n");
-    for (name, value) in headers {
-        request.push_str(&format!("{name}: {value}\r\n"));
-    }
-    request.push_str("Content-Length: 0\r\n\r\n");
-    stream
-        .write_all(request.as_bytes())
-        .expect("write an HTTP request");
-    let mut bytes = Vec::new();
-    stream
-        .read_to_end(&mut bytes)
-        .expect("read an HTTP response");
+    let bytes = tokio::time::timeout(RESPONSE_TIMEOUT, async {
+        let mut stream = TcpStream::connect(address)
+            .await
+            .expect("connect to the Web server");
+        let mut request =
+            format!("{method} {path} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\n");
+        for (name, value) in headers {
+            request.push_str(&format!("{name}: {value}\r\n"));
+        }
+        request.push_str("Content-Length: 0\r\n\r\n");
+        stream
+            .write_all(request.as_bytes())
+            .await
+            .expect("write an HTTP request");
+        let mut bytes = Vec::new();
+        stream
+            .read_to_end(&mut bytes)
+            .await
+            .expect("read an HTTP response");
+        bytes
+    })
+    .await
+    .expect("receive an HTTP response before the deadline");
     HttpResponse::parse(&bytes)
 }
 
-fn request_with_declared_length(
+async fn request_with_declared_length(
     address: SocketAddr,
     path: &str,
     content_length: usize,
 ) -> HttpResponse {
-    let mut stream = TcpStream::connect(address).expect("connect to the Web server");
-    write!(
-        stream,
-        "POST {path} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nContent-Length: {content_length}\r\n\r\n"
-    )
-    .expect("write an HTTP request");
-    let mut bytes = Vec::new();
-    stream
-        .read_to_end(&mut bytes)
-        .expect("read an HTTP response");
+    let bytes = tokio::time::timeout(RESPONSE_TIMEOUT, async {
+        let mut stream = TcpStream::connect(address)
+            .await
+            .expect("connect to the Web server");
+        let request = format!(
+            "POST {path} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nContent-Length: {content_length}\r\n\r\n"
+        );
+        stream
+            .write_all(request.as_bytes())
+            .await
+            .expect("write an HTTP request");
+        let mut bytes = Vec::new();
+        stream
+            .read_to_end(&mut bytes)
+            .await
+            .expect("read an HTTP response");
+        bytes
+    })
+    .await
+    .expect("receive an HTTP response before the deadline");
     HttpResponse::parse(&bytes)
 }