From dbea3c1340c36ff65d7f212346ede350e1d64684 Mon Sep 17 00:00:00 2001 From: Threated Date: Fri, 17 Jul 2026 11:30:33 +0200 Subject: [PATCH 1/2] feat: reload key id when our own cert expired --- proxy/src/config.rs | 21 +++++++++++++-- proxy/src/crypto.rs | 26 +++++++++++------- proxy/src/main.rs | 2 +- proxy/src/serve_tasks.rs | 57 ++++++++++++++++++++++++++++------------ 4 files changed, 77 insertions(+), 29 deletions(-) diff --git a/proxy/src/config.rs b/proxy/src/config.rs index 5baee9ae..ce63b680 100644 --- a/proxy/src/config.rs +++ b/proxy/src/config.rs @@ -3,6 +3,7 @@ use regex::Regex; use reqwest::Url; use rsa::{pkcs1::DecodeRsaPrivateKey, pkcs8::DecodePrivateKey, RsaPrivateKey}; use shared::{errors::SamplyBeamError, jwt_simple::prelude::RS256KeyPair, logger::LogOptions, openssl::x509::X509, reqwest}; +use tokio::sync::RwLock; use std::{ collections::HashMap, @@ -11,6 +12,7 @@ use std::{ path::{Path, PathBuf}, process::exit, str::FromStr, + sync::Arc, }; use axum::http::HeaderValue; @@ -33,10 +35,25 @@ pub struct Config { #[derive(Debug, Clone)] pub struct ConfigCrypto { - pub privkey_rs256: RS256KeyPair, + pub privkey_rs256: Arc>, pub privkey_rsa: RsaPrivateKey, } +impl ConfigCrypto { + pub async fn reload_public_key_id(&self, config: &Config) -> Result<(), SamplyBeamError> { + let mut key_id = self.privkey_rs256.write().await; + let new = crate::crypto::load_public_crypto_for_proxy( + &*key_id, + self.privkey_rsa.clone(), + &config.proxy_id, + ) + .await? + .1; + *key_id = Arc::into_inner(new.privkey_rs256).unwrap().into_inner(); + Ok(()) + } +} + pub type ApiKey = String; #[derive(Parser, Debug)] @@ -167,7 +184,7 @@ fn load_private_crypto_for_proxy(privkey_file: &PathBuf, proxy_id: &ProxyId) -> )) })?; Ok(ConfigCrypto { - privkey_rs256, + privkey_rs256: Arc::new(RwLock::new(privkey_rs256)), privkey_rsa, }) } diff --git a/proxy/src/crypto.rs b/proxy/src/crypto.rs index 1dc313b6..e2237526 100644 --- a/proxy/src/crypto.rs +++ b/proxy/src/crypto.rs @@ -4,7 +4,7 @@ use axum::{body::Bytes, http::{header, request, Method, Request, StatusCode, Uri use beam_lib::{AppOrProxyId, ProxyId}; use rsa::{pkcs1::{DecodeRsaPrivateKey, DecodeRsaPublicKey}, pkcs8::DecodePrivateKey, RsaPrivateKey, RsaPublicKey}; use shared::{ - async_trait, crypto::{self, asn_str_to_vault_str, get_all_certs_and_clients_by_cname_as_pemstr, get_best_own_certificate, x509_cert_to_x509_public_key, CryptoPublicPortion, GetCerts, ProxyCertInfo}, errors::{CertificateInvalidReason, SamplyBeamError}, http_client::SamplyHttpClient, jwt_simple::prelude::RS256KeyPair, openssl::x509::X509, reqwest, EncryptedMessage, MsgEmpty + EncryptedMessage, MsgEmpty, async_trait, crypto::{self, CryptoPublicPortion, GetCerts, ProxyCertInfo, asn_str_to_vault_str, get_all_certs_and_clients_by_cname_as_pemstr, get_best_own_certificate, x509_cert_to_x509_public_key}, errors::{CertificateInvalidReason, SamplyBeamError}, http_client::SamplyHttpClient, jwt_simple::algorithms::RS256KeyPair, openssl::{pkey::Private, x509::X509}, reqwest }; use tracing::{debug, info, warn, error}; @@ -39,7 +39,7 @@ impl GetCertsFromBroker { .expect("To build request successfully") .into_parts(); - let req = sign_request(body, parts, &self.config) + let req = sign_request(&body, parts, &self.config) .await .map_err(|(_, msg)| SamplyBeamError::SignEncryptError(msg.into()))?; Ok(self.client.execute(req).await?.into()) @@ -102,16 +102,22 @@ impl GetCerts for GetCertsFromBroker { pub async fn init_public_crypto_for_proxy( config: &Config ) -> Result<(ProxyCertInfo, config::ConfigCrypto), SamplyBeamError> { - let (public_info, new_crypto) = load_public_crypto_for_proxy(config).await?; + let (public_info, new_crypto) = load_public_crypto_for_proxy( + &*config.crypto.privkey_rs256.read().await, + config.crypto.privkey_rsa.clone(), + &config.proxy_id + ).await?; let cert_info = ProxyCertInfo::try_from(&public_info.cert)?; Ok((cert_info, new_crypto)) } pub async fn load_public_crypto_for_proxy( - config: &Config, + signer: &RS256KeyPair, + privkey_rsa: RsaPrivateKey, + proxy_id: &ProxyId, ) -> Result<(CryptoPublicPortion, config::ConfigCrypto), SamplyBeamError> { - let publics: Vec = get_all_certs_and_clients_by_cname_as_pemstr(&config.proxy_id) + let publics: Vec = get_all_certs_and_clients_by_cname_as_pemstr(proxy_id) .await .into_iter() .filter_map(|r| { @@ -119,13 +125,15 @@ pub async fn load_public_crypto_for_proxy( .ok() }) .collect(); - let public = get_best_own_certificate(publics, &config.crypto.privkey_rsa).ok_or( + let public = get_best_own_certificate(publics, &privkey_rsa).ok_or( SamplyBeamError::SignEncryptError( "Unable to choose valid, newest certificate for this proxy".into(), ), )?; let serial = asn_str_to_vault_str(public.cert.serial_number())?; - let mut crypto_with_kid = config.crypto.clone(); - crypto_with_kid.privkey_rs256 = crypto_with_kid.privkey_rs256.with_key_id(&serial); + let crypto_with_kid = config::ConfigCrypto { + privkey_rs256: std::sync::Arc::new(tokio::sync::RwLock::new(signer.clone().with_key_id(&serial))), + privkey_rsa, + }; Ok((public, crypto_with_kid)) -} \ No newline at end of file +} diff --git a/proxy/src/main.rs b/proxy/src/main.rs index 1659db7d..dd0f2e7c 100644 --- a/proxy/src/main.rs +++ b/proxy/src/main.rs @@ -175,7 +175,7 @@ fn spawn_controller_polling(client: SamplyHttpClient, config: &'static Config) { .expect("To build request successfully") .into_parts(); - let req = sign_request(body, parts, &config).await.expect("Unable to sign request; this should always work"); + let req = sign_request(&body, parts, &config).await.expect("Unable to sign request; this should always work"); // In the future this will poll actual control related tasks let res = match client.execute(req).await { Ok(res) if res.status() == StatusCode::CONFLICT => { diff --git a/proxy/src/serve_tasks.rs b/proxy/src/serve_tasks.rs index fb7f36d1..b1ac9083 100644 --- a/proxy/src/serve_tasks.rs +++ b/proxy/src/serve_tasks.rs @@ -16,6 +16,7 @@ use rsa::{pkcs8::DecodePublicKey, RsaPublicKey}; use serde::{de::DeserializeOwned, Deserialize, Serialize}; use serde_json::Value; use beam_lib::{AppId, AppOrProxyId, ProxyId}; +use shared::jwt_simple::prelude::RSAKeyPairLike; use shared::{ DecryptableMsg, EncryptableMsg, EncryptedMessage, EncryptedMsgTaskRequest, EncryptedMsgTaskResult, MessageType, Msg, MsgEmpty, MsgId, MsgSigned, MsgTaskRequest, MsgTaskResult, PlainMessage, crypto::{self, CryptoPublicPortion}, crypto_jwt, errors::SamplyBeamError, format_to_without_broker, http_client::SamplyHttpClient, reqwest, sse_event::SseEventType }; @@ -91,20 +92,42 @@ pub(crate) async fn forward_request( MessageType::MsgSocketRequest(socket_req) => info!(from = %socket_req.get_from().hide_broker(), to = %format_to_without_broker(&socket_req.get_to()), id = %socket_req.id, "Submitting socket request"), MessageType::MsgEmpty(..) => {}, }; - let req = sign_request(encrypted_msg, parts, &config).await.map_err(IntoResponse::into_response)?; - trace!("Requesting: {:?}", req); - let resp = client.execute(req).await.map_err(|e| { - if e.is_timeout() { - debug!("Request to broker timed out after set proxy timeout of {PROXY_TIMEOUT}s"); - (StatusCode::GATEWAY_TIMEOUT, "Request to broker timed out ") - } else { - warn!("Request to broker failed: {}", e.to_string()); - (StatusCode::BAD_GATEWAY, "Upstream error; see server logs.") - }.into_response() - })?; + async fn execute_broker_request( + client: &SamplyHttpClient, + req: reqwest::Request, + ) -> Result { + trace!("Requesting: {req:?}"); + client.execute(req).await.map_err(|e| { + if e.is_timeout() { + debug!("Request to broker timed out after set proxy timeout of {PROXY_TIMEOUT}s"); + (StatusCode::GATEWAY_TIMEOUT, "Request to broker timed out ") + } else { + warn!("Request to broker failed: {}", e.to_string()); + (StatusCode::BAD_GATEWAY, "Upstream error; see server logs.") + } + .into_response() + }) + } + let retry_parts = parts.clone(); + let req = sign_request(&encrypted_msg, parts, config) + .await + .map_err(IntoResponse::into_response)?; + let mut resp = execute_broker_request(client, req).await?; if resp.status() == StatusCode::UNAUTHORIZED { - error!("The Broker has rejected our request with 401 Unauthorized. This is likely because our beam certificate expired."); - std::process::exit(401); + warn!("The Broker has rejected our request with 401 Unauthorized. Checking whether our beam certificate was extended."); + if let Err(e) = config.crypto.reload_public_key_id(config).await { + error!("Failed to reload public key: {e:#}. Aborting."); + std::process::exit(401); + } + info!("Loaded the extended beam certificate. Retrying the request."); + let req = sign_request(&encrypted_msg, retry_parts, config) + .await + .map_err(IntoResponse::into_response)?; + resp = execute_broker_request(client, req).await?; + if resp.status() == StatusCode::UNAUTHORIZED { + error!("The Broker rejected our request after trying to reload the beam certificate."); + std::process::exit(401); + } } Ok(resp) } @@ -324,15 +347,15 @@ pub(crate) fn to_server_error(res: Result) -> Result Result { let from = body.get_from(); - let token_without_extended_signature = crypto_jwt::sign_to_jwt(&body, &config.crypto.privkey_rs256) + let signer = config.crypto.privkey_rs256.read().await; + let token_without_extended_signature = crypto_jwt::sign_to_jwt(&body, &signer) .await .map_err(|e| { error!("Crypto failed: {}", e); @@ -356,7 +379,7 @@ pub async fn sign_request( let digest = crypto_jwt::make_extra_fields_digest(&parts.method, &parts.uri, &headers_mut, sig, &from) .map_err(|_| ERR_INTERNALCRYPTO)?; - let token_with_extended_signature = crypto_jwt::sign_to_jwt(&digest, &config.crypto.privkey_rs256) + let token_with_extended_signature = crypto_jwt::sign_to_jwt(&digest, &signer) .await .map_err(|e| { error!("Crypto failed: {}", e); From ec92261eefa58bfc82b4e835ee440cec08e42fa6 Mon Sep 17 00:00:00 2001 From: Threated Date: Mon, 20 Jul 2026 11:19:50 +0200 Subject: [PATCH 2/2] Add test --- proxy/src/serve_tasks.rs | 210 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 210 insertions(+) diff --git a/proxy/src/serve_tasks.rs b/proxy/src/serve_tasks.rs index b1ac9083..1946708e 100644 --- a/proxy/src/serve_tasks.rs +++ b/proxy/src/serve_tasks.rs @@ -503,3 +503,213 @@ async fn encrypt_msg(msg: M) -> Result Result, SamplyBeamError> { + Ok(vec![self.serial.clone()]) + } + + async fn certificate_by_serial_as_pem( + &self, + serial: &str, + ) -> Result { + assert_eq!(serial, self.serial); + Ok(self.certificate.clone()) + } + + async fn im_certificate_as_pem(&self) -> Result { + Ok(self.intermediate.clone()) + } + } + + fn build_certificate( + common_name: &str, + serial: u32, + key: &PKey, + issuer: Option<(&X509, &PKey)>, + ) -> X509 { + let mut name = X509NameBuilder::new().unwrap(); + name.append_entry_by_text("CN", common_name).unwrap(); + let name = name.build(); + + let mut builder = X509::builder().unwrap(); + builder.set_version(2).unwrap(); + let serial = Asn1Integer::from_bn(&BigNum::from_u32(serial).unwrap()).unwrap(); + builder.set_serial_number(&serial).unwrap(); + builder.set_subject_name(&name).unwrap(); + builder.set_pubkey(key).unwrap(); + if let Some((issuer_certificate, _)) = issuer { + builder + .set_issuer_name(issuer_certificate.subject_name()) + .unwrap(); + } else { + builder.set_issuer_name(&name).unwrap(); + } + + let now = SystemTime::now() + .duration_since(SystemTime::UNIX_EPOCH) + .unwrap() + .as_secs() as i64; + let not_before = Asn1Time::from_unix(now - 60).unwrap(); + let not_after = Asn1Time::from_unix(now + 3600).unwrap(); + builder.set_not_before(¬_before).unwrap(); + builder.set_not_after(¬_after).unwrap(); + builder + .sign( + issuer.map(|(_, issuer_key)| issuer_key).unwrap_or(key), + MessageDigest::sha256(), + ) + .unwrap(); + builder.build() + } + + #[tokio::test] + async fn reloads_certificate_serial_and_retries_unauthorized_request() { + beam_lib::set_broker_id("broker.samply.de".into()); + + let key_pair = RS256KeyPair::generate(2048).unwrap(); + let key_pem = key_pair.to_pem().unwrap(); + let private_key = RsaPrivateKey::from_pkcs1_pem(&key_pem) + .or_else(|_| RsaPrivateKey::from_pkcs8_pem(&key_pem)) + .unwrap(); + let leaf_key = PKey::private_key_from_pem(key_pem.as_bytes()).unwrap(); + let root_key = PKey::from_rsa(shared::openssl::rsa::Rsa::generate(2048).unwrap()).unwrap(); + let root_certificate = build_certificate("root", 1, &root_key, None); + let intermediate_key = + PKey::from_rsa(shared::openssl::rsa::Rsa::generate(2048).unwrap()).unwrap(); + let intermediate_certificate = build_certificate( + "intermediate", + 2, + &intermediate_key, + Some((&root_certificate, &root_key)), + ); + let renewed_certificate = build_certificate( + "proxy1.broker.samply.de", + 3, + &leaf_key, + Some((&intermediate_certificate, &intermediate_key)), + ); + let renewed_serial = + asn_str_to_vault_str(renewed_certificate.serial_number()).unwrap(); + + shared::crypto::init_cert_getter(TestCertGetter { + serial: renewed_serial.clone(), + certificate: String::from_utf8(renewed_certificate.to_pem().unwrap()).unwrap(), + intermediate: String::from_utf8(intermediate_certificate.to_pem().unwrap()).unwrap(), + }); + shared::crypto::init_ca_chain(&root_certificate) + .await + .unwrap(); + + let crypto = ConfigCrypto { + privkey_rs256: Arc::new(tokio::sync::RwLock::new( + key_pair.with_key_id(EXPIRED_SERIAL), + )), + privkey_rsa: private_key, + }; + + let seen_serials = Arc::new(Mutex::new(Vec::new())); + let broker = Router::new() + .fallback( + |State(seen_serials): State>>>, + headers: HeaderMap| async move { + let jwt = headers[header::AUTHORIZATION] + .to_str() + .unwrap() + .strip_prefix("SamplyJWT ") + .unwrap(); + let serial = Token::decode_metadata(jwt) + .unwrap() + .key_id() + .unwrap() + .to_owned(); + let mut seen_serials = seen_serials.lock().await; + seen_serials.push(serial); + if seen_serials.len() == 1 { + StatusCode::UNAUTHORIZED + } else { + StatusCode::OK + } + }, + ) + .with_state(seen_serials.clone()); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let broker_addr = listener.local_addr().unwrap(); + let broker_task = tokio::spawn(async move { + axum::serve(listener, broker).await.unwrap(); + }); + + let config = Config { + broker_uri: format!("http://{broker_addr}/").parse().unwrap(), + broker_host_header: HeaderValue::from_str(&broker_addr.to_string()).unwrap(), + bind_addr: "127.0.0.1:0".parse().unwrap(), + proxy_id: beam_lib::ProxyId::new("proxy1.broker.samply.de").unwrap(), + api_keys: HashMap::new(), + tls_ca_certificates: Vec::new(), + crypto, + rootcert: root_certificate, + }; + let sender = beam_lib::AppId::new("app1.proxy1.broker.samply.de").unwrap(); + let request = Request::get("/v1/tasks") + .body(Body::empty()) + .unwrap(); + let client = reqwest::Client::new(); + + let response = tokio::time::timeout( + Duration::from_secs(10), + super::forward_request(request, &config, &sender, &client), + ) + .await + .expect("certificate reload deadlocked") + .expect("request failed"); + + assert_eq!(response.status(), StatusCode::OK); + assert_eq!( + *seen_serials.lock().await, + [EXPIRED_SERIAL, renewed_serial.as_str()] + ); + broker_task.abort(); + } +} \ No newline at end of file