diff --git a/protobufs/operator/vault/unseal.proto b/protobufs/operator/vault/unseal.proto index 9378d5d..6f59675 100644 --- a/protobufs/operator/vault/unseal.proto +++ b/protobufs/operator/vault/unseal.proto @@ -20,6 +20,7 @@ enum UnsealResult { UNSEAL_RESULT_SUCCESS = 1; UNSEAL_RESULT_INVALID_KEY = 2; UNSEAL_RESULT_UNBOOTSTRAPPED = 3; + UNSEAL_RESULT_LOCKED_OUT = 4; } message Request { diff --git a/server/Cargo.lock b/server/Cargo.lock index e36c264..dcf3fa1 100644 --- a/server/Cargo.lock +++ b/server/Cargo.lock @@ -771,6 +771,7 @@ dependencies = [ "proptest", "prost-types", "rand 0.10.1", + "rand_core 0.10.1", "rcgen", "restructed", "rstest", @@ -6261,18 +6262,18 @@ dependencies = [ [[package]] name = "zeroize" -version = "1.8.2" +version = "1.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0" +checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" dependencies = [ "zeroize_derive", ] [[package]] name = "zeroize_derive" -version = "1.4.3" +version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85a5b4158499876c763cb03bc4e49185d3cccbabb15b33c627f7884f43db852e" +checksum = "3c50655cbb0fe3fc43170059e702f1ce5e19b84cec58dc87b037a09935c2f328" dependencies = [ "proc-macro2", "quote", diff --git a/server/Cargo.toml b/server/Cargo.toml index cd588a8..ab81ee3 100644 --- a/server/Cargo.toml +++ b/server/Cargo.toml @@ -21,6 +21,7 @@ mutants = "0.0.4" prost = "0.14.3" prost-types = { version = "0.14.3", features = ["chrono"] } rand = "0.10.1" +rand_core = "0.10.1" rcgen = { version = "0.14.7", features = [ "aws_lc_rs", "pem", "x509-parser", "zeroize" ], default-features = false } rstest = "0.26.1" rustls = { version = "0.23.40", features = ["aws-lc-rs", "logging", "prefer-post-quantum", "std"], default-features = false } diff --git a/server/crates/arbiter-server/Cargo.toml b/server/crates/arbiter-server/Cargo.toml index 7790bd6..e3a202c 100644 --- a/server/crates/arbiter-server/Cargo.toml +++ b/server/crates/arbiter-server/Cargo.toml @@ -31,6 +31,7 @@ diesel_migrations = { version = "2.3.2", features = ["sqlite"] } async-trait.workspace = true tokio-stream.workspace = true rand.workspace = true +rand_core.workspace = true rcgen.workspace = true chrono.workspace = true kameo.workspace = true diff --git a/server/crates/arbiter-server/src/actors/bootstrap.rs b/server/crates/arbiter-server/src/actors/bootstrap.rs index 8ec0059..4859fdd 100644 --- a/server/crates/arbiter-server/src/actors/bootstrap.rs +++ b/server/crates/arbiter-server/src/actors/bootstrap.rs @@ -1,29 +1,48 @@ use crate::db::{self, DatabasePool, schema}; +use arbiter_crypto::safecell::{SafeCell, SafeCellHandle as _}; use arbiter_proto::{BOOTSTRAP_PATH, home_path}; use diesel::QueryDsl; use diesel_async::RunQueryDsl; use kameo::{Actor, messages}; -use rand::{RngExt, distr::Alphanumeric, make_rng, rngs::StdRng}; +use rand::{RngExt, distr::Alphanumeric, rngs::SysRng}; +use rand_core::UnwrapErr; +use std::path::{Path, PathBuf}; use subtle::ConstantTimeEq as _; use thiserror::Error; +use tracing::warn; const TOKEN_LENGTH: usize = 64; -pub async fn generate_token() -> Result { - let rng: StdRng = make_rng(); +async fn write_token_file(path: &Path, content: &str) -> Result<(), std::io::Error> { + tokio::fs::write(path, content.as_bytes()).await?; - let token = rng.sample_iter(Alphanumeric).take(TOKEN_LENGTH).fold( - String::default(), - |mut accum, char| { - accum += char.to_string().as_str(); - accum - }, - ); + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt as _; + tokio::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600)).await?; + } - tokio::fs::write(home_path()?.join(BOOTSTRAP_PATH), token.as_str()).await?; + Ok(()) +} - Ok(token) +async fn generate_token(path: &Path) -> Result, std::io::Error> { + let mut cell = SafeCell::new([0u8; TOKEN_LENGTH]); + { + let mut buf = cell.write(); + for (slot, b) in buf + .iter_mut() + .zip(UnwrapErr(SysRng).sample_iter(Alphanumeric)) + { + *slot = b; + } + } + + let token_str = cell.read_inline(|buf| String::from_utf8_lossy(buf.as_ref()).into_owned()); + + write_token_file(path, &token_str).await?; + + Ok(cell) } #[derive(Error, Debug)] @@ -40,7 +59,8 @@ pub enum Error { #[derive(Actor)] pub struct Bootstrapper { - token: Option, + token: Option>, + token_path: Option, } impl Bootstrapper { @@ -54,34 +74,37 @@ impl Bootstrapper { .await? }; - let token = if row_count == 0 { - let token = generate_token().await?; - Some(token) + let (token, token_path) = if row_count == 0 { + let path = home_path()?.join(BOOTSTRAP_PATH); + let token = generate_token(&path).await?; + (Some(token), Some(path)) } else { - None + (None, None) }; - Ok(Self { token }) + Ok(Self { token, token_path }) + } +} + +impl Bootstrapper { + fn is_correct_token(&mut self, token: &[u8]) -> bool { + self.token.as_mut().is_some_and(|expected| { + expected.read_inline(|exp| bool::from(exp.as_ref().ct_eq(token))) + }) } } #[messages] impl Bootstrapper { #[message] - pub fn is_correct_token(&self, token: String) -> bool { - self.token.as_ref().is_some_and(|expected| { - let expected_bytes = expected.as_bytes(); - let token_bytes = token.as_bytes(); - - let choice = expected_bytes.ct_eq(token_bytes); - bool::from(choice) - }) - } - - #[message] - pub fn consume_token(&mut self, token: String) -> bool { - if self.is_correct_token(token) { + pub async fn consume_token(&mut self, token: Vec) -> bool { + if self.is_correct_token(&token) { self.token = None; + if let Some(path) = self.token_path.take() + && let Err(e) = tokio::fs::remove_file(&path).await + { + warn!(error = ?e, path = ?path, "Failed to delete bootstrap token file after consumption"); + } true } else { false @@ -92,7 +115,9 @@ impl Bootstrapper { #[messages] impl Bootstrapper { #[message] - pub fn get_token(&self) -> Option { - self.token.clone() + pub fn get_token(&mut self) -> Option { + self.token + .as_mut() + .map(|cell| cell.read_inline(|buf| String::from_utf8_lossy(buf.as_ref()).into_owned())) } } diff --git a/server/crates/arbiter-server/src/actors/evm/mod.rs b/server/crates/arbiter-server/src/actors/evm/mod.rs index 06e672c..c2ea4f9 100644 --- a/server/crates/arbiter-server/src/actors/evm/mod.rs +++ b/server/crates/arbiter-server/src/actors/evm/mod.rs @@ -1,6 +1,6 @@ use crate::{ actors::vault::{CreateNew, Decrypt, Vault}, - crypto::integrity, + crypto::integrity::{self, Integrable}, db::{ DatabaseError, DatabasePool, models::{self}, @@ -25,14 +25,35 @@ use diesel::{ use diesel_async::RunQueryDsl; use kameo::{Actor, actor::ActorRef, messages}; use rand::{SeedableRng, rng, rngs::StdRng}; +use tracing::error; pub use crate::evm::safe_signer; +/// Integrity guard that binds a wallet's encrypted key ID to its Ethereum address. +/// Both fields are included in the HMAC — swapping `aead_encrypted_id` in the DB +/// invalidates the envelope MAC, and the AEAD ciphertext is also bound to `address` +/// as AAD, so decryption fails too. +#[derive(arbiter_macros::Hashable)] +struct EvmWalletIntegrity { + aead_encrypted_id: i32, + address: Address, +} + +impl Integrable for EvmWalletIntegrity { + const KIND: &'static str = "evm_wallet"; +} + #[derive(Debug, thiserror::Error)] pub enum SignTransactionError { #[error("Wallet not found")] WalletNotFound, + #[error("Decrypted key does not match requested wallet address")] + KeyAddressMismatch, + + #[error("Internal signing error")] + Internal, + #[error("Database error: {0}")] Database(#[from] DatabaseError), @@ -64,6 +85,12 @@ pub enum Error { Integrity(#[from] integrity::Error), } +impl From for Error { + fn from(e: diesel::result::Error) -> Self { + Self::Database(DatabaseError::from(e)) + } +} + #[derive(Actor)] pub struct EvmActor { pub vault: ActorRef, @@ -97,20 +124,39 @@ impl EvmActor { let aead_id: i32 = self .vault - .ask(CreateNew { plaintext }) + .ask(CreateNew { + plaintext, + aad: address.as_slice().to_vec(), + }) .await .map_err(|_| Error::VaultSend)?; let mut conn = self.db.get().await.map_err(DatabaseError::from)?; - let wallet_id = insert_into(schema::evm_wallet::table) - .values(&models::NewEvmWallet { - address: address.as_slice().to_vec(), - aead_encrypted_id: aead_id, + let wallet_id = conn + .exclusive_transaction(async |conn| { + let wallet_id: i32 = insert_into(schema::evm_wallet::table) + .values(&models::NewEvmWallet { + address: address.as_slice().to_vec(), + aead_encrypted_id: aead_id, + }) + .returning(schema::evm_wallet::id) + .get_result(conn) + .await + .map_err(DatabaseError::from) + .map_err(Error::Database)?; + + integrity::sign_entity( + conn, + &self.vault, + &EvmWalletIntegrity { address, aead_encrypted_id: aead_id }, + wallet_id, + ) + .await + .map_err(Error::Integrity)?; + + Ok::(wallet_id) }) - .returning(schema::evm_wallet::id) - .get_result(&mut conn) - .await - .map_err(DatabaseError::from)?; + .await?; Ok((wallet_id, address)) } @@ -241,16 +287,51 @@ impl EvmActor { .ok_or(SignTransactionError::WalletNotFound)?; drop(conn); + let mut conn = self.db.get().await.map_err(DatabaseError::from)?; + let attestation = integrity::verify_entity( + &mut conn, + &self.vault, + &EvmWalletIntegrity { + address: wallet_address, + aead_encrypted_id: wallet.aead_encrypted_id, + }, + wallet.id, + ) + .await + .map_err(|e| { + error!(?e, wallet_id = wallet.id, "EVM wallet integrity check failed"); + SignTransactionError::Internal + })?; + drop(conn); + + if attestation != integrity::AttestationStatus::Attested { + error!( + wallet_id = wallet.id, + "EVM wallet integrity unavailable; refusing to sign" + ); + return Err(SignTransactionError::Internal); + } + let raw_key: SafeCell> = self .vault .ask(Decrypt { aead_id: wallet.aead_encrypted_id, + aad: wallet.address.clone(), }) .await .map_err(|_| SignTransactionError::VaultSend)?; let signer = safe_signer::SafeSigner::from_cell(raw_key)?; + if signer.address() != wallet_address { + error!( + expected = %wallet_address, + actual = %signer.address(), + "Decrypted private key address does not match requested wallet" + ); + return Err(SignTransactionError::KeyAddressMismatch); + } + self.engine .evaluate_transaction(wallet_access, transaction.clone(), RunKind::Execution) .await?; diff --git a/server/crates/arbiter-server/src/actors/flow_coordinator/mod.rs b/server/crates/arbiter-server/src/actors/flow_coordinator/mod.rs index fa31334..5a8dc0e 100644 --- a/server/crates/arbiter-server/src/actors/flow_coordinator/mod.rs +++ b/server/crates/arbiter-server/src/actors/flow_coordinator/mod.rs @@ -20,6 +20,8 @@ pub mod client_connect_approval; pub struct FlowCoordinator { pub clients: HashMap>, + /// Maps DB `client_id` → `ActorId` for fast connected-client lookup. + client_ids: HashMap, operator_registry: ActorRef, } @@ -27,6 +29,7 @@ impl FlowCoordinator { pub fn new(operator_registry: ActorRef) -> Self { Self { clients: HashMap::default(), + client_ids: HashMap::default(), operator_registry, } } @@ -48,6 +51,7 @@ impl Actor for FlowCoordinator { _: ActorStopReason, ) -> Result, Self::Error> { if self.clients.remove(&id).is_some() { + self.client_ids.retain(|_, actor_id| *actor_id != id); info!( ?id, actor = "FlowCoordinator", @@ -75,14 +79,28 @@ impl FlowCoordinator { #[message(ctx)] pub async fn register_client( &mut self, + client_id: i32, actor: ActorRef, ctx: &mut Context, ) { - info!(id = %actor.id(), actor = "FlowCoordinator", event = "client.connected"); + info!(id = %actor.id(), client_id, actor = "FlowCoordinator", event = "client.connected"); ctx.actor_ref().link(&actor).await; + self.client_ids.insert(client_id, actor.id()); self.clients.insert(actor.id(), actor); } + #[message] + pub fn is_client_connected(&self, client_id: i32) -> bool { + self.client_ids.contains_key(&client_id) + } + + /// Returns the DB `client_ids` of all currently connected SDK clients. + /// Used by operator sessions on startup to seed their approved-client set. + #[message] + pub fn get_connected_client_ids(&self) -> Vec { + self.client_ids.keys().copied().collect() + } + #[message(ctx)] pub async fn request_client_approval( &mut self, diff --git a/server/crates/arbiter-server/src/actors/vault/mod.rs b/server/crates/arbiter-server/src/actors/vault/mod.rs index 74aa15e..ec4c9fc 100644 --- a/server/crates/arbiter-server/src/actors/vault/mod.rs +++ b/server/crates/arbiter-server/src/actors/vault/mod.rs @@ -22,7 +22,7 @@ use hmac::{KeyInit as _, Mac as _}; use kameo::{Actor, Reply, actor::ActorRef, messages}; use kameo_actors::message_bus::{MessageBus, Publish}; use strum::{EnumDiscriminants, IntoDiscriminant}; -use tracing::{error, info}; +use tracing::{error, info, warn}; pub mod events { @@ -46,6 +46,8 @@ pub enum Error { Sealed, #[error("Invalid key provided")] InvalidKey, + #[error("Vault locked: too many failed unseal attempts")] + LockedOut, #[error("Requested aead entry not found")] NotFound, @@ -61,6 +63,9 @@ pub enum Error { #[error("Broken database")] BrokenDatabase, + + #[error("Integrity key version mismatch: envelope uses key {envelope}, current key is {current}")] + KeyVersionMismatch { envelope: i32, current: i32 }, } struct Unsealed { @@ -79,6 +84,8 @@ enum State { Unsealed(Unsealed), } +const MAX_UNSEAL_ATTEMPTS: u32 = 5; + /// Manages vault root key and tracks current state of the vault (bootstrapped/unbootstrapped, sealed/unsealed). /// /// Provides API for encrypting and decrypting data using the vault root key. @@ -88,6 +95,7 @@ pub struct Vault { db: db::DatabasePool, state: State, events: ActorRef, + unseal_failures: u32, } #[messages] @@ -110,7 +118,7 @@ impl Vault { } }; - Ok(Self { db, state, events }) + Ok(Self { db, state, events, unseal_failures: 0 }) } // Exclusive transaction to avoid race condtions if multiple vaults write @@ -219,6 +227,10 @@ impl Vault { #[message] pub async fn try_unseal(&mut self, seal_key_raw: SafeCell>) -> Result<(), Error> { + if self.unseal_failures >= MAX_UNSEAL_ATTEMPTS { + return Err(Error::LockedOut); + } + let State::Sealed { root_key_history_id, } = &self.state @@ -251,13 +263,27 @@ impl Vault { Error::BrokenDatabase })?; - seal_key + if seal_key .decrypt_in_place(&nonce, v1::ROOT_KEY_TAG, &mut root_key) - .map_err(|err| { - error!(?err, "Failed to unseal root key: invalid seal key"); - Error::InvalidKey - })?; + .is_err() + { + self.unseal_failures += 1; + if self.unseal_failures >= MAX_UNSEAL_ATTEMPTS { + error!( + attempts = self.unseal_failures, + "Vault locked: maximum failed unseal attempts reached" + ); + } else { + warn!( + attempts = self.unseal_failures, + remaining = MAX_UNSEAL_ATTEMPTS - self.unseal_failures, + "Failed unseal attempt" + ); + } + return Err(Error::InvalidKey); + } + self.unseal_failures = 0; self.state = State::Unsealed(Unsealed { root_key_history_id: current_key.id, root_key: KeyCell::try_from(root_key).map_err(|err| { @@ -272,8 +298,10 @@ impl Vault { Ok(()) } + /// Decrypts an AEAD entry. The `aad` must match the value used at encryption time; + /// a mismatch causes authentication failure, preventing cross-wallet key swaps. #[message] - pub async fn decrypt(&mut self, aead_id: i32) -> Result>, Error> { + pub async fn decrypt(&mut self, aead_id: i32, aad: Vec) -> Result>, Error> { let Unsealed { root_key, .. } = Self::expect_unsealed(&mut self.state)?; let row: models::AeadEncrypted = { @@ -295,13 +323,15 @@ impl Vault { Error::BrokenDatabase })?; let mut output = SafeCell::new(row.ciphertext); - root_key.decrypt_in_place(&nonce, v1::TAG, &mut output)?; + root_key.decrypt_in_place(&nonce, &aad, &mut output)?; Ok(output) } + /// Creates a new `aead_encrypted` entry and returns its ID. + /// The `aad` is bound into the ciphertext and must be reproduced exactly at decryption time. // Creates new `aead_encrypted` entry in the database and returns it's ID #[message] - pub async fn create_new(&mut self, mut plaintext: SafeCell>) -> Result { + pub async fn create_new(&mut self, mut plaintext: SafeCell>, aad: Vec) -> Result { let Unsealed { root_key, root_key_history_id, @@ -313,7 +343,7 @@ impl Vault { let mut ciphertext_buffer = plaintext.write(); let ciphertext_buffer: &mut Vec = ciphertext_buffer.as_mut(); - root_key.encrypt_in_place(&nonce, v1::TAG, &mut *ciphertext_buffer)?; + root_key.encrypt_in_place(&nonce, &aad, &mut *ciphertext_buffer)?; let ciphertext = std::mem::take(ciphertext_buffer); @@ -370,7 +400,10 @@ impl Vault { } = Self::expect_unsealed(&mut self.state)?; if *root_key_history_id != key_version { - return Ok(false); + return Err(Error::KeyVersionMismatch { + envelope: key_version, + current: *root_key_history_id, + }); } let mut hmac = root_key.0.read_inline(|k| { @@ -445,7 +478,7 @@ mod tests { assert_eq!(root_row.data_encryption_nonce, n2.to_vec()); let id = actor - .create_new(SafeCell::new(b"post-interleave".to_vec())) + .create_new(SafeCell::new(b"post-interleave".to_vec()), b"test-aad".to_vec()) .await .unwrap(); let row: models::AeadEncrypted = schema::aead_encrypted::table diff --git a/server/crates/arbiter-server/src/crypto/integrity/v1.rs b/server/crates/arbiter-server/src/crypto/integrity/v1.rs index edb2274..0b053d9 100644 --- a/server/crates/arbiter-server/src/crypto/integrity/v1.rs +++ b/server/crates/arbiter-server/src/crypto/integrity/v1.rs @@ -192,7 +192,9 @@ pub async fn verify_entity( Ok(false) => Err(Error::MacMismatch { entity_kind: E::KIND, }), - Err(SendError::HandlerError(vault::Error::Sealed)) => Ok(AttestationStatus::Unavailable), + Err(SendError::HandlerError( + vault::Error::Sealed | vault::Error::KeyVersionMismatch { .. }, + )) => Ok(AttestationStatus::Unavailable), Err(_) => Err(Error::VaultSend), } } @@ -331,4 +333,47 @@ mod tests { .unwrap_err(); assert!(matches!(err, Error::MacMismatch { .. })); } + + #[tokio::test] + async fn key_version_mismatch_returns_unavailable_not_mac_mismatch() { + use crate::db::schema::integrity_envelope; + use super::AttestationStatus; + + const ENTITY_ID: &[u8] = b"entity-id-rotation-test"; + + let db = db::create_test_pool().await; + let vault = bootstrapped_vault(&db).await; + let mut conn = db.get().await.unwrap(); + + let entity = DummyEntity { + payload_version: 1, + payload: b"payload-v1".to_vec(), + }; + + sign_entity(&mut conn, &vault, &entity, ENTITY_ID) + .await + .unwrap(); + + // Simulate key rotation: update the stored key_version to a stale value. + // After real rotation the vault's root_key_history_id would advance, but + // here we achieve the same mismatch by back-dating the envelope's key_version. + diesel::update(integrity_envelope::table) + .filter(integrity_envelope::entity_kind.eq("dummy_entity")) + .filter(integrity_envelope::entity_id.eq(ENTITY_ID)) + .set(integrity_envelope::key_version.eq(0)) + .execute(&mut conn) + .await + .unwrap(); + + // Must NOT error — version mismatch is Unavailable, not tampered. + let status = verify_entity(&mut conn, &vault, &entity, ENTITY_ID) + .await + .expect("key version mismatch must not be treated as an error"); + + assert_eq!( + status, + AttestationStatus::Unavailable, + "stale key_version must yield Unavailable, not MacMismatch" + ); + } } diff --git a/server/crates/arbiter-server/src/grpc/operator/auth.rs b/server/crates/arbiter-server/src/grpc/operator/auth.rs index fa15310..bedf1ab 100644 --- a/server/crates/arbiter-server/src/grpc/operator/auth.rs +++ b/server/crates/arbiter-server/src/grpc/operator/auth.rs @@ -171,7 +171,7 @@ impl Receiver for AuthTransportAdapter<'_> { Some(auth::Inbound::AuthChallengeRequest { pubkey, - bootstrap_token, + bootstrap_token: bootstrap_token.map(String::into_bytes), }) } AuthRequestPayload::ChallengeSolution(ProtoAuthChallengeSolution { signature }) => { diff --git a/server/crates/arbiter-server/src/grpc/operator/evm.rs b/server/crates/arbiter-server/src/grpc/operator/evm.rs index 0b3ac2c..2a70abb 100644 --- a/server/crates/arbiter-server/src/grpc/operator/evm.rs +++ b/server/crates/arbiter-server/src/grpc/operator/evm.rs @@ -217,6 +217,11 @@ async fn handle_sign_transaction( result: Some(vet_error.convert()), } } + Err(kameo::error::SendError::HandlerError( + SessionSignTransactionError::ClientNotConnected, + )) => { + return Err(Status::permission_denied("client not connected")); + } Err(kameo::error::SendError::HandlerError(SessionSignTransactionError::Internal)) => { EvmSignTransactionResponse { result: Some(EvmSignTransactionResult::Error( diff --git a/server/crates/arbiter-server/src/grpc/operator/vault_gate/outbound.rs b/server/crates/arbiter-server/src/grpc/operator/vault_gate/outbound.rs index 4a2f072..539d672 100644 --- a/server/crates/arbiter-server/src/grpc/operator/vault_gate/outbound.rs +++ b/server/crates/arbiter-server/src/grpc/operator/vault_gate/outbound.rs @@ -87,6 +87,7 @@ impl TryConvert for vault_gate::Outbound { let proto_result = match result { Ok(()) => ProtoUnsealResult::Success, Err(vault_gate::Error::InvalidKey) => ProtoUnsealResult::InvalidKey, + Err(vault_gate::Error::LockedOut) => ProtoUnsealResult::LockedOut, Err(err) => { warn!(?err, "unseal failed"); return Err(Status::internal("Failed to unseal vault")); diff --git a/server/crates/arbiter-server/src/peers/client/auth.rs b/server/crates/arbiter-server/src/peers/client/auth.rs index f488161..e4b6fb3 100644 --- a/server/crates/arbiter-server/src/peers/client/auth.rs +++ b/server/crates/arbiter-server/src/peers/client/auth.rs @@ -8,7 +8,7 @@ use crate::{ crypto::integrity::{self, AttestationStatus}, db::{ self, - models::{ProgramClientMetadata, SqliteTimestamp}, + models::ProgramClientMetadata, schema::program_client, }, }; @@ -18,14 +18,13 @@ use arbiter_proto::{ transport::{Bi, expect_message}, }; -use chrono::Utc; use diesel::{ ExpressionMethods as _, OptionalExtension as _, QueryDsl as _, SelectableHelper as _, - dsl::insert_into, update, + dsl::insert_into, }; use diesel_async::RunQueryDsl as _; use kameo::{actor::ActorRef, error::SendError}; -use tracing::error; +use tracing::{error, warn}; #[derive(thiserror::Error, Debug, Clone, PartialEq, Eq)] pub enum Error { @@ -211,71 +210,47 @@ async fn insert_client( .await } -async fn sync_client_metadata( +/// Compares stored metadata against what a reconnecting client presents. +/// Metadata is frozen after initial operator approval and must not be silently +/// overwritten. Doing so would let an approved client forge its displayed +/// identity in later approval prompts. Drift is logged and ignored. +async fn check_metadata_drift( db: &db::DatabasePool, client_id: i32, - metadata: &ClientMetadata, + presented: &ClientMetadata, ) -> Result<(), Error> { - use crate::db::schema::{client_metadata, client_metadata_history}; - - let now = SqliteTimestamp(Utc::now()); + use crate::db::schema::client_metadata; let mut conn = db.get().await.map_err(|e| { error!(error = ?e, "Database pool error"); Error::DatabasePoolUnavailable })?; - conn.exclusive_transaction(async |conn| { - let (current_metadata_id, current): (i32, ProgramClientMetadata) = program_client::table - .find(client_id) - .inner_join(client_metadata::table) - .select(( - program_client::metadata_id, - ProgramClientMetadata::as_select(), - )) - .first(&mut *conn) - .await?; + let current: ProgramClientMetadata = program_client::table + .find(client_id) + .inner_join(client_metadata::table) + .select(ProgramClientMetadata::as_select()) + .first(&mut conn) + .await + .map_err(|e| { + error!(error = ?e, "Database error"); + Error::DatabaseOperationFailed + })?; - let unchanged = current.name == metadata.name - && current.description == metadata.description - && current.version == metadata.version; - if unchanged { - return Ok(()); - } + let changed = current.name != presented.name + || current.description != presented.description + || current.version != presented.version; - insert_into(client_metadata_history::table) - .values(( - client_metadata_history::metadata_id.eq(current_metadata_id), - client_metadata_history::client_id.eq(client_id), - )) - .execute(&mut *conn) - .await?; + if changed { + warn!( + client_id, + stored_name = %current.name, + presented_name = %presented.name, + "reconnecting client presented different metadata; ignoring - metadata is frozen after operator approval" + ); + } - let metadata_id = insert_into(client_metadata::table) - .values(( - client_metadata::name.eq(&metadata.name), - client_metadata::description.eq(&metadata.description), - client_metadata::version.eq(&metadata.version), - )) - .returning(client_metadata::id) - .get_result::(&mut *conn) - .await?; - - update(program_client::table.find(client_id)) - .set(( - program_client::metadata_id.eq(metadata_id), - program_client::updated_at.eq(now), - )) - .execute(&mut *conn) - .await?; - - Ok::<(), diesel::result::Error>(()) - }) - .await - .map_err(|e| { - error!(error = ?e, "Database error"); - Error::DatabaseOperationFailed - }) + Ok(()) } async fn challenge_client( @@ -324,6 +299,7 @@ where let client_id = if let Some(id) = get_client_id(&props.db, &pubkey).await? { verify_integrity(&props.db, &props.actors.vault, &pubkey).await?; + check_metadata_drift(&props.db, id, &metadata).await?; id } else { approve_new_client( @@ -337,8 +313,6 @@ where insert_client(&props.db, &props.actors.vault, &pubkey, &metadata).await? }; - sync_client_metadata(&props.db, client_id, &metadata).await?; - let challenge = AuthChallenge::generate(&mut rand::rng()); challenge_client(transport, pubkey, challenge).await?; diff --git a/server/crates/arbiter-server/src/peers/client/session.rs b/server/crates/arbiter-server/src/peers/client/session.rs index 23ebf3c..6106a6d 100644 --- a/server/crates/arbiter-server/src/peers/client/session.rs +++ b/server/crates/arbiter-server/src/peers/client/session.rs @@ -83,7 +83,7 @@ impl Actor for ClientSession { args.props .actors .flow_coordinator - .ask(RegisterClient { actor: this }) + .ask(RegisterClient { client_id: args.client_id, actor: this }) .await .map_err(|_| Error::ConnectionRegistrationFailed)?; Ok(args) diff --git a/server/crates/arbiter-server/src/peers/operator/auth/mod.rs b/server/crates/arbiter-server/src/peers/operator/auth/mod.rs index 8bea8a0..4b2ecb6 100644 --- a/server/crates/arbiter-server/src/peers/operator/auth/mod.rs +++ b/server/crates/arbiter-server/src/peers/operator/auth/mod.rs @@ -14,7 +14,7 @@ mod state; pub enum Inbound { AuthChallengeRequest { pubkey: authn::PublicKey, - bootstrap_token: Option, + bootstrap_token: Option>, }, AuthChallengeSolution { signature: Vec, diff --git a/server/crates/arbiter-server/src/peers/operator/auth/state.rs b/server/crates/arbiter-server/src/peers/operator/auth/state.rs index 38f1ecd..3208502 100644 --- a/server/crates/arbiter-server/src/peers/operator/auth/state.rs +++ b/server/crates/arbiter-server/src/peers/operator/auth/state.rs @@ -16,13 +16,12 @@ use tracing::error; pub(super) struct ChallengeRequest { pub(super) pubkey: authn::PublicKey, - pub(super) bootstrap_token: Option, + pub(super) bootstrap_token: Option>, } pub struct ChallengeContext { pub(super) challenge: AuthChallenge, pub(super) pubkey: authn::PublicKey, - pub(super) bootstrap_token: Option, } pub(super) struct ChallengeSolution { @@ -79,11 +78,16 @@ async fn register_key(db: &DatabasePool, pubkey: &authn::PublicKey) -> Result { pub(super) conn: &'a mut OperatorConnection, pub(super) transport: &'a mut T, + bootstrap_token: Option>, } impl<'a, T: ?Sized> AuthContext<'a, T> { pub(super) const fn new(conn: &'a mut OperatorConnection, transport: &'a mut T) -> Self { - Self { conn, transport } + Self { + conn, + transport, + bootstrap_token: None, + } } } @@ -108,6 +112,8 @@ where } } + self.bootstrap_token = bootstrap_token; + let challenge = AuthChallenge::generate(&mut rand::rng()); self.transport @@ -120,20 +126,12 @@ where Error::Transport })?; - Ok(ChallengeContext { - challenge, - pubkey, - bootstrap_token, - }) + Ok(ChallengeContext { challenge, pubkey }) } async fn verify_solution( &mut self, - ChallengeContext { - challenge, - pubkey, - bootstrap_token, - }: &ChallengeContext, + ChallengeContext { challenge, pubkey }: &ChallengeContext, ChallengeSolution { solution }: ChallengeSolution, ) -> Result { let signature = authn::Signature::try_from(solution.as_slice()).map_err(|()| { @@ -152,15 +150,13 @@ where } // Resolve client id: bootstrap (consume token + register) or lookup - let id = match bootstrap_token { + let id = match self.bootstrap_token.take() { Some(token) => { let token_ok: bool = self .conn .actors .bootstrapper - .ask(ConsumeToken { - token: token.clone(), - }) + .ask(ConsumeToken { token }) .await .map_err(|e| { error!(?e, "Failed to consume bootstrap token"); diff --git a/server/crates/arbiter-server/src/peers/operator/session/handlers.rs b/server/crates/arbiter-server/src/peers/operator/session/handlers.rs index 5ac6cb4..8240817 100644 --- a/server/crates/arbiter-server/src/peers/operator/session/handlers.rs +++ b/server/crates/arbiter-server/src/peers/operator/session/handlers.rs @@ -4,9 +4,12 @@ use crate::{ ClientSignTransaction, Generate, ListWallets, OperatorCreateGrant, OperatorListGrants, SignTransactionError as EvmSignError, }, - actors::flow_coordinator::client_connect_approval::ClientApprovalAnswer, + actors::flow_coordinator::{IsClientConnected, client_connect_approval::ClientApprovalAnswer}, actors::vault::VaultState, - db::models::{EvmWalletAccess, NewEvmWalletAccess, ProgramClient, ProgramClientMetadata}, + db::{ + models::{EvmWalletAccess, NewEvmWalletAccess, ProgramClient, ProgramClientMetadata}, + schema::program_client, + }, evm::policies::{Grant, SpecificGrant}, }; use arbiter_crypto::authn; @@ -15,13 +18,16 @@ use alloy::{consensus::TxEip1559, primitives::Address, signers::Signature}; use diesel::{ExpressionMethods as _, QueryDsl as _, SelectableHelper}; use diesel_async::{AsyncConnection, RunQueryDsl}; use kameo::{error::SendError, messages, prelude::Context}; -use tracing::error; +use tracing::{error, info, warn}; #[derive(Debug, Error)] pub enum SignTransactionError { #[error("Policy evaluation failed")] Vet(#[from] crate::evm::VetError), + #[error("Client not connected")] + ClientNotConnected, + #[error("Internal signing error")] Internal, } @@ -141,6 +147,30 @@ impl OperatorSession { wallet_address: Address, transaction: TxEip1559, ) -> Result { + if !self.approved_client_ids.contains(&client_id) { + warn!( + client_id, + "operator attempted to sign for client not in its approved set" + ); + return Err(SignTransactionError::ClientNotConnected); + } + + let connected = self + .props + .actors + .flow_coordinator + .ask(IsClientConnected { client_id }) + .await + .unwrap_or(false); + + if !connected { + self.approved_client_ids.remove(&client_id); + warn!(client_id, "operator attempted to sign for disconnected client"); + return Err(SignTransactionError::ClientNotConnected); + } + + info!(client_id, event = "sign_transaction", "operator.sign_transaction"); + match self .props .actors @@ -196,7 +226,7 @@ impl OperatorSession { use crate::db::schema::evm_wallet_access; for entry in entries { diesel::delete(evm_wallet_access::table) - .filter(evm_wallet_access::wallet_id.eq(entry)) + .filter(evm_wallet_access::id.eq(entry)) .execute(&mut *conn) .await?; } @@ -249,6 +279,30 @@ impl OperatorSession { ctx.actor_ref().unlink(&pending_approval.controller).await; + if approved { + let pubkey_bytes = pending_approval.pubkey.to_bytes(); + match self.props.db.get().await { + Ok(mut conn) => { + match program_client::table + .filter(program_client::public_key.eq(pubkey_bytes.as_slice())) + .select(program_client::id) + .first::(&mut conn) + .await + { + Ok(client_id) => { + self.approved_client_ids.insert(client_id); + } + Err(err) => { + error!(?err, "Failed to look up client_id for approved pubkey"); + } + } + } + Err(err) => { + error!(?err, "DB pool error after client approval"); + } + } + } + Ok(()) } @@ -271,3 +325,142 @@ impl OperatorSession { Ok(clients) } } + +#[cfg(test)] +mod tests { + use crate::db::{self, models::NewEvmWalletAccess, schema::evm_wallet_access}; + use diesel::{ExpressionMethods as _, QueryDsl as _, SelectableHelper}; + use diesel_async::{AsyncConnection, RunQueryDsl}; + + /// Regression test: revocation must delete by access-entry `id`, not by `wallet_id`. + /// + /// Before the fix, revoking `entry_id=1` would delete all rows where `wallet_id=1`, + /// wiping out every client's access to wallet #1. + #[tokio::test] + async fn revoke_deletes_by_entry_id_not_wallet_id() { + use crate::db::models::EvmWalletAccess; + + let pool = db::create_test_pool().await; + let mut conn = pool.get().await.expect("pool connection"); + + // Insert two access entries for the same wallet but different clients. + // entry A: id will be 1, wallet_id=1, client_id=10 + // entry B: id will be 2, wallet_id=1, client_id=20 + let entry_a = diesel::insert_into(evm_wallet_access::table) + .values(NewEvmWalletAccess { + wallet_id: 1, + client_id: 10, + }) + .returning(EvmWalletAccess::as_select()) + .get_result(&mut *conn) + .await + .expect("insert entry A"); + + let entry_b = diesel::insert_into(evm_wallet_access::table) + .values(NewEvmWalletAccess { + wallet_id: 1, + client_id: 20, + }) + .returning(EvmWalletAccess::as_select()) + .get_result(&mut *conn) + .await + .expect("insert entry B"); + + // Revoke only entry A by its primary key id. + conn.transaction(async |conn| { + diesel::delete(evm_wallet_access::table) + .filter(evm_wallet_access::id.eq(entry_a.id)) + .execute(&mut *conn) + .await + }) + .await + .expect("revoke entry A"); + + // Entry A must be gone. + let gone = evm_wallet_access::table + .filter(evm_wallet_access::id.eq(entry_a.id)) + .count() + .get_result::(&mut *conn) + .await + .expect("count entry A"); + assert_eq!(gone, 0, "revoked entry must be deleted"); + + // Entry B (same wallet, different client) must still exist. + let still_there = evm_wallet_access::table + .filter(evm_wallet_access::id.eq(entry_b.id)) + .count() + .get_result::(&mut *conn) + .await + .expect("count entry B"); + assert_eq!(still_there, 1, "unrelated entry must not be deleted"); + } + + /// Regression test: when `entry_id` and `wallet_id` differ, only the correct row is removed. + /// + /// This specifically catches the case where `entry.id=5` and `wallet_id=1` are different values; + /// the old bug would delete by `wallet_id`, potentially matching a completely different entry. + #[tokio::test] + async fn revoke_with_mismatched_wallet_and_entry_ids() { + use crate::db::models::EvmWalletAccess; + + let pool = db::create_test_pool().await; + let mut conn = pool.get().await.expect("pool connection"); + + // Insert entries to force auto-increment IDs to diverge from wallet_ids. + // We'll insert 5 placeholder entries first so that the real entry gets id=6. + for i in 1_i32..=5 { + diesel::insert_into(evm_wallet_access::table) + .values(NewEvmWalletAccess { + wallet_id: 99, + client_id: i, + }) + .execute(&mut *conn) + .await + .expect("insert placeholder"); + } + + // Real target: wallet_id=1, will get id=6. + let target = diesel::insert_into(evm_wallet_access::table) + .values(NewEvmWalletAccess { + wallet_id: 1, + client_id: 1, + }) + .returning(EvmWalletAccess::as_select()) + .get_result(&mut *conn) + .await + .expect("insert target"); + + // Sanity: target.id != target.wallet_id + assert_ne!( + target.id, target.wallet_id, + "test prerequisite: id and wallet_id must differ" + ); + + // Revoke by entry id. + conn.transaction(async |conn| { + diesel::delete(evm_wallet_access::table) + .filter(evm_wallet_access::id.eq(target.id)) + .execute(&mut *conn) + .await + }) + .await + .expect("revoke target"); + + let remaining = evm_wallet_access::table + .filter(evm_wallet_access::id.eq(target.id)) + .count() + .get_result::(&mut *conn) + .await + .expect("count target"); + assert_eq!(remaining, 0, "target must be deleted by its entry id"); + + // Placeholders for wallet_id=99 must be untouched. + let placeholders = evm_wallet_access::table + .filter(evm_wallet_access::wallet_id.eq(99)) + .count() + .get_result::(&mut *conn) + .await + .expect("count placeholders"); + assert_eq!(placeholders, 5, "unrelated entries must survive"); + } +} diff --git a/server/crates/arbiter-server/src/peers/operator/session/mod.rs b/server/crates/arbiter-server/src/peers/operator/session/mod.rs index d7566ac..bae29d0 100644 --- a/server/crates/arbiter-server/src/peers/operator/session/mod.rs +++ b/server/crates/arbiter-server/src/peers/operator/session/mod.rs @@ -1,16 +1,14 @@ use super::{OutOfBand, OperatorConnection}; use crate::{ actors::{ - flow_coordinator::client_connect_approval::{ClientApprovalAnswer, ClientApprovalController}, - operator_registry::ConnectOperator, - }, - peers::client::ClientProfile, + flow_coordinator::{GetConnectedClientIds, client_connect_approval::{ClientApprovalAnswer, ClientApprovalController}}, operator_registry::ConnectOperator, + }, peers::client::ClientProfile, }; use arbiter_crypto::authn; use arbiter_proto::transport::Sender; use kameo::{Actor, actor::ActorRef, messages}; -use std::{borrow::Cow, collections::HashMap}; +use std::{borrow::Cow, collections::{HashMap, HashSet}}; use thiserror::Error; use tracing::error; @@ -54,6 +52,10 @@ pub struct OperatorSession { sender: Box>, pending_client_approvals: HashMap, PendingClientApproval>, + /// DB `client_ids` this operator session is allowed to sign for. + /// Seeded from currently-connected clients on start, then updated as + /// approvals are granted or denied during the session lifetime. + approved_client_ids: HashSet, } pub mod handlers; @@ -64,6 +66,7 @@ impl OperatorSession { props, sender, pending_client_approvals: HashMap::default(), + approved_client_ids: HashSet::default(), } } } @@ -107,7 +110,7 @@ impl Actor for OperatorSession { type Error = Error; - async fn on_start(args: Self::Args, this: ActorRef) -> Result { + async fn on_start(mut args: Self::Args, this: ActorRef) -> Result { args.props .actors .operator_registry @@ -122,6 +125,16 @@ impl Actor for OperatorSession { ); Error::internal("Failed to register operator connection with operator registry") })?; + + // Seed approved set with clients already connected when this session starts. + // New clients will be added via handle_new_client_approve as they are approved. + match args.props.actors.flow_coordinator.ask(GetConnectedClientIds {}).await { + Ok(ids) => args.approved_client_ids.extend(ids), + Err(err) => { + error!(?err, "Failed to fetch connected client IDs on operator session start"); + } + } + Ok(args) } diff --git a/server/crates/arbiter-server/src/peers/operator/vault_gate/mod.rs b/server/crates/arbiter-server/src/peers/operator/vault_gate/mod.rs index 6a8a265..d0e38e3 100644 --- a/server/crates/arbiter-server/src/peers/operator/vault_gate/mod.rs +++ b/server/crates/arbiter-server/src/peers/operator/vault_gate/mod.rs @@ -25,6 +25,8 @@ pub enum Error { AlreadyBootstrapped, #[error("Invalid key provided")] InvalidKey, + #[error("Vault locked: too many failed unseal attempts")] + LockedOut, #[error("State transition failed")] State, @@ -170,6 +172,7 @@ impl VaultGate { Ok(()) } Err(SendError::HandlerError(vault::Error::InvalidKey)) => Err(Error::InvalidKey), + Err(SendError::HandlerError(vault::Error::LockedOut)) => Err(Error::LockedOut), Err(SendError::HandlerError(err)) => { error!(?err, "Vault failed to unseal key"); Err(Error::InvalidKey) diff --git a/server/crates/arbiter-server/tests/client/auth.rs b/server/crates/arbiter-server/tests/client/auth.rs index 90763a8..b753427 100644 --- a/server/crates/arbiter-server/tests/client/auth.rs +++ b/server/crates/arbiter-server/tests/client/auth.rs @@ -266,7 +266,7 @@ pub async fn metadata_unchanged_does_not_append_history() { #[tokio::test] #[test_log::test] -pub async fn metadata_change_appends_history_and_repoints_binding() { +pub async fn metadata_frozen_after_approval_ignores_reconnect_changes() { let db = db::create_test_pool().await; let actors = spawn_test_actors(&db).await; let new_key = MlDsa87::key_gen(&mut rand::rng()); @@ -287,6 +287,7 @@ pub async fn metadata_change_appends_history_and_repoints_binding() { connect_client(props, &mut server_transport).await; }); + // Reconnect presenting different metadata — must be silently ignored. test_transport .send(auth::Inbound::AuthChallengeRequest { pubkey: verifying_key(&new_key).into(), @@ -313,6 +314,7 @@ pub async fn metadata_change_appends_history_and_repoints_binding() { client_metadata, client_metadata_history, program_client, }; let mut conn = db.get().await.unwrap(); + // Metadata is frozen: no new row, no history entry. let metadata_count: i64 = client_metadata::table .count() .get_result(&mut conn) @@ -338,15 +340,16 @@ pub async fn metadata_change_appends_history_and_repoints_binding() { .first::<(String, Option, Option)>(&mut conn) .await .unwrap(); - assert_eq!(metadata_count, 2); - assert_eq!(history_count, 1); + assert_eq!(metadata_count, 1, "frozen: no new metadata row on reconnect"); + assert_eq!(history_count, 0, "frozen: no history entry on reconnect"); assert_eq!( current, ( "client".to_owned(), - Some("new".to_owned()), - Some("2.0.0".to_owned()) - ) + Some("old".to_owned()), + Some("1.0.0".to_owned()) + ), + "frozen: original metadata must be preserved" ); } } diff --git a/server/crates/arbiter-server/tests/operator/auth.rs b/server/crates/arbiter-server/tests/operator/auth.rs index cc1f8f3..aa38bf9 100644 --- a/server/crates/arbiter-server/tests/operator/auth.rs +++ b/server/crates/arbiter-server/tests/operator/auth.rs @@ -174,7 +174,7 @@ pub async fn bootstrap_token_auth() { test_transport .send(auth::Inbound::AuthChallengeRequest { pubkey: verifying_key(&new_key).into(), - bootstrap_token: Some(token), + bootstrap_token: Some(token.into_bytes()), }) .await .unwrap(); @@ -231,7 +231,7 @@ pub async fn bootstrap_invalid_token_auth() { test_transport .send(auth::Inbound::AuthChallengeRequest { pubkey: verifying_key(&new_key).into(), - bootstrap_token: Some("invalid_token".to_owned()), + bootstrap_token: Some(b"invalid_token".to_vec()), }) .await .unwrap(); diff --git a/server/crates/arbiter-server/tests/vault/concurrency.rs b/server/crates/arbiter-server/tests/vault/concurrency.rs index ee84f4a..ee77d19 100644 --- a/server/crates/arbiter-server/tests/vault/concurrency.rs +++ b/server/crates/arbiter-server/tests/vault/concurrency.rs @@ -14,6 +14,8 @@ use kameo::actor::{ActorRef, Spawn as _}; use std::collections::{HashMap, HashSet}; use tokio::task::JoinSet; +const TEST_AAD: &[u8] = b"test-aad"; + async fn write_concurrently( actor: ActorRef, prefix: &'static str, @@ -27,6 +29,7 @@ async fn write_concurrently( let id = actor .ask(CreateNew { plaintext: SafeCell::new(plaintext.clone()), + aad: TEST_AAD.to_vec(), }) .await .unwrap(); @@ -120,7 +123,7 @@ async fn insert_failure_does_not_create_partial_row() { drop(conn); let err = actor - .create_new(SafeCell::new(b"should fail".to_vec())) + .create_new(SafeCell::new(b"should fail".to_vec()), TEST_AAD.to_vec()) .await .unwrap_err(); assert!(matches!(err, Error::DatabaseTransaction(_))); @@ -171,7 +174,7 @@ async fn decrypt_roundtrip_after_high_concurrency() { .unwrap(); for (id, plaintext) in expected { - let mut decrypted = decryptor.decrypt(id).await.unwrap(); + let mut decrypted = decryptor.decrypt(id, TEST_AAD.to_vec()).await.unwrap(); assert_eq!(*decrypted.read(), plaintext); } } diff --git a/server/crates/arbiter-server/tests/vault/lifecycle.rs b/server/crates/arbiter-server/tests/vault/lifecycle.rs index c4ee7da..c7096ee 100644 --- a/server/crates/arbiter-server/tests/vault/lifecycle.rs +++ b/server/crates/arbiter-server/tests/vault/lifecycle.rs @@ -12,6 +12,8 @@ use arbiter_server::{ use diesel::{QueryDsl, SelectableHelper}; use diesel_async::RunQueryDsl; +const TEST_AAD: &[u8] = b"test-aad"; + #[tokio::test] #[test_log::test] async fn bootstrap() { @@ -57,7 +59,7 @@ async fn create_new_before_bootstrap_fails() { .unwrap(); let err = actor - .create_new(SafeCell::new(b"data".to_vec())) + .create_new(SafeCell::new(b"data".to_vec()), TEST_AAD.to_vec()) .await .unwrap_err(); assert!(matches!(err, Error::NotBootstrapped)); @@ -71,7 +73,7 @@ async fn decrypt_before_bootstrap_fails() { .await .unwrap(); - let err = actor.decrypt(1).await.unwrap_err(); + let err = actor.decrypt(1, TEST_AAD.to_vec()).await.unwrap_err(); assert!(matches!(err, Error::NotBootstrapped)); } @@ -85,7 +87,7 @@ async fn new_restores_sealed_state() { let mut actor2 = Vault::new(db, GlobalActors::spawn_message_bus()) .await .unwrap(); - let err = actor2.decrypt(1).await.unwrap_err(); + let err = actor2.decrypt(1, TEST_AAD.to_vec()).await.unwrap_err(); assert!(matches!(err, Error::Sealed)); } @@ -97,7 +99,7 @@ async fn unseal_correct_password() { let plaintext = b"survive a restart"; let aead_id = actor - .create_new(SafeCell::new(plaintext.to_vec())) + .create_new(SafeCell::new(plaintext.to_vec()), TEST_AAD.to_vec()) .await .unwrap(); drop(actor); @@ -108,7 +110,7 @@ async fn unseal_correct_password() { let seal_key = SafeCell::new(b"test-seal-key".to_vec()); actor.try_unseal(seal_key).await.unwrap(); - let mut decrypted = actor.decrypt(aead_id).await.unwrap(); + let mut decrypted = actor.decrypt(aead_id, TEST_AAD.to_vec()).await.unwrap(); assert_eq!(*decrypted.read(), plaintext); } @@ -120,7 +122,7 @@ async fn unseal_wrong_then_correct_password() { let plaintext = b"important data"; let aead_id = actor - .create_new(SafeCell::new(plaintext.to_vec())) + .create_new(SafeCell::new(plaintext.to_vec()), TEST_AAD.to_vec()) .await .unwrap(); drop(actor); @@ -136,6 +138,6 @@ async fn unseal_wrong_then_correct_password() { let good_key = SafeCell::new(b"test-seal-key".to_vec()); actor.try_unseal(good_key).await.unwrap(); - let mut decrypted = actor.decrypt(aead_id).await.unwrap(); + let mut decrypted = actor.decrypt(aead_id, TEST_AAD.to_vec()).await.unwrap(); assert_eq!(*decrypted.read(), plaintext); } diff --git a/server/crates/arbiter-server/tests/vault/storage.rs b/server/crates/arbiter-server/tests/vault/storage.rs index c1bd321..ed1cc8a 100644 --- a/server/crates/arbiter-server/tests/vault/storage.rs +++ b/server/crates/arbiter-server/tests/vault/storage.rs @@ -10,6 +10,8 @@ use diesel::{ExpressionMethods as _, QueryDsl, SelectableHelper, dsl::update}; use diesel_async::RunQueryDsl; use std::collections::HashSet; +const TEST_AAD: &[u8] = b"test-aad"; + #[tokio::test] #[test_log::test] async fn create_decrypt_roundtrip() { @@ -18,11 +20,11 @@ async fn create_decrypt_roundtrip() { let plaintext = b"hello arbiter"; let aead_id = actor - .create_new(SafeCell::new(plaintext.to_vec())) + .create_new(SafeCell::new(plaintext.to_vec()), TEST_AAD.to_vec()) .await .unwrap(); - let mut decrypted = actor.decrypt(aead_id).await.unwrap(); + let mut decrypted = actor.decrypt(aead_id, TEST_AAD.to_vec()).await.unwrap(); assert_eq!(*decrypted.read(), plaintext); } @@ -32,7 +34,7 @@ async fn decrypt_nonexistent_returns_not_found() { let db = db::create_test_pool().await; let mut actor = common::bootstrapped_vault(&db).await; - let err = actor.decrypt(9999).await.unwrap_err(); + let err = actor.decrypt(9999, TEST_AAD.to_vec()).await.unwrap_err(); assert!(matches!(err, Error::NotFound)); } @@ -44,11 +46,11 @@ async fn ciphertext_differs_across_entries() { let plaintext = b"same content"; let id1 = actor - .create_new(SafeCell::new(plaintext.to_vec())) + .create_new(SafeCell::new(plaintext.to_vec()), TEST_AAD.to_vec()) .await .unwrap(); let id2 = actor - .create_new(SafeCell::new(plaintext.to_vec())) + .create_new(SafeCell::new(plaintext.to_vec()), TEST_AAD.to_vec()) .await .unwrap(); @@ -68,8 +70,8 @@ async fn ciphertext_differs_across_entries() { assert_ne!(row1.ciphertext, row2.ciphertext); - let mut d1 = actor.decrypt(id1).await.unwrap(); - let mut d2 = actor.decrypt(id2).await.unwrap(); + let mut d1 = actor.decrypt(id1, TEST_AAD.to_vec()).await.unwrap(); + let mut d2 = actor.decrypt(id2, TEST_AAD.to_vec()).await.unwrap(); assert_eq!(*d1.read(), plaintext); assert_eq!(*d2.read(), plaintext); } @@ -83,7 +85,7 @@ async fn nonce_never_reused() { let n = 5; for i in 0..n { actor - .create_new(SafeCell::new(format!("secret {i}").into_bytes())) + .create_new(SafeCell::new(format!("secret {i}").into_bytes()), TEST_AAD.to_vec()) .await .unwrap(); } @@ -137,7 +139,7 @@ async fn broken_db_nonce_format_fails_closed() { drop(conn); let err = actor - .create_new(SafeCell::new(b"must fail".to_vec())) + .create_new(SafeCell::new(b"must fail".to_vec()), TEST_AAD.to_vec()) .await .unwrap_err(); assert!(matches!(err, Error::BrokenDatabase)); @@ -145,7 +147,7 @@ async fn broken_db_nonce_format_fails_closed() { let db = db::create_test_pool().await; let mut actor = common::bootstrapped_vault(&db).await; let id = actor - .create_new(SafeCell::new(b"decrypt target".to_vec())) + .create_new(SafeCell::new(b"decrypt target".to_vec()), TEST_AAD.to_vec()) .await .unwrap(); let mut conn = db.get().await.unwrap(); @@ -156,6 +158,6 @@ async fn broken_db_nonce_format_fails_closed() { .unwrap(); drop(conn); - let err = actor.decrypt(id).await.unwrap_err(); + let err = actor.decrypt(id, TEST_AAD.to_vec()).await.unwrap_err(); assert!(matches!(err, Error::BrokenDatabase)); }