diff --git a/Extensions/Firefox b/Extensions/Firefox index 3301425..bde440f 160000 --- a/Extensions/Firefox +++ b/Extensions/Firefox @@ -1 +1 @@ -Subproject commit 330142590bb464708ec074a44929fa03b1fc7b89 +Subproject commit bde440f5e11620a5155a434b4263aca55504493d diff --git a/src-tauri/src/extension_server.rs b/src-tauri/src/extension_server.rs index 1cfa02a..c8e59d6 100644 --- a/src-tauri/src/extension_server.rs +++ b/src-tauri/src/extension_server.rs @@ -1,7 +1,9 @@ use axum::{ - body::Bytes, + body::{Body, Bytes}, extract::State, - http::{HeaderMap, Method, StatusCode}, + http::{HeaderMap, HeaderValue, Method, Request, StatusCode}, + middleware::{self, Next}, + response::Response, routing::{get, post}, Router, }; @@ -23,6 +25,7 @@ pub const EXTENSION_SERVER_PORT: u16 = 6412; pub const EXTENSION_SERVER_PORT_RANGE: std::ops::RangeInclusive = EXTENSION_SERVER_PORT..=6422; const MAX_URL_COUNT: usize = 200; const SIGNATURE_MAX_AGE_MS: u64 = 60_000; +const SERVER_HEADER: &str = "x-firelink-server"; type HmacSha256 = Hmac; pub type SharedExtensionToken = Arc>; @@ -90,6 +93,7 @@ pub async fn start_server( .route("/ping", get(ping_handler)) .route("/download", post(download_handler)) .layer(cors) + .layer(middleware::from_fn(add_server_identity)) .with_state(state); let (port, listener) = bind_extension_listener().await?; @@ -116,6 +120,14 @@ pub async fn start_server( server_result } +async fn add_server_identity(request: Request, next: Next) -> Response { + let mut response = next.run(request).await; + response + .headers_mut() + .insert(SERVER_HEADER, HeaderValue::from_static("1")); + response +} + async fn bind_extension_listener() -> Result<(u16, tokio::net::TcpListener), String> { let mut errors = Vec::new(); for port in EXTENSION_SERVER_PORT_RANGE { @@ -359,3 +371,29 @@ fn is_allowed_origin(origin: &str) -> bool { .ok() .is_some_and(|url| matches!(url.scheme(), "moz-extension" | "chrome-extension")) } + +#[cfg(test)] +mod tests { + use super::{add_server_identity, SERVER_HEADER}; + use axum::{http::StatusCode, middleware, routing::get, Router}; + + #[tokio::test] + async fn identifies_every_extension_server_response() { + let app = Router::new() + .route("/ping", get(|| async { StatusCode::FORBIDDEN })) + .layer(middleware::from_fn(add_server_identity)); + let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + let response = reqwest::get(format!("http://{address}/ping")).await.unwrap(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); + assert_eq!(response.headers().get(SERVER_HEADER).unwrap(), "1"); + + server.abort(); + } +} diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 872f346..0bdd60b 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -3676,6 +3676,16 @@ pub fn run() { let database = crate::db::init(app.handle()) .map_err(|error| format!("failed to initialize persistence: {error}"))?; + let initial_pairing_token = { + let mut connection = database.lock()?; + crate::db::hydrate_pairing_token(&mut connection)?.0 + }; + { + let mut pairing_token = extension_pairing_token + .write() + .map_err(|_| "extension pairing token lock is unavailable".to_string())?; + *pairing_token = initial_pairing_token; + } app.manage(database); let max_concurrent = {