server: serve xrpc over http/3

Lewis: May this revision serve well! <lu5a@proton.me>
This commit is contained in:
Lewis
2026-06-10 13:18:49 +03:00
committed by Tangled
parent b009ccdaf2
commit 5bbe2146ff
15 changed files with 963 additions and 50 deletions
Generated
+169 -32
View File
@@ -210,7 +210,7 @@ version = "0.6.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5493c3bedbacf7fd7382c6346bbd66687d12bbaad3a89a2d2c303ee6cf20b048"
dependencies = [
"asn1-rs-derive",
"asn1-rs-derive 0.5.1",
"asn1-rs-impl",
"displaydoc",
"nom 7.1.3",
@@ -220,6 +220,22 @@ dependencies = [
"time",
]
[[package]]
name = "asn1-rs"
version = "0.7.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b7f43a50ac4fdca5df8e885c21b835997f0a1cdee65494a6847694a98652d9d8"
dependencies = [
"asn1-rs-derive 0.6.0",
"asn1-rs-impl",
"displaydoc",
"nom 7.1.3",
"num-traits",
"rusticata-macros",
"thiserror 2.0.18",
"time",
]
[[package]]
name = "asn1-rs-derive"
version = "0.5.1"
@@ -232,6 +248,18 @@ dependencies = [
"synstructure",
]
[[package]]
name = "asn1-rs-derive"
version = "0.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3109e49b1e4909e9db6515a30c633684d68cdeaa252f215214cb4fa1a5bfee2c"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.117",
"synstructure",
]
[[package]]
name = "asn1-rs-impl"
version = "0.2.0"
@@ -1046,7 +1074,7 @@ version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3"
dependencies = [
"bit-vec",
"bit-vec 0.8.0",
]
[[package]]
@@ -1055,6 +1083,15 @@ version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7"
[[package]]
name = "bit-vec"
version = "0.9.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b71798fca2c1fe1086445a7258a4bc81e6e49dcd24c8d0dd9a1e57395b603f51"
dependencies = [
"serde",
]
[[package]]
name = "bitflags"
version = "2.11.0"
@@ -1970,7 +2007,21 @@ version = "9.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5cd0a5c643689626bec213c4d8bd4d96acc8ffdb4ad4bb6bc16abf27d5f4b553"
dependencies = [
"asn1-rs",
"asn1-rs 0.6.2",
"displaydoc",
"nom 7.1.3",
"num-bigint",
"num-traits",
"rusticata-macros",
]
[[package]]
name = "der-parser"
version = "10.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "07da5016415d5a3c4dd39b11ed26f915f52fc4e0dc197d87908bc916e51bc1a6"
dependencies = [
"asn1-rs 0.7.2",
"displaydoc",
"nom 7.1.3",
"num-bigint",
@@ -2834,6 +2885,34 @@ dependencies = [
"tracing",
]
[[package]]
name = "h3"
version = "0.0.8"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "10872b55cfb02a821b69dc7cf8dc6a71d6af25eb9a79662bec4a9d016056b3be"
dependencies = [
"bytes",
"fastrand",
"futures-util",
"http 1.4.0",
"pin-project-lite",
"tokio",
]
[[package]]
name = "h3-quinn"
version = "0.0.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8b2e732c8d91a74731663ac8479ab505042fbf547b9a207213ab7fbcbfc4f8b4"
dependencies = [
"bytes",
"futures",
"h3",
"quinn",
"tokio",
"tokio-util",
]
[[package]]
name = "half"
version = "2.7.1"
@@ -4614,7 +4693,16 @@ version = "0.7.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a8d8034d9489cdaf79228eb9f6a3b8d7bb32ba00d6645ebd48eef4077ceb5bd9"
dependencies = [
"asn1-rs",
"asn1-rs 0.6.2",
]
[[package]]
name = "oid-registry"
version = "0.8.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "12f40cff3dde1b6087cc5d5f5d4d65712f34016a03ed60e9c08dcc392736b5b7"
dependencies = [
"asn1-rs 0.7.2",
]
[[package]]
@@ -5184,7 +5272,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "37566cb3fdacef14c0737f9546df7cfeadbfbc9fef10991038bf5015d0c80532"
dependencies = [
"bit-set",
"bit-vec",
"bit-vec 0.8.0",
"bitflags",
"num-traits",
"rand 0.9.2",
@@ -5427,6 +5515,7 @@ checksum = "b9e20a958963c291dc322d98411f541009df2ced7b5a4f2bd52337638cfccf20"
dependencies = [
"bytes",
"cfg_aliases",
"futures-io",
"pin-project-lite",
"quinn-proto",
"quinn-udp",
@@ -5630,6 +5719,20 @@ dependencies = [
"crossbeam-utils",
]
[[package]]
name = "rcgen"
version = "0.14.8"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "57f6d249aad744e274e682777a50283a225a32705394ee6d5fcc01efa25e4055"
dependencies = [
"pem",
"ring",
"rustls-pki-types",
"time",
"x509-parser 0.18.1",
"yasna",
]
[[package]]
name = "redis"
version = "1.1.0"
@@ -7527,7 +7630,7 @@ dependencies = [
[[package]]
name = "tranquil-api"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"anyhow",
"axum",
@@ -7578,7 +7681,7 @@ dependencies = [
[[package]]
name = "tranquil-auth"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"anyhow",
"base32",
@@ -7601,7 +7704,7 @@ dependencies = [
[[package]]
name = "tranquil-cache"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"async-trait",
"base64 0.22.1",
@@ -7615,7 +7718,7 @@ dependencies = [
[[package]]
name = "tranquil-comms"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"async-trait",
"base64 0.22.1",
@@ -7641,7 +7744,7 @@ dependencies = [
[[package]]
name = "tranquil-config"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"confique",
"serde",
@@ -7649,7 +7752,7 @@ dependencies = [
[[package]]
name = "tranquil-crypto"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"aes-gcm",
"base64 0.22.1",
@@ -7665,7 +7768,7 @@ dependencies = [
[[package]]
name = "tranquil-db"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"async-trait",
"chrono",
@@ -7682,7 +7785,7 @@ dependencies = [
[[package]]
name = "tranquil-db-traits"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"async-trait",
"base64 0.22.1",
@@ -7698,7 +7801,7 @@ dependencies = [
[[package]]
name = "tranquil-infra"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"async-trait",
"bytes",
@@ -7709,7 +7812,7 @@ dependencies = [
[[package]]
name = "tranquil-lexicon"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"chrono",
"futures",
@@ -7728,7 +7831,7 @@ dependencies = [
[[package]]
name = "tranquil-oauth"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"anyhow",
"axum",
@@ -7751,7 +7854,7 @@ dependencies = [
[[package]]
name = "tranquil-oauth-server"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"axum",
"base64 0.22.1",
@@ -7784,7 +7887,7 @@ dependencies = [
[[package]]
name = "tranquil-pds"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"aes-gcm",
"anyhow",
@@ -7878,7 +7981,7 @@ dependencies = [
[[package]]
name = "tranquil-repo"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"bytes",
"cid",
@@ -7890,7 +7993,7 @@ dependencies = [
[[package]]
name = "tranquil-ripple"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"async-trait",
"backon",
@@ -7915,7 +8018,7 @@ dependencies = [
[[package]]
name = "tranquil-scopes"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"axum",
"futures",
@@ -7931,17 +8034,23 @@ dependencies = [
[[package]]
name = "tranquil-server"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"arc-swap",
"axum",
"bytes",
"clap",
"dotenvy",
"ed25519-dalek",
"futures-util",
"h3",
"h3-quinn",
"hex",
"http 1.4.0",
"hyper 1.8.1",
"hyper-util",
"quinn",
"rcgen",
"rustls 0.23.37",
"rustls-pemfile",
"thiserror 2.0.18",
@@ -7961,7 +8070,7 @@ dependencies = [
[[package]]
name = "tranquil-signal"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"async-trait",
"chrono",
@@ -7984,7 +8093,7 @@ dependencies = [
[[package]]
name = "tranquil-storage"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"async-trait",
"aws-config",
@@ -8001,7 +8110,7 @@ dependencies = [
[[package]]
name = "tranquil-store"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"async-trait",
"bytes",
@@ -8050,7 +8159,7 @@ dependencies = [
[[package]]
name = "tranquil-sync"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"anyhow",
"axum",
@@ -8072,7 +8181,7 @@ dependencies = [
[[package]]
name = "tranquil-types"
version = "0.6.4"
version = "0.6.5"
dependencies = [
"chrono",
"cid",
@@ -8585,7 +8694,7 @@ checksum = "15784340a24c170ce60567282fb956a0938742dbfbf9eff5df793a686a009b8b"
dependencies = [
"base64 0.21.7",
"base64urlsafedata",
"der-parser",
"der-parser 9.0.0",
"hex",
"nom 7.1.3",
"openssl",
@@ -8601,7 +8710,7 @@ dependencies = [
"uuid",
"webauthn-attestation-ca",
"webauthn-rs-proto",
"x509-parser",
"x509-parser 0.16.0",
]
[[package]]
@@ -9155,17 +9264,35 @@ version = "0.16.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fcbc162f30700d6f3f82a24bf7cc62ffe7caea42c0b2cba8bf7f3ae50cf51f69"
dependencies = [
"asn1-rs",
"asn1-rs 0.6.2",
"data-encoding",
"der-parser",
"der-parser 9.0.0",
"lazy_static",
"nom 7.1.3",
"oid-registry",
"oid-registry 0.7.1",
"rusticata-macros",
"thiserror 1.0.69",
"time",
]
[[package]]
name = "x509-parser"
version = "0.18.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d43b0f71ce057da06bc0851b23ee24f3f86190b07203dd8f567d0b706a185202"
dependencies = [
"asn1-rs 0.7.2",
"data-encoding",
"der-parser 10.0.0",
"lazy_static",
"nom 7.1.3",
"oid-registry 0.8.1",
"ring",
"rusticata-macros",
"thiserror 2.0.18",
"time",
]
[[package]]
name = "xattr"
version = "1.6.1"
@@ -9194,6 +9321,16 @@ version = "1.0.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cfe53a6657fd280eaa890a3bc59152892ffa3e30101319d168b781ed6529b049"
[[package]]
name = "yasna"
version = "0.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b5f6765e852b9b4dc8e2a76843e4d64d1cea8e79bcde0b6901aea8e7c7f08282"
dependencies = [
"bit-vec 0.9.1",
"time",
]
[[package]]
name = "yoke"
version = "0.8.1"
+5 -1
View File
@@ -26,7 +26,7 @@ members = [
]
[workspace.package]
version = "0.6.4"
version = "0.6.5"
edition = "2024"
license = "AGPL-3.0-or-later"
@@ -82,6 +82,8 @@ foca = { version = "1", features = ["bincode-codec", "tracing"] }
futures = "0.3"
futures-util = "0.3"
governor = "0.10"
h3 = "0.0.8"
h3-quinn = "0.0.10"
hex = "0.4"
hickory-resolver = { version = "0.24", features = ["tokio-runtime"] }
hkdf = "0.12"
@@ -106,7 +108,9 @@ parking_lot = "0.12"
multihash = "0.19"
p256 = { version = "0.13", features = ["ecdsa"] }
p384 = { version = "0.13", features = ["ecdsa"] }
quinn = { version = "0.11", default-features = false, features = ["runtime-tokio", "rustls-ring", "log"] }
rand = "0.8"
rcgen = { version = "0.14", default-features = false, features = ["ring", "pem"] }
redis = { version = "1.0", features = ["tokio-comp", "connection-manager"] }
regex = "1"
rsa = "0.9"
+12 -2
View File
@@ -15,8 +15,12 @@ pub mod temp;
use tranquil_pds::state::AppState;
pub fn api_routes() -> axum::Router<AppState> {
use axum::extract::DefaultBodyLimit;
use axum::routing::{get, post};
let blob_body_limit =
DefaultBodyLimit::max(tranquil_config::get().server.max_blob_size as usize);
axum::Router::new()
.route("/_health", get(server::health))
.route(
@@ -68,7 +72,10 @@ pub fn api_routes() -> axum::Router<AppState> {
.route("/com.atproto.repo.deleteRecord", post(repo::delete_record))
.route("/com.atproto.repo.listRecords", get(repo::list_records))
.route("/com.atproto.repo.describeRepo", get(repo::describe_repo))
.route("/com.atproto.repo.uploadBlob", post(repo::upload_blob))
.route(
"/com.atproto.repo.uploadBlob",
post(repo::upload_blob).layer(blob_body_limit),
)
.route("/com.atproto.repo.applyWrites", post(repo::apply_writes))
.route(
"/com.atproto.server.checkAccountStatus",
@@ -247,7 +254,10 @@ pub fn api_routes() -> axum::Router<AppState> {
"/_identity.verifyHandleOwnership",
post(identity::verify_handle_ownership),
)
.route("/com.atproto.repo.importRepo", post(repo::import_repo))
.route(
"/com.atproto.repo.importRepo",
post(repo::import_repo).layer(blob_body_limit),
)
.route(
"/com.atproto.admin.deleteAccount",
post(admin::delete_account),
+3 -1
View File
@@ -197,7 +197,9 @@ pub async fn create_session(
&row.did,
&twofa_ctx,
input.auth_factor_token.as_deref(),
async |t: &str| crate::server::totp::verify_totp_or_backup_for_user(&state, &row.did, t).await,
async |t: &str| {
crate::server::totp::verify_totp_or_backup_for_user(&state, &row.did, t).await
},
)
.await
{
+30
View File
@@ -513,6 +513,11 @@ pub struct TlsConfig {
/// Path to the TLS private key.
#[config(env = "TLS_KEY_PATH")]
pub key_path: Option<String>,
/// Serve HTTP/3 over QUIC on the same UDP port as the TCP listener.
/// Requires cert_path and key_path.
#[config(env = "TLS_HTTP3", default = false)]
pub http3: bool,
}
impl TlsConfig {
@@ -532,6 +537,13 @@ impl TlsConfig {
.to_string(),
);
}
if self.http3 && self.material().is_none() {
errors.push(
"server.tls.http3 (TLS_HTTP3) requires server.tls.cert_path \
and erver.tls.key_path"
.to_string(),
);
}
}
}
@@ -1799,6 +1811,7 @@ port = 587
TlsConfig {
cert_path: None,
key_path: None,
http3: false,
}
.validate(&mut errors);
assert!(errors.is_empty(), "expected no errors, got {errors:?}");
@@ -1810,6 +1823,7 @@ port = 587
TlsConfig {
cert_path: Some("/etc/tranquil/cert.pem".to_string()),
key_path: Some("/etc/tranquil/key.pem".to_string()),
http3: false,
}
.validate(&mut errors);
assert!(errors.is_empty(), "expected no errors, got {errors:?}");
@@ -1821,6 +1835,7 @@ port = 587
TlsConfig {
cert_path: Some("/etc/tranquil/cert.pem".to_string()),
key_path: None,
http3: false,
}
.validate(&mut errors);
assert!(
@@ -1829,6 +1844,21 @@ port = 587
);
}
#[test]
fn tls_validate_rejects_http3_without_material() {
let mut errors = Vec::new();
TlsConfig {
cert_path: None,
key_path: None,
http3: true,
}
.validate(&mut errors);
assert!(
errors.iter().any(|e| e.contains("http3")),
"expected http3 error, got {errors:?}"
);
}
#[derive(Default)]
struct EmailOverrides {
from_address: Option<&'static str>,
+1 -2
View File
@@ -133,8 +133,7 @@ fn current_timestamp() -> u64 {
pub fn looks_like_totp_token(code: &str) -> bool {
let c = code.trim();
(c.len() == 6 && c.bytes().all(|b| b.is_ascii_digit()))
|| crate::auth::is_backup_code_format(c)
(c.len() == 6 && c.bytes().all(|b| b.is_ascii_digit())) || crate::auth::is_backup_code_format(c)
}
pub fn used_totp_factor(has_totp: bool, auth_factor_token: Option<&str>) -> bool {
+3 -3
View File
@@ -51,6 +51,8 @@ pub const BUILD_VERSION: &str = concat!(
#[cfg(not(debug_assertions))]
pub const BUILD_VERSION: &str = env!("CARGO_PKG_VERSION");
pub const GENERAL_BODY_LIMIT: usize = 16 * 1024 * 1024;
pub struct ExternalRoutes {
pub xrpc: Router<AppState>,
pub oauth: Router<AppState>,
@@ -97,9 +99,7 @@ pub fn app_with_routes(state: AppState, external: ExternalRoutes) -> Router {
.nest("/.well-known", well_known_router)
.route("/metrics", get(metrics::metrics_handler))
.merge(external.extra)
.layer(DefaultBodyLimit::max(
tranquil_config::get().server.max_blob_size as usize,
))
.layer(DefaultBodyLimit::max(GENERAL_BODY_LIMIT))
.layer(axum::middleware::map_response(rewrite_extractor_errors))
.layer(middleware::from_fn(metrics::metrics_middleware))
.layer(
@@ -80,7 +80,11 @@ async fn car_export_missing_block_is_repairable() {
let root = build_tree(&source).await;
let pristine = open_store(&dir.path().join("pristine"));
let head_block = source.get(&root).await.expect("read root").expect("root present");
let head_block = source
.get(&root)
.await
.expect("read root")
.expect("root present");
pristine.put(&head_block).await.expect("seed root only");
let err = generate_repo_car(&pristine, &root)
+3 -1
View File
@@ -933,7 +933,9 @@ pub async fn get_test_block_store() -> &'static tranquil_pds::repo::AnyBlockStor
#[allow(dead_code)]
pub async fn get_test_app_state() -> &'static AppState {
base_url().await;
TEST_APP_STATE.get().expect("TEST_APP_STATE not initialized")
TEST_APP_STATE
.get()
.expect("TEST_APP_STATE not initialized")
}
#[allow(dead_code)]
+15 -3
View File
@@ -86,7 +86,11 @@ async fn repair_fails_loud_on_missing_leaf_block() {
assert!(!records.is_empty(), "repo must contain records");
let leaf_cid = Cid::from_str(records[0].record_cid.as_str()).expect("parse leaf cid");
assert!(
block_store.get(&leaf_cid).await.expect("read leaf").is_some(),
block_store
.get(&leaf_cid)
.await
.expect("read leaf")
.is_some(),
"leaf must be present before corruption"
);
@@ -108,7 +112,11 @@ async fn repair_fails_loud_on_missing_leaf_block() {
.expect("delete leaf record block");
assert!(
block_store.get(&mst_root_cid).await.expect("read").is_none(),
block_store
.get(&mst_root_cid)
.await
.expect("read")
.is_none(),
"mst root node must be gone to force a structural repair"
);
assert!(
@@ -126,7 +134,11 @@ async fn repair_fails_loud_on_missing_leaf_block() {
);
assert!(
block_store.get(&mst_root_cid).await.expect("read").is_some(),
block_store
.get(&mst_root_cid)
.await
.expect("read")
.is_some(),
"structural repair must still re-insert the regenerable MST node"
);
+8
View File
@@ -14,13 +14,18 @@ tranquil-signal = { workspace = true }
arc-swap = { workspace = true }
axum = { workspace = true }
bytes = { workspace = true }
clap = { workspace = true }
dotenvy = { workspace = true }
ed25519-dalek = { workspace = true }
futures-util = { workspace = true }
h3 = { workspace = true }
h3-quinn = { workspace = true }
hex = { workspace = true }
http = { workspace = true }
hyper = { workspace = true }
hyper-util = { workspace = true }
quinn = { workspace = true }
rustls = { workspace = true }
rustls-pemfile = { workspace = true }
thiserror = { workspace = true }
@@ -31,6 +36,9 @@ tower = { workspace = true }
tracing = { workspace = true }
tracing-subscriber = { workspace = true }
[dev-dependencies]
rcgen = { workspace = true }
[features]
default = ["frontend", "s3", "valkey"]
frontend = ["tranquil-pds/frontend"]
+666
View File
@@ -0,0 +1,666 @@
use std::net::SocketAddr;
use std::sync::Arc;
use std::time::Duration;
use axum::Router;
use axum::body::Body;
use axum::extract::ConnectInfo;
use bytes::{Buf, Bytes};
use futures_util::StreamExt;
use http::header::ALT_SVC;
use http::{HeaderValue, Request, Response, StatusCode};
use quinn::crypto::rustls::QuicServerConfig;
use quinn::{Endpoint, Incoming, ServerConfig, TransportConfig, VarInt};
use tokio::sync::Semaphore;
use tokio_util::sync::CancellationToken;
use tokio_util::task::TaskTracker;
use tower::ServiceExt;
use tracing::debug;
use crate::tls::{ReloadableCertResolver, TlsError};
const MAX_CONCURRENT_BIDI_STREAMS: u32 = 256;
const MAX_CONCURRENT_CONNECTIONS: usize = 512;
const IDLE_TIMEOUT: Duration = Duration::from_secs(30);
const SHUTDOWN_GRACE: Duration = Duration::from_secs(10);
const ALT_SVC_MAX_AGE_SECS: u32 = 86_400;
pub fn build_quic_server_config(
resolver: Arc<ReloadableCertResolver>,
) -> Result<ServerConfig, TlsError> {
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut crypto = rustls::ServerConfig::builder_with_provider(provider)
.with_protocol_versions(&[&rustls::version::TLS13])
.map_err(|e| TlsError::Config(e.to_string()))?
.with_no_client_auth()
.with_cert_resolver(resolver);
crypto.alpn_protocols = vec![b"h3".to_vec()];
let quic_crypto =
QuicServerConfig::try_from(crypto).map_err(|e| TlsError::Config(e.to_string()))?;
let mut config = ServerConfig::with_crypto(Arc::new(quic_crypto));
let mut transport = TransportConfig::default();
transport.max_concurrent_bidi_streams(VarInt::from_u32(MAX_CONCURRENT_BIDI_STREAMS));
transport.max_idle_timeout(Some(
IDLE_TIMEOUT
.try_into()
.expect("idle timeout fits in varint"),
));
config.transport_config(Arc::new(transport));
Ok(config)
}
pub fn alt_svc_header(port: u16) -> HeaderValue {
HeaderValue::from_str(&format!("h3=\":{port}\"; ma={ALT_SVC_MAX_AGE_SECS}"))
.expect("alt-svc header value is valid ascii")
}
pub fn with_alt_svc(app: Router, port: u16) -> Router {
let value = alt_svc_header(port);
app.layer(axum::middleware::map_response(
move |mut response: Response<Body>| {
let value = value.clone();
async move {
if response.status() != StatusCode::SWITCHING_PROTOCOLS {
response.headers_mut().insert(ALT_SVC, value);
}
response
}
},
))
}
pub fn with_host_from_authority(app: Router) -> Router {
app.layer(axum::middleware::map_request(
|mut request: Request<Body>| async move {
let authority = request
.uri()
.authority()
.map(|a| HeaderValue::from_str(a.as_str()));
match (request.headers().contains_key(http::header::HOST), authority) {
(false, Some(Ok(value))) => {
request.headers_mut().insert(http::header::HOST, value);
request
}
_ => request,
}
},
))
}
pub async fn serve_http3(endpoint: Endpoint, app: Router, shutdown: CancellationToken) {
let tracker = TaskTracker::new();
let conn_limiter = Arc::new(Semaphore::new(MAX_CONCURRENT_CONNECTIONS));
loop {
tokio::select! {
_ = shutdown.cancelled() => break,
incoming = endpoint.accept() => {
let Some(incoming) = incoming else { break };
let Ok(permit) = conn_limiter.clone().try_acquire_owned() else {
debug!(
peer = %incoming.remote_address(),
max = MAX_CONCURRENT_CONNECTIONS,
"refusing h3 connection: limit reached"
);
incoming.refuse();
continue;
};
let app = app.clone();
let conn_shutdown = shutdown.clone();
let conn_tracker = tracker.clone();
tracker.spawn(async move {
let _permit = permit;
if let Err(e) = serve_connection(incoming, app, conn_shutdown, conn_tracker).await {
debug!("h3 connection ended: {e}");
}
});
}
}
}
tracker.close();
if tokio::time::timeout(SHUTDOWN_GRACE, tracker.wait())
.await
.is_err()
{
debug!("h3 connections did not drain within grace, closing");
}
endpoint.close(0u32.into(), b"shutdown");
endpoint.wait_idle().await;
}
async fn serve_connection(
incoming: Incoming,
app: Router,
shutdown: CancellationToken,
tracker: TaskTracker,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let conn = incoming.await?;
let remote = conn.remote_address();
let mut h3_conn =
h3::server::Connection::<_, Bytes>::new(h3_quinn::Connection::new(conn)).await?;
let mut draining = false;
loop {
tokio::select! {
_ = shutdown.cancelled(), if !draining => {
draining = true;
let _ = h3_conn.shutdown(0).await;
}
resolved = h3_conn.accept() => match resolved {
Ok(Some(resolver)) => {
let app = app.clone();
tracker.spawn(async move {
if let Err(e) = serve_request(resolver, app, remote).await {
debug!(peer = %remote, "h3 request failed: {e}");
}
});
}
Ok(None) => break,
Err(e) => {
debug!(peer = %remote, "h3 accept error: {e}");
break;
}
}
}
}
Ok(())
}
async fn serve_request(
resolver: h3::server::RequestResolver<h3_quinn::Connection, Bytes>,
app: Router,
remote: SocketAddr,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let (req, stream) = resolver.resolve_request().await?;
let (mut send, recv) = stream.split();
let (mut parts, ()) = req.into_parts();
parts.extensions.insert(ConnectInfo(remote));
let request = Request::from_parts(parts, request_body(recv));
let response = match app.oneshot(request).await {
Ok(response) => response,
Err(infallible) => match infallible {},
};
let (parts, body) = response.into_parts();
send.send_response(Response::from_parts(parts, ())).await?;
let mut data = body.into_data_stream();
while let Some(chunk) = data.next().await {
match chunk {
Ok(bytes) if bytes.has_remaining() => send.send_data(bytes).await?,
Ok(_) => {}
Err(e) => {
debug!(peer = %remote, "h3 response body error: {e}");
send.stop_stream(h3::error::Code::H3_INTERNAL_ERROR);
return Ok(());
}
}
}
send.finish().await?;
Ok(())
}
struct RecvGuard {
stream: h3::server::RequestStream<h3_quinn::RecvStream, Bytes>,
ended: bool,
}
impl Drop for RecvGuard {
fn drop(&mut self) {
if !self.ended {
self.stream.stop_sending(h3::error::Code::H3_NO_ERROR);
}
}
}
fn request_body(recv: h3::server::RequestStream<h3_quinn::RecvStream, Bytes>) -> Body {
let guard = RecvGuard {
stream: recv,
ended: false,
};
let stream = futures_util::stream::unfold(Some(guard), |state| async move {
let mut guard = state?;
match guard.stream.recv_data().await {
Ok(Some(mut buf)) => {
let bytes = buf.copy_to_bytes(buf.remaining());
Some((Ok::<Bytes, std::io::Error>(bytes), Some(guard)))
}
Ok(None) => {
guard.ended = true;
None
}
Err(e) => {
guard.ended = true;
Some((Err(std::io::Error::other(e.to_string())), None))
}
}
});
Body::from_stream(stream)
}
#[cfg(test)]
mod tests {
use super::*;
use axum::routing::get;
use rustls::pki_types::{
CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, ServerName, UnixTime,
};
fn self_signed_resolver() -> Arc<ReloadableCertResolver> {
let cert = rcgen::generate_simple_self_signed(vec!["localhost".to_string()]).unwrap();
let cert_der = cert.cert.der().clone();
let key_der =
PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der()));
let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der).unwrap();
let certified = rustls::sign::CertifiedKey::new(vec![cert_der], signing_key);
Arc::new(ReloadableCertResolver::new(certified))
}
#[derive(Debug)]
struct SkipServerVerification(Arc<rustls::crypto::CryptoProvider>);
impl rustls::client::danger::ServerCertVerifier for SkipServerVerification {
fn verify_server_cert(
&self,
_end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp_response: &[u8],
_now: UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&self.0.signature_verification_algorithms,
)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&self.0.signature_verification_algorithms,
)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
self.0.signature_verification_algorithms.supported_schemes()
}
}
fn client_endpoint() -> Endpoint {
let provider = Arc::new(rustls::crypto::ring::default_provider());
let mut crypto = rustls::ClientConfig::builder_with_provider(provider.clone())
.with_protocol_versions(&[&rustls::version::TLS13])
.unwrap()
.dangerous()
.with_custom_certificate_verifier(Arc::new(SkipServerVerification(provider)))
.with_no_client_auth();
crypto.alpn_protocols = vec![b"h3".to_vec()];
let quic = quinn::crypto::rustls::QuicClientConfig::try_from(crypto).unwrap();
let mut endpoint = Endpoint::client("0.0.0.0:0".parse().unwrap()).unwrap();
endpoint.set_default_client_config(quinn::ClientConfig::new(Arc::new(quic)));
endpoint
}
#[tokio::test]
async fn h3_get_roundtrips_through_router() {
let app = Router::new().route("/", get(|| async { "ok" }));
let server_config = build_quic_server_config(self_signed_resolver()).unwrap();
let server = Endpoint::server(server_config, "127.0.0.1:0".parse().unwrap()).unwrap();
let addr = server.local_addr().unwrap();
let shutdown = CancellationToken::new();
tokio::spawn(serve_http3(server, app, shutdown.clone()));
let client = client_endpoint();
let conn = client.connect(addr, "localhost").unwrap().await.unwrap();
let (mut driver, mut send_request) = h3::client::new(h3_quinn::Connection::new(conn))
.await
.unwrap();
let drive =
tokio::spawn(async move { std::future::poll_fn(|cx| driver.poll_close(cx)).await });
let req = Request::get("https://localhost/").body(()).unwrap();
let mut stream = send_request.send_request(req).await.unwrap();
stream.finish().await.unwrap();
let response = stream.recv_response().await.unwrap();
assert_eq!(response.status(), 200);
let mut body = Vec::new();
while let Some(mut chunk) = stream.recv_data().await.unwrap() {
let bytes = chunk.copy_to_bytes(chunk.remaining());
body.extend_from_slice(&bytes);
}
assert_eq!(body.as_slice(), b"ok");
shutdown.cancel();
drive.abort();
}
#[test]
fn alt_svc_header_advertises_h3() {
assert_eq!(
alt_svc_header(443).to_str().unwrap(),
"h3=\":443\"; ma=86400"
);
}
#[tokio::test]
async fn alt_svc_added_to_responses_except_switching_protocols() {
let app = with_alt_svc(
Router::new().route("/ok", get(|| async { "ok" })).route(
"/upgrade",
get(|| async {
Response::builder()
.status(StatusCode::SWITCHING_PROTOCOLS)
.body(Body::empty())
.unwrap()
}),
),
443,
);
let normal = app
.clone()
.oneshot(Request::get("/ok").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(
normal.headers().get(ALT_SVC).and_then(|v| v.to_str().ok()),
Some("h3=\":443\"; ma=86400")
);
let upgrade = app
.oneshot(Request::get("/upgrade").body(Body::empty()).unwrap())
.await
.unwrap();
assert!(
upgrade.headers().get(ALT_SVC).is_none(),
"101 responses must not carry Alt-Svc"
);
}
fn make_cert(dns: &str) -> (rustls::sign::CertifiedKey, CertificateDer<'static>) {
let cert = rcgen::generate_simple_self_signed(vec![dns.to_string()]).unwrap();
let cert_der = cert.cert.der().clone();
let key_der =
PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der()));
let signing_key = rustls::crypto::ring::sign::any_supported_type(&key_der).unwrap();
let certified = rustls::sign::CertifiedKey::new(vec![cert_der.clone()], signing_key);
(certified, cert_der)
}
#[derive(Debug)]
struct RecordingVerifier {
provider: Arc<rustls::crypto::CryptoProvider>,
seen: Arc<std::sync::Mutex<Vec<u8>>>,
}
impl rustls::client::danger::ServerCertVerifier for RecordingVerifier {
fn verify_server_cert(
&self,
end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp_response: &[u8],
_now: UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
*self.seen.lock().unwrap() = end_entity.as_ref().to_vec();
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&self.provider.signature_verification_algorithms,
)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&self.provider.signature_verification_algorithms,
)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
self.provider
.signature_verification_algorithms
.supported_schemes()
}
}
async fn observe_server_cert(addr: SocketAddr) -> Vec<u8> {
let provider = Arc::new(rustls::crypto::ring::default_provider());
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let verifier = Arc::new(RecordingVerifier {
provider: provider.clone(),
seen: seen.clone(),
});
let mut crypto = rustls::ClientConfig::builder_with_provider(provider)
.with_protocol_versions(&[&rustls::version::TLS13])
.unwrap()
.dangerous()
.with_custom_certificate_verifier(verifier)
.with_no_client_auth();
crypto.alpn_protocols = vec![b"h3".to_vec()];
let quic = quinn::crypto::rustls::QuicClientConfig::try_from(crypto).unwrap();
let mut endpoint = Endpoint::client("127.0.0.1:0".parse().unwrap()).unwrap();
endpoint.set_default_client_config(quinn::ClientConfig::new(Arc::new(quic)));
let conn = endpoint.connect(addr, "localhost").unwrap().await.unwrap();
conn.close(0u32.into(), b"done");
endpoint.wait_idle().await;
seen.lock().unwrap().clone()
}
#[tokio::test]
async fn quic_handshake_observes_reloaded_certificate() {
let (cert_a, der_a) = make_cert("localhost");
let (cert_b, der_b) = make_cert("localhost");
assert_ne!(der_a, der_b, "test must use two distinct certs");
let resolver = Arc::new(ReloadableCertResolver::new(cert_a));
let server_config = build_quic_server_config(resolver.clone()).unwrap();
let server = Endpoint::server(server_config, "127.0.0.1:0".parse().unwrap()).unwrap();
let addr = server.local_addr().unwrap();
let shutdown = CancellationToken::new();
let app = Router::new().route("/", get(|| async { "ok" }));
tokio::spawn(serve_http3(server, app, shutdown.clone()));
let before = observe_server_cert(addr).await;
assert_eq!(
before.as_slice(),
der_a.as_ref(),
"first handshake must present the original cert"
);
resolver.store(cert_b);
let after = observe_server_cert(addr).await;
assert_eq!(
after.as_slice(),
der_b.as_ref(),
"handshake after reload must present the new cert"
);
assert_ne!(before, after, "reload must change the presented cert");
shutdown.cancel();
}
#[tokio::test]
async fn h3_requests_carry_remote_addr_connect_info() {
let app = Router::new().route(
"/",
get(|ConnectInfo(addr): ConnectInfo<SocketAddr>| async move { addr.to_string() }),
);
let server_config = build_quic_server_config(self_signed_resolver()).unwrap();
let server = Endpoint::server(server_config, "127.0.0.1:0".parse().unwrap()).unwrap();
let addr = server.local_addr().unwrap();
let shutdown = CancellationToken::new();
tokio::spawn(serve_http3(server, app, shutdown.clone()));
let client = client_endpoint();
let client_port = client.local_addr().unwrap().port();
let conn = client.connect(addr, "localhost").unwrap().await.unwrap();
let (mut driver, mut send_request) = h3::client::new(h3_quinn::Connection::new(conn))
.await
.unwrap();
let drive =
tokio::spawn(async move { std::future::poll_fn(|cx| driver.poll_close(cx)).await });
let req = Request::get("https://localhost/").body(()).unwrap();
let mut stream = send_request.send_request(req).await.unwrap();
stream.finish().await.unwrap();
assert_eq!(stream.recv_response().await.unwrap().status(), 200);
let mut body = Vec::new();
while let Some(mut chunk) = stream.recv_data().await.unwrap() {
let bytes = chunk.copy_to_bytes(chunk.remaining());
body.extend_from_slice(&bytes);
}
let reported: SocketAddr = String::from_utf8(body).unwrap().parse().unwrap();
assert!(reported.ip().is_loopback());
assert_eq!(
reported.port(),
client_port,
"handlers must see the QUIC remote address via ConnectInfo"
);
shutdown.cancel();
drive.abort();
}
#[tokio::test]
async fn host_header_filled_from_authority() {
let app = with_host_from_authority(Router::new().route(
"/",
get(|headers: http::HeaderMap| async move {
headers
.get(http::header::HOST)
.and_then(|v| v.to_str().ok())
.map(str::to_owned)
.unwrap_or_default()
}),
));
let server_config = build_quic_server_config(self_signed_resolver()).unwrap();
let server = Endpoint::server(server_config, "127.0.0.1:0".parse().unwrap()).unwrap();
let addr = server.local_addr().unwrap();
let shutdown = CancellationToken::new();
tokio::spawn(serve_http3(server, app, shutdown.clone()));
let client = client_endpoint();
let conn = client.connect(addr, "localhost").unwrap().await.unwrap();
let (mut driver, mut send_request) = h3::client::new(h3_quinn::Connection::new(conn))
.await
.unwrap();
let drive =
tokio::spawn(async move { std::future::poll_fn(|cx| driver.poll_close(cx)).await });
let req = Request::get("https://localhost/").body(()).unwrap();
let mut stream = send_request.send_request(req).await.unwrap();
stream.finish().await.unwrap();
assert_eq!(stream.recv_response().await.unwrap().status(), 200);
let mut body = Vec::new();
while let Some(mut chunk) = stream.recv_data().await.unwrap() {
let bytes = chunk.copy_to_bytes(chunk.remaining());
body.extend_from_slice(&bytes);
}
assert_eq!(
String::from_utf8(body).unwrap(),
"localhost",
"handlers must see the authority as the Host header"
);
shutdown.cancel();
drive.abort();
}
#[tokio::test]
async fn h3_body_error_resets_stream_instead_of_truncating() {
let app = Router::new().route(
"/",
get(|| async {
Body::from_stream(futures_util::stream::iter(vec![
Ok::<Bytes, std::io::Error>(Bytes::from_static(b"partial")),
Err(std::io::Error::other("body source failed")),
]))
}),
);
let server_config = build_quic_server_config(self_signed_resolver()).unwrap();
let server = Endpoint::server(server_config, "127.0.0.1:0".parse().unwrap()).unwrap();
let addr = server.local_addr().unwrap();
let shutdown = CancellationToken::new();
tokio::spawn(serve_http3(server, app, shutdown.clone()));
let client = client_endpoint();
let conn = client.connect(addr, "localhost").unwrap().await.unwrap();
let (mut driver, mut send_request) = h3::client::new(h3_quinn::Connection::new(conn))
.await
.unwrap();
let drive =
tokio::spawn(async move { std::future::poll_fn(|cx| driver.poll_close(cx)).await });
let req = Request::get("https://localhost/").body(()).unwrap();
let mut stream = send_request.send_request(req).await.unwrap();
stream.finish().await.unwrap();
let outcome = async {
stream.recv_response().await?;
let mut body = Vec::new();
loop {
match stream.recv_data().await {
Ok(Some(mut chunk)) => {
let bytes = chunk.copy_to_bytes(chunk.remaining());
body.extend_from_slice(&bytes);
}
Ok(None) => return Ok(body),
Err(e) => return Err(e),
}
}
}
.await;
assert!(
outcome.is_err(),
"a mid-body error must reset the stream, not end the body cleanly after {} bytes",
outcome.map(|b| b.len()).unwrap_or(0)
);
shutdown.cancel();
drive.abort();
}
}
+32 -4
View File
@@ -14,6 +14,7 @@ use tranquil_pds::scheduled::{
};
use tranquil_pds::state::AppState;
mod http3;
mod tls;
#[derive(Parser)]
@@ -259,7 +260,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
shutdown.clone(),
));
let app = tranquil_pds::app_with_routes(
let app = http3::with_host_from_authority(tranquil_pds::app_with_routes(
state,
tranquil_pds::ExternalRoutes {
xrpc: tranquil_api::api_routes().merge(tranquil_sync::sync_routes()),
@@ -270,7 +271,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.merge(tranquil_api::webhook_routes())
.merge(tranquil_oauth_server::frontend_client_metadata_route()),
},
);
));
let cfg = tranquil_config::get();
let host = &cfg.server.host;
@@ -286,6 +287,8 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.await
.map_err(|e| format!("Failed to bind to {}: {}", addr, e))?;
let mut http3_handle: Option<tokio::task::JoinHandle<()>> = None;
let server_handle = match cfg.server.tls.material() {
Some((cert_path, key_path)) => {
let initial = tls::load_certified_key(cert_path, key_path)
@@ -296,14 +299,35 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.map_err(|e| format!("Failed to build TLS configuration: {e}"))?,
);
tls::spawn_reload_handler(
resolver,
resolver.clone(),
cert_path.to_string(),
key_path.to_string(),
shutdown.clone(),
);
let tcp_app = if cfg.server.tls.http3 {
let quic_config = http3::build_quic_server_config(resolver)
.map_err(|e| format!("Failed to build HTTP/3 configuration: {e}"))?;
let endpoint = quinn::Endpoint::server(quic_config, addr)
.map_err(|e| format!("Failed to bind HTTP/3 endpoint on {addr}: {e}"))?;
let h3_port = endpoint
.local_addr()
.map(|a| a.port())
.map_err(|e| format!("Failed to read HTTP/3 local address: {e}"))?;
info!("HTTP/3 enabled on udp/{h3_port}");
http3_handle = Some(tokio::spawn(http3::serve_http3(
endpoint,
app.clone(),
shutdown.clone(),
)));
http3::with_alt_svc(app, h3_port)
} else {
app
};
info!("TLS termination enabled (h2, http/1.1), reload with SIGHUP");
let shutdown = shutdown.clone();
tokio::spawn(tls::serve_tls(listener, app, server_config, shutdown))
tokio::spawn(tls::serve_tls(listener, tcp_app, server_config, shutdown))
}
None => {
let make_service = app.into_make_service_with_connect_info::<SocketAddr>();
@@ -332,6 +356,10 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.await
.map_err(|e| format!("Server task panicked: {}", e))?;
if let Some(handle) = http3_handle {
handle.await.ok();
}
comms_handle.await.ok();
if let Some(handle) = crawlers_handle {
+3
View File
@@ -211,6 +211,9 @@ pub async fn serve_tls(
async move {
match accepted {
Ok((tcp, peer)) => {
if let Err(e) = tcp.set_nodelay(true) {
debug!("failed to set nodelay for {peer}: {e}");
}
let permit = tokio::select! {
biased;
_ = conn_shutdown.cancelled() => return,
+8
View File
@@ -128,6 +128,14 @@
# Can also be specified via environment variable `TLS_KEY_PATH`.
#key_path =
# Serve HTTP/3 over QUIC on the same UDP port as the TCP listener.
# Requires cert_path and key_path.
#
# Can also be specified via environment variable `TLS_HTTP3`.
#
# Default value: false
#http3 = false
[frontend]
# Whether to enable the built in serving of the frontend.
#