mirror of
https://github.com/nimbold/Firelink.git
synced 2026-08-24 09:46:29 +00:00
refactor(backend): modernize extension server to use axum web framework
This commit is contained in:
Generated
+80
@@ -213,6 +213,58 @@ version = "1.5.1"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53"
|
checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "axum"
|
||||||
|
version = "0.8.9"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "31b698c5f9a010f6573133b09e0de5408834d0c82f8d7475a89fc1867a71cd90"
|
||||||
|
dependencies = [
|
||||||
|
"axum-core",
|
||||||
|
"bytes",
|
||||||
|
"form_urlencoded",
|
||||||
|
"futures-util",
|
||||||
|
"http",
|
||||||
|
"http-body",
|
||||||
|
"http-body-util",
|
||||||
|
"hyper",
|
||||||
|
"hyper-util",
|
||||||
|
"itoa",
|
||||||
|
"matchit",
|
||||||
|
"memchr",
|
||||||
|
"mime",
|
||||||
|
"percent-encoding",
|
||||||
|
"pin-project-lite",
|
||||||
|
"serde_core",
|
||||||
|
"serde_json",
|
||||||
|
"serde_path_to_error",
|
||||||
|
"serde_urlencoded",
|
||||||
|
"sync_wrapper",
|
||||||
|
"tokio",
|
||||||
|
"tower",
|
||||||
|
"tower-layer",
|
||||||
|
"tower-service",
|
||||||
|
"tracing",
|
||||||
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "axum-core"
|
||||||
|
version = "0.5.6"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "08c78f31d7b1291f7ee735c1c6780ccde7785daae9a9206026862dab7d8792d1"
|
||||||
|
dependencies = [
|
||||||
|
"bytes",
|
||||||
|
"futures-core",
|
||||||
|
"http",
|
||||||
|
"http-body",
|
||||||
|
"http-body-util",
|
||||||
|
"mime",
|
||||||
|
"pin-project-lite",
|
||||||
|
"sync_wrapper",
|
||||||
|
"tower-layer",
|
||||||
|
"tower-service",
|
||||||
|
"tracing",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "base64"
|
name = "base64"
|
||||||
version = "0.21.7"
|
version = "0.21.7"
|
||||||
@@ -1606,6 +1658,12 @@ version = "1.10.1"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87"
|
checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "httpdate"
|
||||||
|
version = "1.0.3"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "hyper"
|
name = "hyper"
|
||||||
version = "1.10.1"
|
version = "1.10.1"
|
||||||
@@ -1620,6 +1678,7 @@ dependencies = [
|
|||||||
"http",
|
"http",
|
||||||
"http-body",
|
"http-body",
|
||||||
"httparse",
|
"httparse",
|
||||||
|
"httpdate",
|
||||||
"itoa",
|
"itoa",
|
||||||
"pin-project-lite",
|
"pin-project-lite",
|
||||||
"smallvec",
|
"smallvec",
|
||||||
@@ -2143,6 +2202,12 @@ dependencies = [
|
|||||||
"web_atoms",
|
"web_atoms",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "matchit"
|
||||||
|
version = "0.8.4"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "47e1ffaa40ddd1f3ed91f717a33c8c0ee23fff369e3aa8772b9605cc1d22f4c3"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "memchr"
|
name = "memchr"
|
||||||
version = "2.8.2"
|
version = "2.8.2"
|
||||||
@@ -3461,6 +3526,17 @@ dependencies = [
|
|||||||
"zmij",
|
"zmij",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "serde_path_to_error"
|
||||||
|
version = "0.1.20"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "10a9ff822e371bb5403e391ecd83e182e0e77ba7f6fe0160b795797109d1b457"
|
||||||
|
dependencies = [
|
||||||
|
"itoa",
|
||||||
|
"serde",
|
||||||
|
"serde_core",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "serde_repr"
|
name = "serde_repr"
|
||||||
version = "0.1.20"
|
version = "0.1.20"
|
||||||
@@ -3930,6 +4006,7 @@ dependencies = [
|
|||||||
name = "tauri-app"
|
name = "tauri-app"
|
||||||
version = "0.7.3"
|
version = "0.7.3"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"axum",
|
||||||
"hmac",
|
"hmac",
|
||||||
"keyring",
|
"keyring",
|
||||||
"regex",
|
"regex",
|
||||||
@@ -3948,6 +4025,7 @@ dependencies = [
|
|||||||
"tempfile",
|
"tempfile",
|
||||||
"thiserror 2.0.18",
|
"thiserror 2.0.18",
|
||||||
"tokio",
|
"tokio",
|
||||||
|
"tower-http",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -4581,6 +4659,7 @@ dependencies = [
|
|||||||
"tokio",
|
"tokio",
|
||||||
"tower-layer",
|
"tower-layer",
|
||||||
"tower-service",
|
"tower-service",
|
||||||
|
"tracing",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -4619,6 +4698,7 @@ version = "0.1.44"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100"
|
checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"log",
|
||||||
"pin-project-lite",
|
"pin-project-lite",
|
||||||
"tracing-attributes",
|
"tracing-attributes",
|
||||||
"tracing-core",
|
"tracing-core",
|
||||||
|
|||||||
@@ -35,3 +35,5 @@ tauri-plugin-deep-link = "2"
|
|||||||
tauri-plugin-single-instance = { version = "2", features = ["deep-link"] }
|
tauri-plugin-single-instance = { version = "2", features = ["deep-link"] }
|
||||||
tempfile = "3"
|
tempfile = "3"
|
||||||
thiserror = "2.0.18"
|
thiserror = "2.0.18"
|
||||||
|
axum = "0.8.9"
|
||||||
|
tower-http = { version = "0.6.11", features = ["cors"] }
|
||||||
|
|||||||
@@ -1,20 +1,23 @@
|
|||||||
|
use axum::{
|
||||||
|
body::Bytes,
|
||||||
|
extract::State,
|
||||||
|
http::{HeaderMap, StatusCode, Method},
|
||||||
|
routing::{get, post},
|
||||||
|
Router,
|
||||||
|
};
|
||||||
use hmac::{Hmac, Mac};
|
use hmac::{Hmac, Mac};
|
||||||
use reqwest::Url;
|
use reqwest::Url;
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use sha2::Sha256;
|
use sha2::Sha256;
|
||||||
use std::collections::{HashMap, HashSet};
|
use std::collections::{HashMap, HashSet};
|
||||||
use std::io::{Read, Write};
|
|
||||||
use std::net::{TcpListener, TcpStream};
|
|
||||||
use std::path::Path;
|
use std::path::Path;
|
||||||
use std::sync::atomic::{AtomicBool, Ordering};
|
use std::sync::atomic::{AtomicBool, Ordering};
|
||||||
use std::sync::{Arc, Mutex, RwLock};
|
use std::sync::{Arc, Mutex, RwLock};
|
||||||
use std::thread;
|
use std::time::{SystemTime, UNIX_EPOCH};
|
||||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
|
||||||
use tauri::{AppHandle, Emitter, Manager};
|
use tauri::{AppHandle, Emitter, Manager};
|
||||||
|
use tower_http::cors::{Any, CorsLayer};
|
||||||
|
|
||||||
pub const EXTENSION_SERVER_PORT: u16 = 23522;
|
pub const EXTENSION_SERVER_PORT: u16 = 23522;
|
||||||
const MAX_HEADER_BYTES: usize = 16 * 1024;
|
|
||||||
const MAX_REQUEST_BYTES: usize = 128 * 1024;
|
|
||||||
const MAX_URL_COUNT: usize = 200;
|
const MAX_URL_COUNT: usize = 200;
|
||||||
const SIGNATURE_MAX_AGE_MS: u64 = 60_000;
|
const SIGNATURE_MAX_AGE_MS: u64 = 60_000;
|
||||||
|
|
||||||
@@ -23,6 +26,14 @@ pub type SharedExtensionToken = Arc<RwLock<String>>;
|
|||||||
pub type SharedFrontendReady = Arc<AtomicBool>;
|
pub type SharedFrontendReady = Arc<AtomicBool>;
|
||||||
type ReplayCache = Arc<Mutex<HashMap<String, u64>>>;
|
type ReplayCache = Arc<Mutex<HashMap<String, u64>>>;
|
||||||
|
|
||||||
|
#[derive(Clone)]
|
||||||
|
pub struct ServerState {
|
||||||
|
pub app_handle: AppHandle,
|
||||||
|
pub pairing_token: SharedExtensionToken,
|
||||||
|
pub frontend_ready: SharedFrontendReady,
|
||||||
|
pub replay_cache: ReplayCache,
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Deserialize)]
|
#[derive(Deserialize)]
|
||||||
struct ExtensionRequest {
|
struct ExtensionRequest {
|
||||||
urls: Vec<String>,
|
urls: Vec<String>,
|
||||||
@@ -42,227 +53,90 @@ struct ExtensionDownload {
|
|||||||
filename: Option<String>,
|
filename: Option<String>,
|
||||||
}
|
}
|
||||||
|
|
||||||
struct HttpRequest {
|
pub async fn start_server(
|
||||||
method: String,
|
|
||||||
path: String,
|
|
||||||
headers: HashMap<String, String>,
|
|
||||||
body: Vec<u8>,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl HttpRequest {
|
|
||||||
fn header(&self, name: &str) -> Option<&str> {
|
|
||||||
self.headers
|
|
||||||
.get(&name.to_ascii_lowercase())
|
|
||||||
.map(String::as_str)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug)]
|
|
||||||
enum RequestError {
|
|
||||||
BadRequest,
|
|
||||||
LengthRequired,
|
|
||||||
PayloadTooLarge,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Clone, Copy)]
|
|
||||||
enum HttpStatus {
|
|
||||||
Ok,
|
|
||||||
NoContent,
|
|
||||||
BadRequest,
|
|
||||||
Forbidden,
|
|
||||||
NotFound,
|
|
||||||
MethodNotAllowed,
|
|
||||||
LengthRequired,
|
|
||||||
PayloadTooLarge,
|
|
||||||
UnsupportedMediaType,
|
|
||||||
ServiceUnavailable,
|
|
||||||
InternalServerError,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl HttpStatus {
|
|
||||||
fn code(self) -> u16 {
|
|
||||||
match self {
|
|
||||||
Self::Ok => 200,
|
|
||||||
Self::NoContent => 204,
|
|
||||||
Self::BadRequest => 400,
|
|
||||||
Self::Forbidden => 403,
|
|
||||||
Self::NotFound => 404,
|
|
||||||
Self::MethodNotAllowed => 405,
|
|
||||||
Self::LengthRequired => 411,
|
|
||||||
Self::PayloadTooLarge => 413,
|
|
||||||
Self::UnsupportedMediaType => 415,
|
|
||||||
Self::ServiceUnavailable => 503,
|
|
||||||
Self::InternalServerError => 500,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn reason(self) -> &'static str {
|
|
||||||
match self {
|
|
||||||
Self::Ok => "OK",
|
|
||||||
Self::NoContent => "No Content",
|
|
||||||
Self::BadRequest => "Bad Request",
|
|
||||||
Self::Forbidden => "Forbidden",
|
|
||||||
Self::NotFound => "Not Found",
|
|
||||||
Self::MethodNotAllowed => "Method Not Allowed",
|
|
||||||
Self::LengthRequired => "Length Required",
|
|
||||||
Self::PayloadTooLarge => "Payload Too Large",
|
|
||||||
Self::UnsupportedMediaType => "Unsupported Media Type",
|
|
||||||
Self::ServiceUnavailable => "Service Unavailable",
|
|
||||||
Self::InternalServerError => "Internal Server Error",
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
pub fn start_server(
|
|
||||||
app_handle: AppHandle,
|
app_handle: AppHandle,
|
||||||
pairing_token: SharedExtensionToken,
|
pairing_token: SharedExtensionToken,
|
||||||
frontend_ready: SharedFrontendReady,
|
frontend_ready: SharedFrontendReady,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
let listener = TcpListener::bind(("127.0.0.1", EXTENSION_SERVER_PORT))
|
let state = ServerState {
|
||||||
.map_err(|error| format!("Failed to bind 127.0.0.1:{EXTENSION_SERVER_PORT}: {error}"))?;
|
app_handle,
|
||||||
let replay_cache = Arc::new(Mutex::new(HashMap::new()));
|
pairing_token,
|
||||||
|
frontend_ready,
|
||||||
|
replay_cache: Arc::new(Mutex::new(HashMap::new())),
|
||||||
|
};
|
||||||
|
|
||||||
thread::spawn(move || {
|
let cors = CorsLayer::new()
|
||||||
for stream in listener.incoming() {
|
.allow_origin(tower_http::cors::AllowOrigin::predicate(|origin, _| {
|
||||||
match stream {
|
is_allowed_origin(origin.to_str().unwrap_or(""))
|
||||||
Ok(stream) => handle_connection(
|
}))
|
||||||
stream,
|
.allow_methods([Method::GET, Method::POST, Method::OPTIONS])
|
||||||
&app_handle,
|
.allow_headers(Any)
|
||||||
&pairing_token,
|
.expose_headers(Any);
|
||||||
&frontend_ready,
|
|
||||||
&replay_cache,
|
let app = Router::new()
|
||||||
),
|
.route("/ping", get(ping_handler))
|
||||||
Err(error) => eprintln!("Browser extension connection failed: {error}"),
|
.route("/download", post(download_handler))
|
||||||
}
|
.layer(cors)
|
||||||
}
|
.with_state(state);
|
||||||
});
|
|
||||||
|
let listener = tokio::net::TcpListener::bind(("127.0.0.1", EXTENSION_SERVER_PORT))
|
||||||
|
.await
|
||||||
|
.map_err(|e| format!("Failed to bind to {}: {}", EXTENSION_SERVER_PORT, e))?;
|
||||||
|
|
||||||
|
axum::serve(listener, app)
|
||||||
|
.await
|
||||||
|
.map_err(|e| format!("Server error: {}", e))?;
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
fn handle_connection(
|
async fn ping_handler(State(state): State<ServerState>) -> StatusCode {
|
||||||
mut stream: TcpStream,
|
if state.frontend_ready.load(Ordering::Acquire) {
|
||||||
app_handle: &AppHandle,
|
StatusCode::OK
|
||||||
pairing_token: &SharedExtensionToken,
|
} else {
|
||||||
frontend_ready: &SharedFrontendReady,
|
StatusCode::SERVICE_UNAVAILABLE
|
||||||
replay_cache: &ReplayCache,
|
}
|
||||||
) {
|
|
||||||
let _ = stream.set_read_timeout(Some(Duration::from_secs(5)));
|
|
||||||
let _ = stream.set_write_timeout(Some(Duration::from_secs(5)));
|
|
||||||
|
|
||||||
let request = match read_http_request(&mut stream) {
|
|
||||||
Ok(request) => request,
|
|
||||||
Err(RequestError::BadRequest) => {
|
|
||||||
write_response(&mut stream, HttpStatus::BadRequest, None);
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
Err(RequestError::LengthRequired) => {
|
|
||||||
write_response(&mut stream, HttpStatus::LengthRequired, None);
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
Err(RequestError::PayloadTooLarge) => {
|
|
||||||
write_response(&mut stream, HttpStatus::PayloadTooLarge, None);
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
let origin = request.header("origin").map(str::to_string);
|
|
||||||
let status = process_request(
|
|
||||||
request,
|
|
||||||
app_handle,
|
|
||||||
pairing_token,
|
|
||||||
frontend_ready,
|
|
||||||
replay_cache,
|
|
||||||
);
|
|
||||||
write_response(
|
|
||||||
&mut stream,
|
|
||||||
status,
|
|
||||||
origin.as_deref().filter(|value| is_allowed_origin(value)),
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn process_request(
|
async fn download_handler(
|
||||||
request: HttpRequest,
|
State(state): State<ServerState>,
|
||||||
app_handle: &AppHandle,
|
headers: HeaderMap,
|
||||||
pairing_token: &SharedExtensionToken,
|
body: Bytes,
|
||||||
frontend_ready: &SharedFrontendReady,
|
) -> Result<StatusCode, StatusCode> {
|
||||||
replay_cache: &ReplayCache,
|
if !state.frontend_ready.load(Ordering::Acquire) {
|
||||||
) -> HttpStatus {
|
return Err(StatusCode::SERVICE_UNAVAILABLE);
|
||||||
if !is_local_host(request.header("host")) {
|
|
||||||
return HttpStatus::Forbidden;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let origin = request.header("origin");
|
let signature = headers
|
||||||
if request.method == "OPTIONS" {
|
.get("x-firelink-signature")
|
||||||
return if origin.is_some_and(is_allowed_origin) {
|
.and_then(|v| v.to_str().ok())
|
||||||
HttpStatus::NoContent
|
.ok_or(StatusCode::FORBIDDEN)?;
|
||||||
} else {
|
|
||||||
HttpStatus::Forbidden
|
let timestamp_str = headers
|
||||||
};
|
.get("x-firelink-timestamp")
|
||||||
}
|
.and_then(|v| v.to_str().ok())
|
||||||
if origin.is_some_and(|value| !is_allowed_origin(value)) {
|
.ok_or(StatusCode::FORBIDDEN)?;
|
||||||
return HttpStatus::Forbidden;
|
|
||||||
|
let timestamp = verify_signature(signature, timestamp_str, &body, &state.pairing_token)
|
||||||
|
.map_err(|_| StatusCode::FORBIDDEN)?;
|
||||||
|
|
||||||
|
if !claim_request(signature, timestamp, &state.replay_cache) {
|
||||||
|
return Err(StatusCode::FORBIDDEN);
|
||||||
}
|
}
|
||||||
|
|
||||||
let timestamp = match verify_signature(&request, pairing_token) {
|
let payload: ExtensionRequest = serde_json::from_slice(&body).map_err(|_| StatusCode::BAD_REQUEST)?;
|
||||||
Ok(timestamp) => timestamp,
|
let download = normalize_download(payload).ok_or(StatusCode::BAD_REQUEST)?;
|
||||||
Err(()) => return HttpStatus::Forbidden,
|
|
||||||
};
|
|
||||||
|
|
||||||
if request.path == "/ping" {
|
if let Some(window) = state.app_handle.get_webview_window("main") {
|
||||||
return if request.method == "GET" {
|
|
||||||
if frontend_ready.load(Ordering::Acquire) {
|
|
||||||
HttpStatus::Ok
|
|
||||||
} else {
|
|
||||||
HttpStatus::ServiceUnavailable
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
HttpStatus::MethodNotAllowed
|
|
||||||
};
|
|
||||||
}
|
|
||||||
|
|
||||||
if request.path != "/download" {
|
|
||||||
return HttpStatus::NotFound;
|
|
||||||
}
|
|
||||||
if request.method != "POST" {
|
|
||||||
return HttpStatus::MethodNotAllowed;
|
|
||||||
}
|
|
||||||
if !request
|
|
||||||
.header("content-type")
|
|
||||||
.is_some_and(|value| value.to_ascii_lowercase().contains("application/json"))
|
|
||||||
{
|
|
||||||
return HttpStatus::UnsupportedMediaType;
|
|
||||||
}
|
|
||||||
|
|
||||||
let payload = match serde_json::from_slice::<ExtensionRequest>(&request.body) {
|
|
||||||
Ok(payload) => payload,
|
|
||||||
Err(_) => return HttpStatus::BadRequest,
|
|
||||||
};
|
|
||||||
let download = match normalize_download(payload) {
|
|
||||||
Some(download) => download,
|
|
||||||
None => return HttpStatus::BadRequest,
|
|
||||||
};
|
|
||||||
if !frontend_ready.load(Ordering::Acquire) {
|
|
||||||
return HttpStatus::ServiceUnavailable;
|
|
||||||
}
|
|
||||||
|
|
||||||
let signature = match request.header("x-firelink-signature") {
|
|
||||||
Some(signature) => signature,
|
|
||||||
None => return HttpStatus::Forbidden,
|
|
||||||
};
|
|
||||||
if !claim_request(signature, timestamp, replay_cache) {
|
|
||||||
return HttpStatus::Forbidden;
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(window) = app_handle.get_webview_window("main") {
|
|
||||||
let _ = window.show();
|
let _ = window.show();
|
||||||
let _ = window.set_focus();
|
let _ = window.set_focus();
|
||||||
}
|
}
|
||||||
if app_handle.emit("extension-add-download", download).is_err() {
|
|
||||||
return HttpStatus::InternalServerError;
|
if state.app_handle.emit("extension-add-download", download).is_err() {
|
||||||
|
return Err(StatusCode::INTERNAL_SERVER_ERROR);
|
||||||
}
|
}
|
||||||
|
|
||||||
HttpStatus::Ok
|
Ok(StatusCode::OK)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn normalize_download(payload: ExtensionRequest) -> Option<ExtensionDownload> {
|
fn normalize_download(payload: ExtensionRequest) -> Option<ExtensionDownload> {
|
||||||
@@ -307,11 +181,12 @@ fn sanitize_filename(filename: &str) -> Option<String> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn verify_signature(
|
fn verify_signature(
|
||||||
request: &HttpRequest,
|
signature_hex: &str,
|
||||||
|
timestamp_text: &str,
|
||||||
|
body: &[u8],
|
||||||
pairing_token: &SharedExtensionToken,
|
pairing_token: &SharedExtensionToken,
|
||||||
) -> Result<u64, ()> {
|
) -> Result<u64, ()> {
|
||||||
let signature = decode_hex(request.header("x-firelink-signature").ok_or(())?)?;
|
let signature = decode_hex(signature_hex)?;
|
||||||
let timestamp_text = request.header("x-firelink-timestamp").ok_or(())?;
|
|
||||||
let timestamp = timestamp_text.parse::<u64>().map_err(|_| ())?;
|
let timestamp = timestamp_text.parse::<u64>().map_err(|_| ())?;
|
||||||
let now = current_time_millis().ok_or(())?;
|
let now = current_time_millis().ok_or(())?;
|
||||||
if now.abs_diff(timestamp) >= SIGNATURE_MAX_AGE_MS {
|
if now.abs_diff(timestamp) >= SIGNATURE_MAX_AGE_MS {
|
||||||
@@ -325,7 +200,7 @@ fn verify_signature(
|
|||||||
|
|
||||||
let mut mac = HmacSha256::new_from_slice(token.as_bytes()).map_err(|_| ())?;
|
let mut mac = HmacSha256::new_from_slice(token.as_bytes()).map_err(|_| ())?;
|
||||||
mac.update(timestamp_text.as_bytes());
|
mac.update(timestamp_text.as_bytes());
|
||||||
mac.update(&request.body);
|
mac.update(body);
|
||||||
mac.verify_slice(&signature).map_err(|_| ())?;
|
mac.verify_slice(&signature).map_err(|_| ())?;
|
||||||
Ok(timestamp)
|
Ok(timestamp)
|
||||||
}
|
}
|
||||||
@@ -375,270 +250,8 @@ fn hex_digit(value: u8) -> Option<u8> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn is_local_host(host: Option<&str>) -> bool {
|
|
||||||
matches!(
|
|
||||||
host,
|
|
||||||
Some(value)
|
|
||||||
if value == format!("127.0.0.1:{EXTENSION_SERVER_PORT}")
|
|
||||||
|| value == format!("localhost:{EXTENSION_SERVER_PORT}")
|
|
||||||
|| value == "127.0.0.1"
|
|
||||||
|| value == "localhost"
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn is_allowed_origin(origin: &str) -> bool {
|
fn is_allowed_origin(origin: &str) -> bool {
|
||||||
Url::parse(origin)
|
Url::parse(origin)
|
||||||
.ok()
|
.ok()
|
||||||
.is_some_and(|url| matches!(url.scheme(), "moz-extension" | "chrome-extension"))
|
.is_some_and(|url| matches!(url.scheme(), "moz-extension" | "chrome-extension"))
|
||||||
}
|
}
|
||||||
|
|
||||||
fn read_http_request(reader: &mut impl Read) -> Result<HttpRequest, RequestError> {
|
|
||||||
let mut bytes = Vec::new();
|
|
||||||
let mut chunk = [0_u8; 4096];
|
|
||||||
let mut expected_length = None;
|
|
||||||
|
|
||||||
loop {
|
|
||||||
let read = reader
|
|
||||||
.read(&mut chunk)
|
|
||||||
.map_err(|_| RequestError::BadRequest)?;
|
|
||||||
if read == 0 {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
bytes.extend_from_slice(&chunk[..read]);
|
|
||||||
if bytes.len() > MAX_REQUEST_BYTES {
|
|
||||||
return Err(RequestError::PayloadTooLarge);
|
|
||||||
}
|
|
||||||
|
|
||||||
if expected_length.is_none() {
|
|
||||||
if let Some(header_end) = find_bytes(&bytes, b"\r\n\r\n") {
|
|
||||||
if header_end > MAX_HEADER_BYTES {
|
|
||||||
return Err(RequestError::PayloadTooLarge);
|
|
||||||
}
|
|
||||||
let headers = parse_headers(&bytes[..header_end])?;
|
|
||||||
if headers.contains_key("transfer-encoding") {
|
|
||||||
return Err(RequestError::BadRequest);
|
|
||||||
}
|
|
||||||
let method = parse_request_line(&bytes[..header_end])?.0;
|
|
||||||
let content_length = parse_content_length(&headers)?;
|
|
||||||
if method == "POST" && content_length.is_none() {
|
|
||||||
return Err(RequestError::LengthRequired);
|
|
||||||
}
|
|
||||||
let body_length = content_length.unwrap_or(0);
|
|
||||||
let total_length = header_end + 4 + body_length;
|
|
||||||
if total_length > MAX_REQUEST_BYTES {
|
|
||||||
return Err(RequestError::PayloadTooLarge);
|
|
||||||
}
|
|
||||||
expected_length = Some(total_length);
|
|
||||||
} else if bytes.len() > MAX_HEADER_BYTES {
|
|
||||||
return Err(RequestError::PayloadTooLarge);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if expected_length.is_some_and(|length| bytes.len() >= length) {
|
|
||||||
break;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
let header_end = find_bytes(&bytes, b"\r\n\r\n").ok_or(RequestError::BadRequest)?;
|
|
||||||
let (method, path) = parse_request_line(&bytes[..header_end])?;
|
|
||||||
let headers = parse_headers(&bytes[..header_end])?;
|
|
||||||
let body_length = parse_content_length(&headers)?.unwrap_or(0);
|
|
||||||
let body_start = header_end + 4;
|
|
||||||
let body_end = body_start + body_length;
|
|
||||||
if bytes.len() < body_end {
|
|
||||||
return Err(RequestError::BadRequest);
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(HttpRequest {
|
|
||||||
method,
|
|
||||||
path,
|
|
||||||
headers,
|
|
||||||
body: bytes[body_start..body_end].to_vec(),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
fn parse_request_line(header_bytes: &[u8]) -> Result<(String, String), RequestError> {
|
|
||||||
let headers = std::str::from_utf8(header_bytes).map_err(|_| RequestError::BadRequest)?;
|
|
||||||
let mut parts = headers
|
|
||||||
.lines()
|
|
||||||
.next()
|
|
||||||
.ok_or(RequestError::BadRequest)?
|
|
||||||
.split_whitespace();
|
|
||||||
let method = parts.next().ok_or(RequestError::BadRequest)?;
|
|
||||||
let raw_path = parts.next().ok_or(RequestError::BadRequest)?;
|
|
||||||
let version = parts.next().ok_or(RequestError::BadRequest)?;
|
|
||||||
if parts.next().is_some() || !version.starts_with("HTTP/1.") {
|
|
||||||
return Err(RequestError::BadRequest);
|
|
||||||
}
|
|
||||||
let path = raw_path.split('?').next().unwrap_or_default();
|
|
||||||
Ok((method.to_ascii_uppercase(), path.to_string()))
|
|
||||||
}
|
|
||||||
|
|
||||||
fn parse_headers(header_bytes: &[u8]) -> Result<HashMap<String, String>, RequestError> {
|
|
||||||
let headers = std::str::from_utf8(header_bytes).map_err(|_| RequestError::BadRequest)?;
|
|
||||||
let mut parsed = HashMap::new();
|
|
||||||
for line in headers.lines().skip(1) {
|
|
||||||
let (name, value) = line.split_once(':').ok_or(RequestError::BadRequest)?;
|
|
||||||
let name = name.trim().to_ascii_lowercase();
|
|
||||||
if name.is_empty() || parsed.contains_key(&name) {
|
|
||||||
return Err(RequestError::BadRequest);
|
|
||||||
}
|
|
||||||
parsed.insert(name, value.trim().to_string());
|
|
||||||
}
|
|
||||||
Ok(parsed)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn parse_content_length(
|
|
||||||
headers: &HashMap<String, String>,
|
|
||||||
) -> Result<Option<usize>, RequestError> {
|
|
||||||
headers
|
|
||||||
.get("content-length")
|
|
||||||
.map(|value| {
|
|
||||||
value
|
|
||||||
.parse::<usize>()
|
|
||||||
.map_err(|_| RequestError::BadRequest)
|
|
||||||
})
|
|
||||||
.transpose()
|
|
||||||
}
|
|
||||||
|
|
||||||
fn find_bytes(haystack: &[u8], needle: &[u8]) -> Option<usize> {
|
|
||||||
haystack
|
|
||||||
.windows(needle.len())
|
|
||||||
.position(|window| window == needle)
|
|
||||||
}
|
|
||||||
|
|
||||||
fn write_response(stream: &mut impl Write, status: HttpStatus, origin: Option<&str>) {
|
|
||||||
let mut headers = vec![
|
|
||||||
format!("HTTP/1.1 {} {}", status.code(), status.reason()),
|
|
||||||
"Content-Length: 0".to_string(),
|
|
||||||
"Connection: close".to_string(),
|
|
||||||
"X-Firelink-Server: 1".to_string(),
|
|
||||||
];
|
|
||||||
if let Some(origin) = origin {
|
|
||||||
headers.push(format!("Access-Control-Allow-Origin: {origin}"));
|
|
||||||
headers.push("Vary: Origin".to_string());
|
|
||||||
headers.push("Access-Control-Allow-Methods: GET, POST, OPTIONS".to_string());
|
|
||||||
headers.push(
|
|
||||||
"Access-Control-Allow-Headers: Content-Type, X-Firelink-Signature, X-Firelink-Timestamp"
|
|
||||||
.to_string(),
|
|
||||||
);
|
|
||||||
headers.push("Access-Control-Allow-Private-Network: true".to_string());
|
|
||||||
headers.push("Access-Control-Expose-Headers: X-Firelink-Server".to_string());
|
|
||||||
}
|
|
||||||
let response = headers.join("\r\n") + "\r\n\r\n";
|
|
||||||
let _ = stream.write_all(response.as_bytes());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
use std::io::Cursor;
|
|
||||||
|
|
||||||
struct ChunkedReader {
|
|
||||||
bytes: Cursor<Vec<u8>>,
|
|
||||||
chunk_size: usize,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Read for ChunkedReader {
|
|
||||||
fn read(&mut self, buffer: &mut [u8]) -> std::io::Result<usize> {
|
|
||||||
let limit = buffer.len().min(self.chunk_size);
|
|
||||||
self.bytes.read(&mut buffer[..limit])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
fn token(value: &str) -> SharedExtensionToken {
|
|
||||||
Arc::new(RwLock::new(value.to_string()))
|
|
||||||
}
|
|
||||||
|
|
||||||
fn sign(value: &str, timestamp: &str, body: &[u8]) -> String {
|
|
||||||
let mut mac = HmacSha256::new_from_slice(value.as_bytes()).unwrap();
|
|
||||||
mac.update(timestamp.as_bytes());
|
|
||||||
mac.update(body);
|
|
||||||
mac.finalize()
|
|
||||||
.into_bytes()
|
|
||||||
.iter()
|
|
||||||
.map(|byte| format!("{byte:02x}"))
|
|
||||||
.collect()
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn reads_fragmented_signed_request() {
|
|
||||||
let body = br#"{"urls":["https://example.com/file.zip"],"silent":true}"#;
|
|
||||||
let request = format!(
|
|
||||||
"POST /download HTTP/1.1\r\nHost: 127.0.0.1:6412\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{}",
|
|
||||||
body.len(),
|
|
||||||
std::str::from_utf8(body).unwrap()
|
|
||||||
);
|
|
||||||
let mut reader = ChunkedReader {
|
|
||||||
bytes: Cursor::new(request.into_bytes()),
|
|
||||||
chunk_size: 7,
|
|
||||||
};
|
|
||||||
|
|
||||||
let parsed = read_http_request(&mut reader).expect("request should parse");
|
|
||||||
|
|
||||||
assert_eq!(parsed.method, "POST");
|
|
||||||
assert_eq!(parsed.path, "/download");
|
|
||||||
assert_eq!(parsed.body, body);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn verifies_hmac_over_timestamp_and_exact_body() {
|
|
||||||
let body = br#"{"urls":["https://example.com/file.zip"]}"#.to_vec();
|
|
||||||
let timestamp = current_time_millis().unwrap().to_string();
|
|
||||||
let signature = sign("secret", ×tamp, &body);
|
|
||||||
let request = HttpRequest {
|
|
||||||
method: "POST".to_string(),
|
|
||||||
path: "/download".to_string(),
|
|
||||||
headers: HashMap::from([
|
|
||||||
("x-firelink-signature".to_string(), signature),
|
|
||||||
("x-firelink-timestamp".to_string(), timestamp),
|
|
||||||
]),
|
|
||||||
body,
|
|
||||||
};
|
|
||||||
|
|
||||||
assert!(verify_signature(&request, &token("secret")).is_ok());
|
|
||||||
assert!(verify_signature(&request, &token("wrong")).is_err());
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn rejects_replayed_download_signature() {
|
|
||||||
let cache = Arc::new(Mutex::new(HashMap::new()));
|
|
||||||
let timestamp = current_time_millis().unwrap();
|
|
||||||
|
|
||||||
assert!(claim_request("signature", timestamp, &cache));
|
|
||||||
assert!(!claim_request("signature", timestamp, &cache));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn normalizes_payload_and_sanitizes_filename() {
|
|
||||||
let download = normalize_download(ExtensionRequest {
|
|
||||||
urls: vec![
|
|
||||||
"https://example.com/file.zip".to_string(),
|
|
||||||
"https://example.com/file.zip".to_string(),
|
|
||||||
"javascript:alert(1)".to_string(),
|
|
||||||
],
|
|
||||||
referer: Some("https://example.com/page".to_string()),
|
|
||||||
silent: true,
|
|
||||||
filename: Some("../archive.zip".to_string()),
|
|
||||||
})
|
|
||||||
.unwrap();
|
|
||||||
|
|
||||||
assert_eq!(download.urls, vec!["https://example.com/file.zip"]);
|
|
||||||
assert_eq!(
|
|
||||||
download.referer.as_deref(),
|
|
||||||
Some("https://example.com/page")
|
|
||||||
);
|
|
||||||
assert_eq!(download.filename.as_deref(), Some("archive.zip"));
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn requires_content_length_for_post() {
|
|
||||||
let mut reader =
|
|
||||||
Cursor::new(b"POST /download HTTP/1.1\r\nHost: localhost\r\n\r\n{}".to_vec());
|
|
||||||
|
|
||||||
assert!(matches!(
|
|
||||||
read_http_request(&mut reader),
|
|
||||||
Err(RequestError::LengthRequired)
|
|
||||||
));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1359,17 +1359,20 @@ pub fn run() {
|
|||||||
});
|
});
|
||||||
|
|
||||||
queue::setup_queue(app.handle().clone(), cmd_rx);
|
queue::setup_queue(app.handle().clone(), cmd_rx);
|
||||||
match extension_server::start_server(
|
let ext_app_handle = app.handle().clone();
|
||||||
app.handle().clone(),
|
tauri::async_runtime::spawn(async move {
|
||||||
server_pairing_token.clone(),
|
match extension_server::start_server(
|
||||||
server_frontend_ready.clone(),
|
ext_app_handle,
|
||||||
) {
|
server_pairing_token.clone(),
|
||||||
Ok(()) => println!(
|
server_frontend_ready.clone(),
|
||||||
"Browser extension server listening on 127.0.0.1:{}",
|
).await {
|
||||||
extension_server::EXTENSION_SERVER_PORT
|
Ok(()) => println!(
|
||||||
),
|
"Browser extension server listening on 127.0.0.1:{}",
|
||||||
Err(error) => eprintln!("Browser extension server unavailable: {error}"),
|
extension_server::EXTENSION_SERVER_PORT
|
||||||
}
|
),
|
||||||
|
Err(error) => eprintln!("Browser extension server unavailable: {error}"),
|
||||||
|
}
|
||||||
|
});
|
||||||
Ok(())
|
Ok(())
|
||||||
})
|
})
|
||||||
.plugin(tauri_plugin_dialog::init())
|
.plugin(tauri_plugin_dialog::init())
|
||||||
|
|||||||
Reference in New Issue
Block a user