diff --git a/desktop/src-tauri/Cargo.lock b/desktop/src-tauri/Cargo.lock index 95bc1a674a6..bd5efca8a21 100644 --- a/desktop/src-tauri/Cargo.lock +++ b/desktop/src-tauri/Cargo.lock @@ -798,6 +798,7 @@ checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ "block-buffer", "crypto-common", + "subtle", ] [[package]] @@ -1583,6 +1584,15 @@ version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" +[[package]] +name = "hmac" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" +dependencies = [ + "digest", +] + [[package]] name = "html5ever" version = "0.38.0" @@ -2543,9 +2553,12 @@ dependencies = [ name = "opencodex-desktop" version = "2.61.0" dependencies = [ + "base64 0.22.1", + "hmac", "reqwest 0.12.24", "serde", "serde_json", + "sha2", "tauri", "tauri-build", "tauri-plugin-autostart", diff --git a/desktop/src-tauri/Cargo.toml b/desktop/src-tauri/Cargo.toml index 2548b000cd0..43c6c5a46f8 100644 --- a/desktop/src-tauri/Cargo.toml +++ b/desktop/src-tauri/Cargo.toml @@ -15,9 +15,12 @@ crate-type = ["staticlib", "cdylib", "rlib"] tauri-build = { version = "=2.6.3", features = [] } [dependencies] +base64 = "=0.22.1" +hmac = "=0.12.1" reqwest = { version = "=0.12.24", default-features = false, features = ["json", "rustls-tls"] } serde = { version = "=1.0.219", features = ["derive"] } serde_json = "=1.0.140" +sha2 = "=0.10.9" uuid = { version = "=1.18.1", features = ["v4"] } tauri = { version = "=2.11.6", features = ["tray-icon", "image-png"] } tauri-plugin-autostart = "=2.5.0" diff --git a/desktop/src-tauri/src/auth.rs b/desktop/src-tauri/src/auth.rs index 81a16467b05..236b5da5592 100644 --- a/desktop/src-tauri/src/auth.rs +++ b/desktop/src-tauri/src/auth.rs @@ -1,5 +1,15 @@ use std::path::PathBuf; +use serde::Deserialize; + +#[derive(Clone, Debug, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "camelCase")] +pub struct RuntimeIdentity { + pub pid: u32, + pub port: u16, + pub attestation_secret: String, +} + #[derive(Clone, Debug)] pub struct Auth { home: PathBuf, @@ -25,6 +35,15 @@ impl Auth { }) } + pub fn runtime_identity(&self) -> Option { + let value = std::fs::read(self.home.join("runtime-port.json")).ok()?; + let identity: RuntimeIdentity = serde_json::from_slice(&value).ok()?; + if identity.pid == 0 || identity.attestation_secret.len() != 43 { + return None; + } + Some(identity) + } + pub fn user_agent() -> &'static str { concat!("OpenCodexDesktop/", env!("CARGO_PKG_VERSION")) } diff --git a/desktop/src-tauri/src/proxy.rs b/desktop/src-tauri/src/proxy.rs index fc8dc71ec5e..a49c8abd42f 100644 --- a/desktop/src-tauri/src/proxy.rs +++ b/desktop/src-tauri/src/proxy.rs @@ -1,6 +1,12 @@ -use crate::{auth::Auth, discovery::ProxyEndpoint}; +use crate::{ + auth::{Auth, RuntimeIdentity}, + discovery::ProxyEndpoint, +}; +use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine}; +use hmac::{Hmac, Mac}; use reqwest::{Client, Method, StatusCode}; use serde_json::Value; +use sha2::Sha256; use std::time::Duration; #[derive(Clone)] @@ -73,6 +79,7 @@ impl ProxyClient { async fn request(&self, method: Method, path: &str) -> Result { let response = self.send(&method, path, None).await?; if response.status() == StatusCode::UNAUTHORIZED { + self.authenticate_target().await?; let token = self.auth.token().ok_or(ProxyError::Unauthorized)?; let response = self.send(&method, path, Some(token)).await?; return decode(response).await; @@ -80,6 +87,42 @@ impl ProxyClient { decode(response).await } + async fn authenticate_target(&self) -> Result<(), ProxyError> { + let identity = self + .auth + .runtime_identity() + .ok_or(ProxyError::Unauthorized)?; + if identity.port != self.endpoint.port { + return Err(ProxyError::Unauthorized); + } + let mut challenge_bytes = [0_u8; 32]; + challenge_bytes[..16].copy_from_slice(uuid::Uuid::new_v4().as_bytes()); + challenge_bytes[16..].copy_from_slice(uuid::Uuid::new_v4().as_bytes()); + let challenge = URL_SAFE_NO_PAD.encode(challenge_bytes); + let response = self + .client + .get(self.endpoint.url("/healthz")) + .header("x-opencodex-attestation-challenge", &challenge) + .send() + .await + .map_err(map_request_error)?; + let proof = response + .headers() + .get("x-opencodex-attestation-proof") + .and_then(|value| value.to_str().ok()) + .map(str::to_owned); + let health: Value = decode(response).await?; + if health.get("service").and_then(Value::as_str) != Some("opencodex") + || health.get("pid").and_then(Value::as_u64) != Some(identity.pid.into()) + || health.get("port").and_then(Value::as_u64) != Some(identity.port.into()) + || !valid_attestation_proof(&identity, &challenge, proof.as_deref()) + || self.auth.runtime_identity().as_ref() != Some(&identity) + { + return Err(ProxyError::Unauthorized); + } + Ok(()) + } + async fn send( &self, method: &Method, @@ -90,13 +133,35 @@ impl ProxyClient { if let Some(value) = token { request = request.header("X-OpenCodex-API-Key", value); } - request.send().await.map_err(|error| { - if error.is_connect() { - ProxyError::Unreachable - } else { - ProxyError::Decode(error) - } - }) + request.send().await.map_err(map_request_error) + } +} + +fn valid_attestation_proof( + identity: &RuntimeIdentity, + challenge: &str, + proof: Option<&str>, +) -> bool { + let Ok(mut mac) = Hmac::::new_from_slice(identity.attestation_secret.as_bytes()) else { + return false; + }; + mac.update( + format!( + "opencodex-local-management-v1\n{challenge}\n{}\n{}", + identity.pid, identity.port + ) + .as_bytes(), + ); + proof + .and_then(|value| URL_SAFE_NO_PAD.decode(value).ok()) + .is_some_and(|value| mac.verify_slice(&value).is_ok()) +} + +fn map_request_error(error: reqwest::Error) -> ProxyError { + if error.is_connect() { + ProxyError::Unreachable + } else { + ProxyError::Decode(error) } } @@ -109,3 +174,32 @@ async fn decode(response: reqwest::Response) -> Result { } response.json().await.map_err(ProxyError::Decode) } + +#[cfg(test)] +mod tests { + use super::*; + + fn identity() -> RuntimeIdentity { + RuntimeIdentity { + pid: 4242, + port: 10100, + attestation_secret: "BwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwcHBwc".into(), + } + } + + #[test] + fn accepts_only_a_proof_bound_to_the_runtime_identity() { + let challenge = "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"; + let proof = "Yr9EKHjeAFfsFMsF8Xsd7J6LxBYnObweKZlLyTMk0Lo"; + assert!(valid_attestation_proof(&identity(), challenge, Some(proof))); + + let mut replacement = identity(); + replacement.pid += 1; + assert!(!valid_attestation_proof( + &replacement, + challenge, + Some(proof) + )); + assert!(!valid_attestation_proof(&identity(), challenge, None)); + } +}