diff --git a/rust/networking/src/reqwest.rs b/rust/networking/src/reqwest.rs index ce2a32e..dc4f373 100644 --- a/rust/networking/src/reqwest.rs +++ b/rust/networking/src/reqwest.rs @@ -1,136 +1,176 @@ -//! An [`http::Client`] implementation that utilizes [`reqwest`]. - -use ::http::{HeaderName, HeaderValue}; -use async_trait::async_trait; -use reqwest::{Certificate, RequestBuilder}; -use std::collections::HashMap; -use std::str::FromStr; -use std::time::Duration; -use tracing::warn; - -use crate::http; - -/// Options for configuring the [`reqwest`] [`Client`]. -#[derive(Debug, Clone)] -pub struct ClientOptions<'a> { - pub additional_root_certs: Vec, - pub timeout: Duration, - pub default_headers: HashMap<&'a str, &'a str>, -} - -impl<'a> Default for ClientOptions<'a> { - fn default() -> Self { - Self { - additional_root_certs: Vec::new(), - timeout: Duration::from_secs(30), - default_headers: HashMap::from([( - "User-Agent", - concat!(env!("CARGO_PKG_NAME"), "/", env!("CARGO_PKG_VERSION")), - )]), - } - } -} - -/// An [`http::Client`] implementation that utilizes [`reqwest`]. -#[derive(Clone, Debug, Default)] -pub struct Client { - // reqwest::Client holds a connection pool. It's reference-counted - // internally, so this field is relatively cheap to clone. - http: reqwest::Client, -} - -impl Client { - pub fn new(options: ClientOptions) -> Self { - let mut b = reqwest::Client::builder() - .timeout(options.timeout) - // The service checker needs access to the server's certificate to - // warn if it will expire soon. - .tls_info(true) - .use_rustls_tls(); - - let mut default_headers = reqwest::header::HeaderMap::new(); - for (key, value) in options.default_headers { - if let (Ok(header_name), Ok(header_value)) = - (HeaderName::from_str(key), HeaderValue::from_str(value)) - { - default_headers.append(header_name, header_value); - } - } - b = b.default_headers(default_headers); - - for c in options.additional_root_certs { - b = b.add_root_certificate(c); - } - Self { - http: b.build().expect("TODO"), - } - } - - pub fn to_reqwest(&self, request: http::Request) -> RequestBuilder { - let mut request_builder = match request.method { - http::Method::Get => self.http.get(request.url), - http::Method::Put => self.http.put(request.url), - http::Method::Post => self.http.post(request.url), - http::Method::Delete => self.http.delete(request.url), - }; - - let mut headers = reqwest::header::HeaderMap::new(); - for (key, value) in request.headers { - if let (Ok(header_name), Ok(header_value)) = - (HeaderName::from_str(&key), HeaderValue::from_str(&value)) - { - headers.append(header_name, header_value); - } - } - request_builder = request_builder.headers(headers); - - if let Some(body) = request.body { - request_builder = request_builder.body(body); - } - - if let Some(timeout) = request.timeout { - request_builder = request_builder.timeout(timeout); - } - request_builder - } - - pub async fn to_response( - &self, - resp: Result, - ) -> Result { - match resp { - Err(err) => { - warn!(%err, "error sending HTTP request"); - Err(err) - } - Ok(response) => { - let status = response.status().as_u16(); - let mut headers = HashMap::new(); - for (header_name, header_value) in response.headers() { - if let Ok(value) = header_value.to_str() { - headers.insert(header_name.to_string(), value.to_owned()); - } - } - match response.bytes().await { - Err(err) => { - warn!(%err, "error receiving HTTP response"); - Err(err) - } - Ok(bytes) => Ok(http::Response { - status_code: status, - headers, - body: bytes.to_vec(), - }), - } - } - } - } -} - -#[async_trait] -impl http::Client for Client { - async fn send(&self, request: http::Request) -> Option { - let resp = self.to_reqwest(request).send().await; - self.to_response(resp).await.ok() - } -} +//! An [`http::Client`] implementation that utilizes [`reqwest`]. + +use ::http::{HeaderName, HeaderValue}; +use async_trait::async_trait; +use reqwest::{Certificate, RequestBuilder}; +use std::collections::HashMap; +use std::str::FromStr; +use std::time::Duration; +use tracing::warn; + +use crate::http; + +/// Options for configuring the [`reqwest`] [`Client`]. +#[derive(Debug, Clone)] +pub struct ClientOptions<'a> { + pub additional_root_certs: Vec, + pub timeout: Duration, + pub default_headers: HashMap<&'a str, &'a str>, +} + +/// Upper bound on the accepted HTTP response body size, in bytes. +/// +/// A misbehaving or compromised realm can stream an unbounded body and +/// otherwise exhaust the client's memory; the request timeout bounds duration, +/// not size. This cap is enforced while streaming, before the full body is +/// buffered. +const MAX_RESPONSE_BODY_BYTES: usize = 16 * 1024 * 1024; + +/// Error variants produced while accumulating a capped response body. +#[derive(Debug, thiserror::Error)] +enum ResponseBodyError { + #[error("HTTP response body too large: {0} bytes (limit {MAX_RESPONSE_BODY_BYTES})")] + TooLarge(usize), + #[error("error receiving HTTP response body: {0}")] + Receive(#[from] reqwest::Error), +} + +/// Streams the response body into memory, aborting once the size cap is hit. +async fn read_body_capped(mut response: reqwest::Response) -> Result, ResponseBodyError> { + if let Some(length) = response.content_length() { + if length as usize > MAX_RESPONSE_BODY_BYTES { + return Err(ResponseBodyError::TooLarge(length as usize)); + } + } + + let capacity = usize::try_from(response.content_length().unwrap_or(0)) + .unwrap_or(0) + .min(MAX_RESPONSE_BODY_BYTES); + let mut body = Vec::with_capacity(capacity); + + while let Some(chunk) = response.chunk().await? { + body.extend_from_slice(&chunk); + if body.len() > MAX_RESPONSE_BODY_BYTES { + return Err(ResponseBodyError::TooLarge(body.len())); + } + } + + Ok(body) +} + +impl<'a> Default for ClientOptions<'a> { + fn default() -> Self { + Self { + additional_root_certs: Vec::new(), + timeout: Duration::from_secs(30), + default_headers: HashMap::from([( + "User-Agent", + concat!(env!("CARGO_PKG_NAME"), "/", env!("CARGO_PKG_VERSION")), + )]), + } + } +} + +/// An [`http::Client`] implementation that utilizes [`reqwest`]. +#[derive(Clone, Debug, Default)] +pub struct Client { + // reqwest::Client holds a connection pool. It's reference-counted + // internally, so this field is relatively cheap to clone. + http: reqwest::Client, +} + +impl Client { + pub fn new(options: ClientOptions) -> Self { + let mut b = reqwest::Client::builder() + .timeout(options.timeout) + // The service checker needs access to the server's certificate to + // warn if it will expire soon. + .tls_info(true) + .use_rustls_tls(); + + let mut default_headers = reqwest::header::HeaderMap::new(); + for (key, value) in options.default_headers { + if let (Ok(header_name), Ok(header_value)) = + (HeaderName::from_str(key), HeaderValue::from_str(value)) + { + default_headers.append(header_name, header_value); + } + } + b = b.default_headers(default_headers); + + for c in options.additional_root_certs { + b = b.add_root_certificate(c); + } + Self { + http: b.build().expect("TODO"), + } + } + + pub fn to_reqwest(&self, request: http::Request) -> RequestBuilder { + let mut request_builder = match request.method { + http::Method::Get => self.http.get(request.url), + http::Method::Put => self.http.put(request.url), + http::Method::Post => self.http.post(request.url), + http::Method::Delete => self.http.delete(request.url), + }; + + let mut headers = reqwest::header::HeaderMap::new(); + for (key, value) in request.headers { + if let (Ok(header_name), Ok(header_value)) = + (HeaderName::from_str(&key), HeaderValue::from_str(&value)) + { + headers.append(header_name, header_value); + } + } + request_builder = request_builder.headers(headers); + + if let Some(body) = request.body { + request_builder = request_builder.body(body); + } + + if let Some(timeout) = request.timeout { + request_builder = request_builder.timeout(timeout); + } + request_builder + } + + pub async fn to_response( + &self, + resp: Result, + ) -> Option { + match resp { + Err(err) => { + warn!(%err, "error sending HTTP request"); + None + } + Ok(response) => { + let status = response.status().as_u16(); + let mut headers = HashMap::new(); + for (header_name, header_value) in response.headers() { + if let Ok(value) = header_value.to_str() { + headers.insert(header_name.to_string(), value.to_owned()); + } + } + match read_body_capped(response).await { + Err(err) => { + warn!(%err, "error receiving HTTP response"); + None + } + Ok(bytes) => Some(http::Response { + status_code: status, + headers, + body: bytes, + }), + } + } + } + } +} + +#[async_trait] +impl http::Client for Client { + async fn send(&self, request: http::Request) -> Option { + let resp = self.to_reqwest(request).send().await; + self.to_response(resp).await + } +} diff --git a/rust/realm/api/src/requests.rs b/rust/realm/api/src/requests.rs index 5778e95..c245728 100644 --- a/rust/realm/api/src/requests.rs +++ b/rust/realm/api/src/requests.rs @@ -1,338 +1,345 @@ -extern crate alloc; - -use alloc::boxed::Box; -use alloc::vec::Vec; -use core::fmt; -use core::time::Duration; -use serde::{Deserialize, Serialize}; - -use crate::signing::OprfSignedPublicKey; -use crate::types::{ - AuthToken, EncryptedUserSecret, EncryptedUserSecretCommitment, Policy, RealmId, - RegistrationVersion, SecretBytesArray, SessionId, UnlockKeyCommitment, UnlockKeyTag, - UserSecretEncryptionKeyScalarShare, -}; -use juicebox_marshalling::{self as marshalling, bytes, DeserializationError, SerializationError}; -use juicebox_noise as noise; -use juicebox_oprf as oprf; - -#[derive(Clone, Debug, Deserialize, Serialize)] -pub struct ClientRequest { - pub realm: RealmId, - pub auth_token: AuthToken, - pub session_id: SessionId, - pub kind: ClientRequestKind, - pub encrypted: NoiseRequest, -} - -/// Used in [`ClientRequest`]. -#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] -pub enum ClientRequestKind { - /// The [`ClientRequest`] contains just a Noise handshake request, without - /// a [`SecretsRequest`]. The server does not need to access the user's - /// record to process this. - HandshakeOnly, - /// The [`ClientRequest`] contains a Noise handshake or transport request - /// with an encrypted [`SecretsRequest`]. The server will need to access - /// the user's record to process this. - SecretsRequest, -} - -#[derive(Debug, Deserialize, Serialize)] -#[allow(clippy::large_enum_variant)] -pub enum ClientResponse { - Ok(NoiseResponse), - /// The appropriate server to handle the request is not currently - /// available. - Unavailable, - /// The request's auth token is not acceptable. - InvalidAuth, - /// The server could not find the Noise session state referenced by the - /// request's session ID. This can occur in normal circumstances when a - /// server restarts or has expired the session. The client should open a - /// new session. - MissingSession, - /// The server could not decrypt the encapsulated Noise request. - SessionError, - // The server could not deserialize the [`ClientRequest`] or the - // encapsulated [`SecretsRequest`]. - DecodingError, - /// The payload sent to the server was too large to be processed. - PayloadTooLarge, - /// The tenant has exceeded their allowed number of operations. Try again - /// later. - RateLimitExceeded, -} - -/// A Noise protocol handshake or transport message. -#[derive(Clone, Deserialize, Serialize)] -pub enum NoiseRequest { - Handshake { - handshake: noise::HandshakeRequest, - }, - Transport { - #[serde(with = "bytes")] - ciphertext: Vec, - }, -} - -impl fmt::Debug for NoiseRequest { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - Self::Handshake { .. } => f - .debug_struct("NoiseRequest::Handshake") - .finish_non_exhaustive(), - Self::Transport { .. } => f - .debug_struct("NoiseRequest::Transport") - .finish_non_exhaustive(), - } - } -} - -#[derive(Deserialize, Serialize)] -pub enum NoiseResponse { - Handshake { - handshake: noise::HandshakeResponse, - /// Once the session becomes inactive for this long, the client should - /// discard the session. - session_lifetime: Duration, - }, - Transport { - #[serde(with = "bytes")] - ciphertext: Vec, - }, -} - -impl fmt::Debug for NoiseResponse { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - Self::Handshake { - session_lifetime, .. - } => f - .debug_struct("NoiseResponse::Handshake") - .field("session_lifetime", &session_lifetime) - .finish_non_exhaustive(), - Self::Transport { .. } => f - .debug_struct("NoiseResponse::Transport") - .finish_non_exhaustive(), - } - } -} - -#[derive(Clone, Debug, Deserialize, Serialize)] -pub enum SecretsRequest { - Register1, - Register2(Box), - Recover1, - Recover2(Recover2Request), - Recover3(Recover3Request), - Delete, -} - -impl SecretsRequest { - /// Returns whether the request type requires forward secrecy. - /// - /// This controls whether the request may be sent as part of a Noise NK - /// handshake request, which does not provide forward secrecy. - /// - /// For more sensitive request types, this returns true, requiring an - /// established Noise session before the request can be sent. Decrypting - /// these requests would require both the server/realm's static secret key - /// and the ephemeral key used only for this session. - /// - /// For less sensitive request types, this returns false, indicating that - /// they can be piggy-backed with the Noise NK handshake request. - /// Decrypting these requests would be possible with just the - /// server/realm's static secret key (even any time in the future). - pub fn needs_forward_secrecy(&self) -> bool { - match self { - Self::Register1 => false, - Self::Register2(_) => true, - Self::Recover1 => false, - Self::Recover2(_) => true, - Self::Recover3(_) => true, - Self::Delete => false, - } - } -} - -#[derive(Debug, Deserialize, Serialize)] -#[allow(clippy::large_enum_variant)] -pub enum SecretsResponse { - Register1(Register1Response), - Register2(Register2Response), - Recover1(Recover1Response), - Recover2(Recover2Response), - Recover3(Recover3Response), - Delete(DeleteResponse), -} - -const MAX_SECRETS_RESPONSE_LENGTH: usize = 436; - -/// A padded representation of a [`SecretsResponse`]. -#[derive(Debug, Deserialize, Serialize)] -pub struct PaddedSecretsResponse { - pub unpadded_length: u16, - pub padded_bytes: SecretBytesArray, -} - -impl TryFrom<&SecretsResponse> for PaddedSecretsResponse { - type Error = SerializationError; - - fn try_from(value: &SecretsResponse) -> Result { - let mut padded_response = marshalling::to_vec(value)?; - assert!(padded_response.len() <= MAX_SECRETS_RESPONSE_LENGTH); - let unpadded_length = padded_response - .len() - .try_into() - .expect("padded length unexpectedly large"); - padded_response.resize(MAX_SECRETS_RESPONSE_LENGTH, 0); - Ok(Self { - unpadded_length, - padded_bytes: SecretBytesArray::try_from(padded_response) - .expect("padded_response unexpectedly wrong length"), - }) - } -} - -impl TryFrom<&PaddedSecretsResponse> for SecretsResponse { - type Error = DeserializationError; - - fn try_from(value: &PaddedSecretsResponse) -> Result { - marshalling::from_slice( - &value.padded_bytes.expose_secret()[..usize::from(value.unpadded_length)], - ) - } -} - -/// Response message for the first phase of registration. -#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] -pub enum Register1Response { - Ok, -} - -/// Request message for the second phase of registration. -#[derive(Clone, Debug, Deserialize, Serialize)] -pub struct Register2Request { - pub version: RegistrationVersion, - pub oprf_private_key: oprf::PrivateKey, - pub oprf_signed_public_key: OprfSignedPublicKey, - pub unlock_key_commitment: UnlockKeyCommitment, - pub unlock_key_tag: UnlockKeyTag, - pub encryption_key_scalar_share: UserSecretEncryptionKeyScalarShare, - pub encrypted_secret: EncryptedUserSecret, - pub encrypted_secret_commitment: EncryptedUserSecretCommitment, - pub policy: Policy, -} - -/// Response message for the second phase of registration. -#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] -pub enum Register2Response { - Ok, -} - -/// Response message for the first phase of recovery. -#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] -pub enum Recover1Response { - Ok { version: RegistrationVersion }, - NotRegistered, - NoGuesses, -} - -/// Request message for the second phase of recovery. -#[derive(Clone, Debug, Deserialize, Serialize)] -pub struct Recover2Request { - pub version: RegistrationVersion, - pub oprf_blinded_input: oprf::BlindedInput, -} - -/// Response message for the second phase of recovery. -#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] -#[allow(clippy::large_enum_variant)] -pub enum Recover2Response { - Ok { - oprf_signed_public_key: OprfSignedPublicKey, - oprf_blinded_result: oprf::BlindedOutput, - oprf_proof: oprf::Proof, - unlock_key_commitment: UnlockKeyCommitment, - num_guesses: u16, - guess_count: u16, - }, - VersionMismatch, - NotRegistered, - NoGuesses, -} - -/// Request message for the third phase of recovery. -#[derive(Clone, Debug, Deserialize, Serialize)] -pub struct Recover3Request { - pub version: RegistrationVersion, - pub unlock_key_tag: UnlockKeyTag, -} - -/// Response message for the third phase of recovery. -#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] -pub enum Recover3Response { - Ok { - encryption_key_scalar_share: UserSecretEncryptionKeyScalarShare, - encrypted_secret: EncryptedUserSecret, - encrypted_secret_commitment: EncryptedUserSecretCommitment, - }, - VersionMismatch, - NotRegistered, - BadUnlockKeyTag { - guesses_remaining: u16, - }, - NoGuesses, -} - -/// Response message to delete registered secrets. -#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] -pub enum DeleteResponse { - Ok, -} - -/// The maximum expected request size from the SDK -pub const BODY_SIZE_LIMIT: usize = 2048; - -#[cfg(test)] -mod tests { - use crate::{ - requests::{Register2Request, SecretsRequest, BODY_SIZE_LIMIT}, - signing::{OprfSignedPublicKey, OprfVerifyingKey}, - types::{ - EncryptedUserSecret, EncryptedUserSecretCommitment, Policy, RegistrationVersion, - SecretBytesArray, UnlockKeyCommitment, UnlockKeyTag, - UserSecretEncryptionKeyScalarShare, - }, - }; - use curve25519_dalek::Scalar; - use juicebox_marshalling as marshalling; - use juicebox_oprf as oprf; - use rand_core::OsRng; - - #[test] - fn test_request_body_size_limit() { - let oprf_private_key = oprf::PrivateKey::random(&mut OsRng); - let oprf_public_key = oprf_private_key.to_public_key(); - let secrets_request = SecretsRequest::Register2(Box::new(Register2Request { - version: RegistrationVersion::from([0xff; 16]), - oprf_private_key, - oprf_signed_public_key: OprfSignedPublicKey { - public_key: oprf_public_key, - verifying_key: OprfVerifyingKey::from([0xff; 32]), - signature: SecretBytesArray::from([0xFF; 64]), - }, - unlock_key_commitment: UnlockKeyCommitment::from([0xff; 32]), - unlock_key_tag: UnlockKeyTag::from([0xff; 16]), - encryption_key_scalar_share: UserSecretEncryptionKeyScalarShare::from(-Scalar::ONE), - encrypted_secret: EncryptedUserSecret::from([0xff; 145]), - encrypted_secret_commitment: EncryptedUserSecretCommitment::from([0xff; 16]), - policy: Policy { - num_guesses: u16::MAX, - }, - })); - let serialized = marshalling::to_vec(&secrets_request).unwrap(); - assert!(serialized.len() < BODY_SIZE_LIMIT); - } -} +extern crate alloc; + +use alloc::boxed::Box; +use alloc::vec::Vec; +use core::fmt; +use core::time::Duration; +use serde::{Deserialize, Serialize}; + +use crate::signing::OprfSignedPublicKey; +use crate::types::{ + AuthToken, EncryptedUserSecret, EncryptedUserSecretCommitment, Policy, RealmId, + RegistrationVersion, SecretBytesArray, SessionId, UnlockKeyCommitment, UnlockKeyTag, + UserSecretEncryptionKeyScalarShare, +}; +use juicebox_marshalling::{self as marshalling, bytes, DeserializationError, SerializationError}; +use juicebox_noise as noise; +use juicebox_oprf as oprf; + +#[derive(Clone, Debug, Deserialize, Serialize)] +pub struct ClientRequest { + pub realm: RealmId, + pub auth_token: AuthToken, + pub session_id: SessionId, + pub kind: ClientRequestKind, + pub encrypted: NoiseRequest, +} + +/// Used in [`ClientRequest`]. +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +pub enum ClientRequestKind { + /// The [`ClientRequest`] contains just a Noise handshake request, without + /// a [`SecretsRequest`]. The server does not need to access the user's + /// record to process this. + HandshakeOnly, + /// The [`ClientRequest`] contains a Noise handshake or transport request + /// with an encrypted [`SecretsRequest`]. The server will need to access + /// the user's record to process this. + SecretsRequest, +} + +#[derive(Debug, Deserialize, Serialize)] +#[allow(clippy::large_enum_variant)] +pub enum ClientResponse { + Ok(NoiseResponse), + /// The appropriate server to handle the request is not currently + /// available. + Unavailable, + /// The request's auth token is not acceptable. + InvalidAuth, + /// The server could not find the Noise session state referenced by the + /// request's session ID. This can occur in normal circumstances when a + /// server restarts or has expired the session. The client should open a + /// new session. + MissingSession, + /// The server could not decrypt the encapsulated Noise request. + SessionError, + // The server could not deserialize the [`ClientRequest`] or the + // encapsulated [`SecretsRequest`]. + DecodingError, + /// The payload sent to the server was too large to be processed. + PayloadTooLarge, + /// The tenant has exceeded their allowed number of operations. Try again + /// later. + RateLimitExceeded, +} + +/// A Noise protocol handshake or transport message. +#[derive(Clone, Deserialize, Serialize)] +pub enum NoiseRequest { + Handshake { + handshake: noise::HandshakeRequest, + }, + Transport { + #[serde(with = "bytes")] + ciphertext: Vec, + }, +} + +impl fmt::Debug for NoiseRequest { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Handshake { .. } => f + .debug_struct("NoiseRequest::Handshake") + .finish_non_exhaustive(), + Self::Transport { .. } => f + .debug_struct("NoiseRequest::Transport") + .finish_non_exhaustive(), + } + } +} + +#[derive(Deserialize, Serialize)] +pub enum NoiseResponse { + Handshake { + handshake: noise::HandshakeResponse, + /// Once the session becomes inactive for this long, the client should + /// discard the session. + session_lifetime: Duration, + }, + Transport { + #[serde(with = "bytes")] + ciphertext: Vec, + }, +} + +impl fmt::Debug for NoiseResponse { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Handshake { + session_lifetime, .. + } => f + .debug_struct("NoiseResponse::Handshake") + .field("session_lifetime", &session_lifetime) + .finish_non_exhaustive(), + Self::Transport { .. } => f + .debug_struct("NoiseResponse::Transport") + .finish_non_exhaustive(), + } + } +} + +#[derive(Clone, Debug, Deserialize, Serialize)] +pub enum SecretsRequest { + Register1, + Register2(Box), + Recover1, + Recover2(Recover2Request), + Recover3(Recover3Request), + Delete, +} + +impl SecretsRequest { + /// Returns whether the request type requires forward secrecy. + /// + /// This controls whether the request may be sent as part of a Noise NK + /// handshake request, which does not provide forward secrecy. + /// + /// For more sensitive request types, this returns true, requiring an + /// established Noise session before the request can be sent. Decrypting + /// these requests would require both the server/realm's static secret key + /// and the ephemeral key used only for this session. + /// + /// For less sensitive request types, this returns false, indicating that + /// they can be piggy-backed with the Noise NK handshake request. + /// Decrypting these requests would be possible with just the + /// server/realm's static secret key (even any time in the future). + pub fn needs_forward_secrecy(&self) -> bool { + match self { + Self::Register1 => false, + Self::Register2(_) => true, + Self::Recover1 => false, + Self::Recover2(_) => true, + Self::Recover3(_) => true, + Self::Delete => false, + } + } +} + +#[derive(Debug, Deserialize, Serialize)] +#[allow(clippy::large_enum_variant)] +pub enum SecretsResponse { + Register1(Register1Response), + Register2(Register2Response), + Recover1(Recover1Response), + Recover2(Recover2Response), + Recover3(Recover3Response), + Delete(DeleteResponse), +} + +const MAX_SECRETS_RESPONSE_LENGTH: usize = 436; + +/// A padded representation of a [`SecretsResponse`]. +#[derive(Debug, Deserialize, Serialize)] +pub struct PaddedSecretsResponse { + pub unpadded_length: u16, + pub padded_bytes: SecretBytesArray, +} + +impl TryFrom<&SecretsResponse> for PaddedSecretsResponse { + type Error = SerializationError; + + fn try_from(value: &SecretsResponse) -> Result { + let mut padded_response = marshalling::to_vec(value)?; + assert!(padded_response.len() <= MAX_SECRETS_RESPONSE_LENGTH); + let unpadded_length = padded_response + .len() + .try_into() + .expect("padded length unexpectedly large"); + padded_response.resize(MAX_SECRETS_RESPONSE_LENGTH, 0); + Ok(Self { + unpadded_length, + padded_bytes: SecretBytesArray::try_from(padded_response) + .expect("padded_response unexpectedly wrong length"), + }) + } +} + +impl TryFrom<&PaddedSecretsResponse> for SecretsResponse { + type Error = DeserializationError; + + fn try_from(value: &PaddedSecretsResponse) -> Result { + let padded_bytes = value.padded_bytes.expose_secret(); + // The `unpadded_length` field originates from a (potentially + // compromised/malicious) realm response rather than from a trusted + // serialization, so clamp it to the actual buffer length before + // slicing. Without this, a value larger than + // MAX_SECRETS_RESPONSE_LENGTH panics on index-out-of-bounds and + // bricks every recovering client. + let unpadded_length = usize::from(value.unpadded_length).min(padded_bytes.len()); + + marshalling::from_slice(&padded_bytes[..unpadded_length]) + } +} + +/// Response message for the first phase of registration. +#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] +pub enum Register1Response { + Ok, +} + +/// Request message for the second phase of registration. +#[derive(Clone, Debug, Deserialize, Serialize)] +pub struct Register2Request { + pub version: RegistrationVersion, + pub oprf_private_key: oprf::PrivateKey, + pub oprf_signed_public_key: OprfSignedPublicKey, + pub unlock_key_commitment: UnlockKeyCommitment, + pub unlock_key_tag: UnlockKeyTag, + pub encryption_key_scalar_share: UserSecretEncryptionKeyScalarShare, + pub encrypted_secret: EncryptedUserSecret, + pub encrypted_secret_commitment: EncryptedUserSecretCommitment, + pub policy: Policy, +} + +/// Response message for the second phase of registration. +#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] +pub enum Register2Response { + Ok, +} + +/// Response message for the first phase of recovery. +#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] +pub enum Recover1Response { + Ok { version: RegistrationVersion }, + NotRegistered, + NoGuesses, +} + +/// Request message for the second phase of recovery. +#[derive(Clone, Debug, Deserialize, Serialize)] +pub struct Recover2Request { + pub version: RegistrationVersion, + pub oprf_blinded_input: oprf::BlindedInput, +} + +/// Response message for the second phase of recovery. +#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] +#[allow(clippy::large_enum_variant)] +pub enum Recover2Response { + Ok { + oprf_signed_public_key: OprfSignedPublicKey, + oprf_blinded_result: oprf::BlindedOutput, + oprf_proof: oprf::Proof, + unlock_key_commitment: UnlockKeyCommitment, + num_guesses: u16, + guess_count: u16, + }, + VersionMismatch, + NotRegistered, + NoGuesses, +} + +/// Request message for the third phase of recovery. +#[derive(Clone, Debug, Deserialize, Serialize)] +pub struct Recover3Request { + pub version: RegistrationVersion, + pub unlock_key_tag: UnlockKeyTag, +} + +/// Response message for the third phase of recovery. +#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] +pub enum Recover3Response { + Ok { + encryption_key_scalar_share: UserSecretEncryptionKeyScalarShare, + encrypted_secret: EncryptedUserSecret, + encrypted_secret_commitment: EncryptedUserSecretCommitment, + }, + VersionMismatch, + NotRegistered, + BadUnlockKeyTag { + guesses_remaining: u16, + }, + NoGuesses, +} + +/// Response message to delete registered secrets. +#[derive(Debug, Deserialize, Eq, PartialEq, Serialize)] +pub enum DeleteResponse { + Ok, +} + +/// The maximum expected request size from the SDK +pub const BODY_SIZE_LIMIT: usize = 2048; + +#[cfg(test)] +mod tests { + use crate::{ + requests::{Register2Request, SecretsRequest, BODY_SIZE_LIMIT}, + signing::{OprfSignedPublicKey, OprfVerifyingKey}, + types::{ + EncryptedUserSecret, EncryptedUserSecretCommitment, Policy, RegistrationVersion, + SecretBytesArray, UnlockKeyCommitment, UnlockKeyTag, + UserSecretEncryptionKeyScalarShare, + }, + }; + use curve25519_dalek::Scalar; + use juicebox_marshalling as marshalling; + use juicebox_oprf as oprf; + use rand_core::OsRng; + + #[test] + fn test_request_body_size_limit() { + let oprf_private_key = oprf::PrivateKey::random(&mut OsRng); + let oprf_public_key = oprf_private_key.to_public_key(); + let secrets_request = SecretsRequest::Register2(Box::new(Register2Request { + version: RegistrationVersion::from([0xff; 16]), + oprf_private_key, + oprf_signed_public_key: OprfSignedPublicKey { + public_key: oprf_public_key, + verifying_key: OprfVerifyingKey::from([0xff; 32]), + signature: SecretBytesArray::from([0xFF; 64]), + }, + unlock_key_commitment: UnlockKeyCommitment::from([0xff; 32]), + unlock_key_tag: UnlockKeyTag::from([0xff; 16]), + encryption_key_scalar_share: UserSecretEncryptionKeyScalarShare::from(-Scalar::ONE), + encrypted_secret: EncryptedUserSecret::from([0xff; 145]), + encrypted_secret_commitment: EncryptedUserSecretCommitment::from([0xff; 16]), + policy: Policy { + num_guesses: u16::MAX, + }, + })); + let serialized = marshalling::to_vec(&secrets_request).unwrap(); + assert!(serialized.len() < BODY_SIZE_LIMIT); + } +} diff --git a/rust/sdk/src/configuration.rs b/rust/sdk/src/configuration.rs index 940e9c4..8997dc0 100644 --- a/rust/sdk/src/configuration.rs +++ b/rust/sdk/src/configuration.rs @@ -1,176 +1,182 @@ -use serde::{Deserialize, Serialize}; -use std::{collections::HashSet, ops::Deref}; - -use crate::{PinHashingMode, Realm}; -use juicebox_realm_api::types::RealmId; -use juicebox_secret_sharing::Index; - -/// The parameters used to configure a [`Client`](crate::Client). -#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] -pub struct Configuration { - /// The remote services that the client interacts with. - /// - /// There must be between `register_threshold` and 255 realms, inclusive. - pub realms: Vec, - - /// A registration will be considered successful if it's successful on at - /// least this many realms. - /// - /// Must be between `recover_threshold` and `realms.len()`, inclusive. - pub register_threshold: u32, - - /// A recovery (or an adversary) will need the cooperation of this many - /// realms to retrieve the secret. - /// - /// Must be between `(realms.len() / 2).ceil()` and `realms.len()`, inclusive. - pub recover_threshold: u32, - - /// Defines how the provided PIN will be hashed before register and recover - /// operations. Changing modes will make previous secrets stored on the realms - /// inaccessible with the same PIN and should not be done without re-registering - /// secrets. - pub pin_hashing_mode: PinHashingMode, -} - -impl Configuration { - pub fn from_json(s: &str) -> Result { - serde_json::from_str(s) - } - - pub fn to_json(&self) -> String { - serde_json::to_string_pretty(self).expect("failed to convert configuration to json") - } -} - -#[derive(Debug)] -pub(crate) struct CheckedConfiguration(Configuration); - -impl CheckedConfiguration { - pub fn from(c: Configuration) -> Self { - assert!( - !c.realms.is_empty(), - "Client needs at least one realm in Configuration" - ); - - assert_eq!( - c.realms - .iter() - .map(|realm| realm.id) - .collect::>() - .len(), - c.realms.len(), - "realm IDs must be unique in Configuration" - ); - - let Ok(realm_count) = u32::try_from(c.realms.len()) else { - panic!("too many realms in Client configuration"); - }; - - for realm in &c.realms { - if let Some(public_key) = realm.public_key.as_ref() { - assert_eq!( - public_key.len(), - 32, - "realm public keys must be 32 bytes" // (x25519 for now) - ); - } - } - - assert!( - 1 <= c.recover_threshold, - "Configuration recover_threshold must be at least 1" - ); - assert!( - c.recover_threshold <= realm_count, - "Configuration recover_threshold cannot exceed number of realms" - ); - assert!( - c.recover_threshold > realm_count / 2, - "Configuration recover_threshold must contain a majority of realms" - ); - - assert!( - c.recover_threshold <= realm_count, - "Configuration register_threshold must be at least recover_threshold" - ); - assert!( - c.register_threshold <= realm_count, - "Configuration register_threshold cannot exceed number of realms" - ); - - // perform a fixed sorting of realms based on their id, so that shares - // are produced in a consistent ordering for a given configuration. - let mut sorted_realms = c.realms.clone(); - sorted_realms.sort_by(|lhs, rhs| lhs.id.cmp(&rhs.id)); - - Self(Configuration { - realms: sorted_realms, - register_threshold: c.register_threshold, - recover_threshold: c.recover_threshold, - pin_hashing_mode: c.pin_hashing_mode, - }) - } -} - -impl CheckedConfiguration { - pub fn share_index(&self, realm: &RealmId) -> Option { - if let Some(index) = self.realms.iter().position(|r| r.id == *realm) { - (index + 1).try_into().map(Index).ok() - } else { - None - } - } - - pub fn share_count(&self) -> u32 { - self.realms.len().try_into().unwrap() - } -} - -impl Deref for CheckedConfiguration { - type Target = Configuration; - - fn deref(&self) -> &Self::Target { - &self.0 - } -} - -#[cfg(test)] -mod tests { - use super::Configuration; - - #[test] - fn test_configuration_json() { - let input = r#"{ - "realms": [ - { - "id": "0102030405060708090a0b0c0d0e0f10", - "address": "https://juicebox.hsm.realm.address/", - "public_key": "0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20" - }, - { - "id": "2102030405060708090a0b0c0d0e0f10", - "address": "https://your.software.realm.address/" - }, - { - "id": "3102030405060708090a0b0c0d0e0f10", - "address": "https://juicebox.software.realm.address/" - } - ], - "register_threshold": 3, - "recover_threshold": 3, - "pin_hashing_mode": "Standard2019" -}"#; - println!("input:"); - println!("{input}"); - - let configuration = Configuration::from_json(input).unwrap(); - println!("parsed:"); - println!("{configuration:#?}"); - - let serialized = configuration.to_json(); - println!("serialized:"); - println!("{serialized}"); - - assert_eq!(input, serialized); - } -} +use serde::{Deserialize, Serialize}; +use std::{collections::HashSet, ops::Deref}; + +use crate::{PinHashingMode, Realm}; +use juicebox_realm_api::types::RealmId; +use juicebox_secret_sharing::Index; + +/// The parameters used to configure a [`Client`](crate::Client). +#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)] +pub struct Configuration { + /// The remote services that the client interacts with. + /// + /// There must be between `register_threshold` and 255 realms, inclusive. + pub realms: Vec, + + /// A registration will be considered successful if it's successful on at + /// least this many realms. + /// + /// Must be between `recover_threshold` and `realms.len()`, inclusive. + pub register_threshold: u32, + + /// A recovery (or an adversary) will need the cooperation of this many + /// realms to retrieve the secret. + /// + /// Must be between `(realms.len() / 2).ceil()` and `realms.len()`, inclusive. + pub recover_threshold: u32, + + /// Defines how the provided PIN will be hashed before register and recover + /// operations. Changing modes will make previous secrets stored on the realms + /// inaccessible with the same PIN and should not be done without re-registering + /// secrets. + pub pin_hashing_mode: PinHashingMode, +} + +impl Configuration { + pub fn from_json(s: &str) -> Result { + serde_json::from_str(s) + } + + pub fn to_json(&self) -> String { + serde_json::to_string_pretty(self).expect("failed to convert configuration to json") + } +} + +#[derive(Debug)] +pub(crate) struct CheckedConfiguration(Configuration); + +impl CheckedConfiguration { + pub fn from(c: Configuration) -> Self { + assert!( + !c.realms.is_empty(), + "Client needs at least one realm in Configuration" + ); + + assert_eq!( + c.realms + .iter() + .map(|realm| realm.id) + .collect::>() + .len(), + c.realms.len(), + "realm IDs must be unique in Configuration" + ); + + let Ok(realm_count) = u32::try_from(c.realms.len()) else { + panic!("too many realms in Client configuration"); + }; + + for realm in &c.realms { + if let Some(public_key) = realm.public_key.as_ref() { + assert_eq!( + public_key.len(), + 32, + "realm public keys must be 32 bytes" // (x25519 for now) + ); + } + } + + assert!( + 1 <= c.recover_threshold, + "Configuration recover_threshold must be at least 1" + ); + assert!( + c.recover_threshold <= realm_count, + "Configuration recover_threshold cannot exceed number of realms" + ); + assert!( + c.recover_threshold > realm_count / 2, + "Configuration recover_threshold must contain a majority of realms" + ); + + // NOTE: checks after this point must use `register_threshold` and not + // repeat the `recover_threshold` conditions above. + assert!( + c.register_threshold >= 1, + "Configuration register_threshold must be at least 1" + ); + assert!( + c.register_threshold <= c.recover_threshold, + "Configuration register_threshold cannot exceed recover_threshold" + ); + assert!( + c.register_threshold <= realm_count, + "Configuration register_threshold cannot exceed number of realms" + ); + + // perform a fixed sorting of realms based on their id, so that shares + // are produced in a consistent ordering for a given configuration. + let mut sorted_realms = c.realms.clone(); + sorted_realms.sort_by(|lhs, rhs| lhs.id.cmp(&rhs.id)); + + Self(Configuration { + realms: sorted_realms, + register_threshold: c.register_threshold, + recover_threshold: c.recover_threshold, + pin_hashing_mode: c.pin_hashing_mode, + }) + } +} + +impl CheckedConfiguration { + pub fn share_index(&self, realm: &RealmId) -> Option { + if let Some(index) = self.realms.iter().position(|r| r.id == *realm) { + (index + 1).try_into().map(Index).ok() + } else { + None + } + } + + pub fn share_count(&self) -> u32 { + self.realms.len().try_into().unwrap() + } +} + +impl Deref for CheckedConfiguration { + type Target = Configuration; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +#[cfg(test)] +mod tests { + use super::Configuration; + + #[test] + fn test_configuration_json() { + let input = r#"{ + "realms": [ + { + "id": "0102030405060708090a0b0c0d0e0f10", + "address": "https://juicebox.hsm.realm.address/", + "public_key": "0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20" + }, + { + "id": "2102030405060708090a0b0c0d0e0f10", + "address": "https://your.software.realm.address/" + }, + { + "id": "3102030405060708090a0b0c0d0e0f10", + "address": "https://juicebox.software.realm.address/" + } + ], + "register_threshold": 3, + "recover_threshold": 3, + "pin_hashing_mode": "Standard2019" +}"#; + println!("input:"); + println!("{input}"); + + let configuration = Configuration::from_json(input).unwrap(); + println!("parsed:"); + println!("{configuration:#?}"); + + let serialized = configuration.to_json(); + println!("serialized:"); + println!("{serialized}"); + + assert_eq!(input, serialized); + } +} diff --git a/rust/sdk/src/recover.rs b/rust/sdk/src/recover.rs index e9be01c..6fe2e01 100644 --- a/rust/sdk/src/recover.rs +++ b/rust/sdk/src/recover.rs @@ -1,447 +1,450 @@ -use curve25519_dalek::{RistrettoPoint, Scalar}; -use rand::rngs::OsRng; -use std::collections::HashMap; -use std::error::Error; -use std::fmt::{Debug, Display}; -use subtle::ConstantTimeEq; -use tracing::instrument; - -use juicebox_oprf as oprf; -use juicebox_realm_api::{ - requests::{ - Recover1Response, Recover2Request, Recover2Response, Recover3Request, Recover3Response, - SecretsRequest, SecretsResponse, - }, - signing::OprfVerifyingKey, - types::{ - EncryptedUserSecret, EncryptedUserSecretCommitment, RegistrationVersion, - UnlockKeyCommitment, UnlockKeyTag, UserSecretEncryptionKeyScalarShare, - }, -}; -use juicebox_secret_sharing::{recover_secret, RecoverSecretError, Share}; - -use crate::{ - auth, - configuration::CheckedConfiguration, - http, - request::{join_at_least_threshold, RequestError}, - types::{ - derive_unlock_key_and_commitment, UserSecretEncryptionKey, UserSecretEncryptionKeyScalar, - }, - Client, Pin, Realm, Sleeper, UserInfo, UserSecret, -}; - -/// Error return type for [`Client::recover`]. -#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)] -pub enum RecoverError { - /// The secret could not be unlocked, but you can try again - /// with a different PIN if you have guesses remaining. If no - /// guesses remain, this secret is locked and inaccessible. - InvalidPin { guesses_remaining: u16 }, - - /// The secret was not registered or not fully registered with the - /// provided realms. - NotRegistered, - - /// A realm rejected the `Client`'s auth token. - InvalidAuth, - - /// The SDK software is too old to communicate with this realm - /// and must be upgraded. - UpgradeRequired, - - /// The tenant has exceeded their allowed number of operations. Try again - /// later. - RateLimitExceeded, - - /// A software error has occurred. This request should not be retried - /// with the same parameters. Verify your inputs, check for software - /// updates and try again. - Assertion, - - /// A transient error in sending or receiving requests to a realm. - /// This request may succeed by trying again with the same parameters. - Transient, -} - -impl Display for RecoverError { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - Debug::fmt(self, f) - } -} - -impl Error for RecoverError {} - -impl Client { - pub(crate) async fn perform_recover( - &self, - pin: &Pin, - info: &UserInfo, - ) -> Result { - let mut configuration = &self.configuration; - let mut iter = self.previous_configurations.iter(); - loop { - return match self - .perform_recover_with_configuration(pin, info, configuration) - .await - { - Ok(secret) => Ok(secret), - Err(RecoverError::NotRegistered) => { - if let Some(next_configuration) = iter.next() { - configuration = next_configuration; - continue; - } - - Err(RecoverError::NotRegistered) - } - Err(err) => Err(err), - }; - } - } - - /// Performs phase 1 of recovery for the parameters specified in a given - /// configuration. If successful, attempts to complete recovery for each - /// subset of realms larger than the recover threshold with matching salts. - #[instrument(level = "trace", skip_all, err(level = "trace", Debug))] - async fn perform_recover_with_configuration( - &self, - pin: &Pin, - info: &UserInfo, - configuration: &CheckedConfiguration, - ) -> Result { - let recover1_requests = configuration - .realms - .iter() - .map(|realm| self.recover1_on_realm(realm)); - - let mut realms_per_version: HashMap> = HashMap::new(); - for (version, realm) in - join_at_least_threshold(recover1_requests, configuration.recover_threshold).await? - { - realms_per_version.entry(version).or_default().push(realm); - } - - realms_per_version - .retain(|_, values| values.len() >= configuration.recover_threshold as usize); - - // We enforce a strict majority for the `recover_threshold`, so there should always - // be one or none realms with consensus on a version available to recover from. - assert!(realms_per_version.len() <= 1); - - let Some((version, realms)) = realms_per_version.into_iter().next() else { - return Err(RecoverError::NotRegistered); - }; - - let (access_key, encryption_key_seed) = pin - .hash(configuration.pin_hashing_mode, &version, info) - .expect("pin hashing failed"); - - let (oprf_blinding_factor, oprf_blinded_input) = - oprf::start(access_key.expose_secret(), &mut OsRng); - - let recover2_requests = realms.iter().map(|realm| { - self.recover2_on_realm(realm, configuration, &version, &oprf_blinded_input) - }); - - let mut oprf_blinded_result_shares_by_commitment_and_verifying_key: HashMap<_, Vec<_>> = - HashMap::new(); - - // TODO: this should stop after finding threshold realms that agree on - // commitment and verifying key - for (oprf_verifying_key, share, commitment, guesses_remaining) in - join_at_least_threshold(recover2_requests, configuration.recover_threshold).await? - { - oprf_blinded_result_shares_by_commitment_and_verifying_key - .entry((commitment, oprf_verifying_key)) - .or_default() - .push((share, guesses_remaining)); - } - - oprf_blinded_result_shares_by_commitment_and_verifying_key - .retain(|_, values| values.len() >= configuration.recover_threshold as usize); - - // We enforce a strict majority for the `recover_threshold`, so there should always - // be one or none realms with consensus on an unlock key commitment and verifying - // key to recover from. - assert!(oprf_blinded_result_shares_by_commitment_and_verifying_key.len() <= 1); - - let Some(((unlock_key_commitment, _), oprf_blinded_result_shares_and_guesses_remaining)) = - oprf_blinded_result_shares_by_commitment_and_verifying_key - .into_iter() - .next() - else { - return Err(RecoverError::Assertion); - }; - - let (oprf_blinded_result_shares, all_guesses_remaining): ( - Vec>, - Vec, - ) = oprf_blinded_result_shares_and_guesses_remaining - .into_iter() - .unzip(); - - let oprf_blinded_result = match recover_secret(&oprf_blinded_result_shares) { - Ok(blinded_result) => oprf::BlindedOutput::from(blinded_result), - Err(RecoverSecretError::DuplicateShares) => return Err(RecoverError::Assertion), - }; - let oprf_result = oprf::finalize( - access_key.expose_secret(), - &oprf_blinding_factor, - &oprf_blinded_result, - ); - - let (unlock_key, our_commitment) = derive_unlock_key_and_commitment(&oprf_result); - if !bool::from(unlock_key_commitment.ct_eq(&our_commitment)) { - let guesses_remaining = all_guesses_remaining.into_iter().min().unwrap(); - return Err(RecoverError::InvalidPin { guesses_remaining }); - } - - let recover3_requests = realms.iter().map(|realm| { - self.recover3_on_realm( - realm, - configuration, - &version, - UnlockKeyTag::derive(&unlock_key, &realm.id), - ) - }); - - let mut encryption_key_scalar_shares_by_encrypted_secret: HashMap< - EncryptedUserSecret, - Vec>, - > = HashMap::new(); - - for (share, encrypted_secret, commitment, realm) in - join_at_least_threshold(recover3_requests, configuration.recover_threshold).await? - { - let our_commitment = EncryptedUserSecretCommitment::derive( - &unlock_key, - &realm.id, - &UserSecretEncryptionKeyScalarShare::from(share.secret), - &encrypted_secret, - ); - - // We can't use the share from this realm, but we continue - // as there may still be enough material from other realms. - if !bool::from(our_commitment.ct_eq(&commitment)) { - continue; - } - - encryption_key_scalar_shares_by_encrypted_secret - .entry(encrypted_secret) - .or_default() - .push(share); - } - - encryption_key_scalar_shares_by_encrypted_secret - .retain(|_, values| values.len() >= configuration.recover_threshold as usize); - - // We enforce a strict majority for the `recover_threshold`, so there should always - // be one or none realms with consensus on an encrypted secret to recover from. - assert!(encryption_key_scalar_shares_by_encrypted_secret.len() <= 1); - - let Some((encrypted_secret, encryption_key_scalar_shares)) = - encryption_key_scalar_shares_by_encrypted_secret - .into_iter() - .next() - else { - return Err(RecoverError::Assertion); - }; - - match recover_secret(&encryption_key_scalar_shares) { - Ok(secret) => { - let scalar = UserSecretEncryptionKeyScalar::new(secret); - let encryption_key = UserSecretEncryptionKey::derive(&encryption_key_seed, &scalar); - - Ok(UserSecret::decrypt(&encrypted_secret, &encryption_key)) - } - Err(_) => Err(RecoverError::Assertion), - } - } - - /// Performs phase 1 of recovery on a particular realm. - #[instrument(level = "trace", skip(self), err(level = "trace", Debug))] - async fn recover1_on_realm( - &self, - realm: &Realm, - ) -> Result<(RegistrationVersion, Realm), RecoverError> { - match self.make_request(realm, SecretsRequest::Recover1).await { - Err(RequestError::UpgradeRequired) => Err(RecoverError::UpgradeRequired), - Err(RequestError::InvalidAuth) => Err(RecoverError::InvalidAuth), - Err(RequestError::Assertion) => Err(RecoverError::Assertion), - Err(RequestError::Transient) => Err(RecoverError::Transient), - Err(RequestError::RateLimitExceeded) => Err(RecoverError::RateLimitExceeded), - - Ok(SecretsResponse::Recover1(response)) => match response { - Recover1Response::Ok { version } => Ok((version, realm.to_owned())), - Recover1Response::NotRegistered => Err(RecoverError::NotRegistered), - Recover1Response::NoGuesses => Err(RecoverError::InvalidPin { - guesses_remaining: 0, - }), - }, - Ok(_) => Err(RecoverError::Assertion), - } - } - - /// Performs phase 2 of recovery on a particular realm. - #[instrument(level = "trace", skip_all, err(level = "trace", Debug))] - async fn recover2_on_realm( - &self, - realm: &Realm, - configuration: &CheckedConfiguration, - version: &RegistrationVersion, - oprf_blinded_input: &oprf::BlindedInput, - ) -> Result< - ( - OprfVerifyingKey, - Share, - UnlockKeyCommitment, - u16, - ), - RecoverError, - > { - let recover2_request = self.make_request( - realm, - SecretsRequest::Recover2(Recover2Request { - version: version.to_owned(), - oprf_blinded_input: oprf_blinded_input.to_owned(), - }), - ); - - let ( - oprf_signed_public_key, - oprf_blinded_result, - oprf_proof, - unlock_key_commitment, - guesses_remaining, - ) = match recover2_request.await { - Err(RequestError::UpgradeRequired) => return Err(RecoverError::UpgradeRequired), - Err(RequestError::Transient) => return Err(RecoverError::Transient), - Err(RequestError::Assertion) => return Err(RecoverError::Assertion), - Err(RequestError::InvalidAuth) => return Err(RecoverError::InvalidAuth), - Err(RequestError::RateLimitExceeded) => return Err(RecoverError::RateLimitExceeded), - - Ok(SecretsResponse::Recover2(rr)) => match rr { - Recover2Response::Ok { - oprf_signed_public_key, - oprf_blinded_result, - oprf_proof, - unlock_key_commitment, - num_guesses, - guess_count, - } => ( - oprf_signed_public_key, - oprf_blinded_result, - oprf_proof, - unlock_key_commitment, - num_guesses - guess_count, - ), - - Recover2Response::VersionMismatch => { - return Err(RecoverError::Assertion); - } - - Recover2Response::NotRegistered => { - return Err(RecoverError::NotRegistered); - } - - Recover2Response::NoGuesses => { - return Err(RecoverError::InvalidPin { - guesses_remaining: 0, - }); - } - }, - - Ok(_) => return Err(RecoverError::Assertion), - }; - - oprf_signed_public_key - .verify(&realm.id) - .map_err(|_| RecoverError::Assertion)?; - - oprf::verify_proof( - oprf_blinded_input, - &oprf_blinded_result, - &oprf_signed_public_key.public_key, - &oprf_proof, - ) - .map_err(|_| RecoverError::Assertion)?; - - let oprf_blinded_result_share = Share { - index: configuration - .share_index(&realm.id) - .ok_or(RecoverError::Assertion)?, - secret: oprf_blinded_result.to_point(), - }; - - Ok(( - oprf_signed_public_key.verifying_key, - oprf_blinded_result_share, - unlock_key_commitment, - guesses_remaining, - )) - } - - /// Performs phase 3 of recovery on a particular realm. - #[instrument(level = "trace", skip_all)] - async fn recover3_on_realm( - &self, - realm: &Realm, - configuration: &CheckedConfiguration, - version: &RegistrationVersion, - unlock_key_tag: UnlockKeyTag, - ) -> Result< - ( - Share, - EncryptedUserSecret, - EncryptedUserSecretCommitment, - Realm, - ), - RecoverError, - > { - let recover3_request = self.make_request( - realm, - SecretsRequest::Recover3(Recover3Request { - version: version.to_owned(), - unlock_key_tag, - }), - ); - - match recover3_request.await { - Err(RequestError::UpgradeRequired) => Err(RecoverError::UpgradeRequired), - Err(RequestError::Transient) => Err(RecoverError::Transient), - Err(RequestError::Assertion) => Err(RecoverError::Assertion), - Err(RequestError::InvalidAuth) => Err(RecoverError::InvalidAuth), - Err(RequestError::RateLimitExceeded) => Err(RecoverError::RateLimitExceeded), - - Ok(SecretsResponse::Recover3(rr)) => match rr { - Recover3Response::Ok { - encryption_key_scalar_share, - encrypted_secret, - encrypted_secret_commitment, - } => { - let secret_share = Share { - index: configuration - .share_index(&realm.id) - .ok_or(RecoverError::Assertion)?, - secret: encryption_key_scalar_share.to_scalar(), - }; - Ok(( - secret_share, - encrypted_secret, - encrypted_secret_commitment, - realm.to_owned(), - )) - } - Recover3Response::NotRegistered => Err(RecoverError::NotRegistered), - Recover3Response::NoGuesses => Err(RecoverError::InvalidPin { - guesses_remaining: 0, - }), - Recover3Response::BadUnlockKeyTag { guesses_remaining } => { - Err(RecoverError::InvalidPin { guesses_remaining }) - } - Recover3Response::VersionMismatch => Err(RecoverError::Assertion), - }, - Ok(_) => Err(RecoverError::Assertion), - } - } -} +use curve25519_dalek::{RistrettoPoint, Scalar}; +use rand::rngs::OsRng; +use std::collections::HashMap; +use std::error::Error; +use std::fmt::{Debug, Display}; +use subtle::ConstantTimeEq; +use tracing::instrument; + +use juicebox_oprf as oprf; +use juicebox_realm_api::{ + requests::{ + Recover1Response, Recover2Request, Recover2Response, Recover3Request, Recover3Response, + SecretsRequest, SecretsResponse, + }, + signing::OprfVerifyingKey, + types::{ + EncryptedUserSecret, EncryptedUserSecretCommitment, RegistrationVersion, + UnlockKeyCommitment, UnlockKeyTag, UserSecretEncryptionKeyScalarShare, + }, +}; +use juicebox_secret_sharing::{recover_secret, RecoverSecretError, Share}; + +use crate::{ + auth, + configuration::CheckedConfiguration, + http, + request::{join_at_least_threshold, RequestError}, + types::{ + derive_unlock_key_and_commitment, UserSecretEncryptionKey, UserSecretEncryptionKeyScalar, + }, + Client, Pin, Realm, Sleeper, UserInfo, UserSecret, +}; + +/// Error return type for [`Client::recover`]. +#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)] +pub enum RecoverError { + /// The secret could not be unlocked, but you can try again + /// with a different PIN if you have guesses remaining. If no + /// guesses remain, this secret is locked and inaccessible. + InvalidPin { guesses_remaining: u16 }, + + /// The secret was not registered or not fully registered with the + /// provided realms. + NotRegistered, + + /// A realm rejected the `Client`'s auth token. + InvalidAuth, + + /// The SDK software is too old to communicate with this realm + /// and must be upgraded. + UpgradeRequired, + + /// The tenant has exceeded their allowed number of operations. Try again + /// later. + RateLimitExceeded, + + /// A software error has occurred. This request should not be retried + /// with the same parameters. Verify your inputs, check for software + /// updates and try again. + Assertion, + + /// A transient error in sending or receiving requests to a realm. + /// This request may succeed by trying again with the same parameters. + Transient, +} + +impl Display for RecoverError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + Debug::fmt(self, f) + } +} + +impl Error for RecoverError {} + +impl Client { + pub(crate) async fn perform_recover( + &self, + pin: &Pin, + info: &UserInfo, + ) -> Result { + let mut configuration = &self.configuration; + let mut iter = self.previous_configurations.iter(); + loop { + return match self + .perform_recover_with_configuration(pin, info, configuration) + .await + { + Ok(secret) => Ok(secret), + Err(RecoverError::NotRegistered) => { + if let Some(next_configuration) = iter.next() { + configuration = next_configuration; + continue; + } + + Err(RecoverError::NotRegistered) + } + Err(err) => Err(err), + }; + } + } + + /// Performs phase 1 of recovery for the parameters specified in a given + /// configuration. If successful, attempts to complete recovery for each + /// subset of realms larger than the recover threshold with matching salts. + #[instrument(level = "trace", skip_all, err(level = "trace", Debug))] + async fn perform_recover_with_configuration( + &self, + pin: &Pin, + info: &UserInfo, + configuration: &CheckedConfiguration, + ) -> Result { + let recover1_requests = configuration + .realms + .iter() + .map(|realm| self.recover1_on_realm(realm)); + + let mut realms_per_version: HashMap> = HashMap::new(); + for (version, realm) in + join_at_least_threshold(recover1_requests, configuration.recover_threshold).await? + { + realms_per_version.entry(version).or_default().push(realm); + } + + realms_per_version + .retain(|_, values| values.len() >= configuration.recover_threshold as usize); + + // We enforce a strict majority for the `recover_threshold`, so there should always + // be one or none realms with consensus on a version available to recover from. + assert!(realms_per_version.len() <= 1); + + let Some((version, realms)) = realms_per_version.into_iter().next() else { + return Err(RecoverError::NotRegistered); + }; + + let (access_key, encryption_key_seed) = pin + .hash(configuration.pin_hashing_mode, &version, info) + .expect("pin hashing failed"); + + let (oprf_blinding_factor, oprf_blinded_input) = + oprf::start(access_key.expose_secret(), &mut OsRng); + + let recover2_requests = realms.iter().map(|realm| { + self.recover2_on_realm(realm, configuration, &version, &oprf_blinded_input) + }); + + let mut oprf_blinded_result_shares_by_commitment_and_verifying_key: HashMap<_, Vec<_>> = + HashMap::new(); + + // TODO: this should stop after finding threshold realms that agree on + // commitment and verifying key + for (oprf_verifying_key, share, commitment, guesses_remaining) in + join_at_least_threshold(recover2_requests, configuration.recover_threshold).await? + { + oprf_blinded_result_shares_by_commitment_and_verifying_key + .entry((commitment, oprf_verifying_key)) + .or_default() + .push((share, guesses_remaining)); + } + + oprf_blinded_result_shares_by_commitment_and_verifying_key + .retain(|_, values| values.len() >= configuration.recover_threshold as usize); + + // We enforce a strict majority for the `recover_threshold`, so there should always + // be one or none realms with consensus on an unlock key commitment and verifying + // key to recover from. + assert!(oprf_blinded_result_shares_by_commitment_and_verifying_key.len() <= 1); + + let Some(((unlock_key_commitment, _), oprf_blinded_result_shares_and_guesses_remaining)) = + oprf_blinded_result_shares_by_commitment_and_verifying_key + .into_iter() + .next() + else { + return Err(RecoverError::Assertion); + }; + + let (oprf_blinded_result_shares, all_guesses_remaining): ( + Vec>, + Vec, + ) = oprf_blinded_result_shares_and_guesses_remaining + .into_iter() + .unzip(); + + let oprf_blinded_result = match recover_secret(&oprf_blinded_result_shares) { + Ok(blinded_result) => oprf::BlindedOutput::from(blinded_result), + Err(RecoverSecretError::DuplicateShares) => return Err(RecoverError::Assertion), + }; + let oprf_result = oprf::finalize( + access_key.expose_secret(), + &oprf_blinding_factor, + &oprf_blinded_result, + ); + + let (unlock_key, our_commitment) = derive_unlock_key_and_commitment(&oprf_result); + if !bool::from(unlock_key_commitment.ct_eq(&our_commitment)) { + let guesses_remaining = all_guesses_remaining.into_iter().min().unwrap(); + return Err(RecoverError::InvalidPin { guesses_remaining }); + } + + let recover3_requests = realms.iter().map(|realm| { + self.recover3_on_realm( + realm, + configuration, + &version, + UnlockKeyTag::derive(&unlock_key, &realm.id), + ) + }); + + let mut encryption_key_scalar_shares_by_encrypted_secret: HashMap< + EncryptedUserSecret, + Vec>, + > = HashMap::new(); + + for (share, encrypted_secret, commitment, realm) in + join_at_least_threshold(recover3_requests, configuration.recover_threshold).await? + { + let our_commitment = EncryptedUserSecretCommitment::derive( + &unlock_key, + &realm.id, + &UserSecretEncryptionKeyScalarShare::from(share.secret), + &encrypted_secret, + ); + + // We can't use the share from this realm, but we continue + // as there may still be enough material from other realms. + if !bool::from(our_commitment.ct_eq(&commitment)) { + continue; + } + + encryption_key_scalar_shares_by_encrypted_secret + .entry(encrypted_secret) + .or_default() + .push(share); + } + + encryption_key_scalar_shares_by_encrypted_secret + .retain(|_, values| values.len() >= configuration.recover_threshold as usize); + + // We enforce a strict majority for the `recover_threshold`, so there should always + // be one or none realms with consensus on an encrypted secret to recover from. + assert!(encryption_key_scalar_shares_by_encrypted_secret.len() <= 1); + + let Some((encrypted_secret, encryption_key_scalar_shares)) = + encryption_key_scalar_shares_by_encrypted_secret + .into_iter() + .next() + else { + return Err(RecoverError::Assertion); + }; + + match recover_secret(&encryption_key_scalar_shares) { + Ok(secret) => { + let scalar = UserSecretEncryptionKeyScalar::new(secret); + let encryption_key = UserSecretEncryptionKey::derive(&encryption_key_seed, &scalar); + + Ok(UserSecret::decrypt(&encrypted_secret, &encryption_key)) + } + Err(_) => Err(RecoverError::Assertion), + } + } + + /// Performs phase 1 of recovery on a particular realm. + #[instrument(level = "trace", skip(self), err(level = "trace", Debug))] + async fn recover1_on_realm( + &self, + realm: &Realm, + ) -> Result<(RegistrationVersion, Realm), RecoverError> { + match self.make_request(realm, SecretsRequest::Recover1).await { + Err(RequestError::UpgradeRequired) => Err(RecoverError::UpgradeRequired), + Err(RequestError::InvalidAuth) => Err(RecoverError::InvalidAuth), + Err(RequestError::Assertion) => Err(RecoverError::Assertion), + Err(RequestError::Transient) => Err(RecoverError::Transient), + Err(RequestError::RateLimitExceeded) => Err(RecoverError::RateLimitExceeded), + + Ok(SecretsResponse::Recover1(response)) => match response { + Recover1Response::Ok { version } => Ok((version, realm.to_owned())), + Recover1Response::NotRegistered => Err(RecoverError::NotRegistered), + Recover1Response::NoGuesses => Err(RecoverError::InvalidPin { + guesses_remaining: 0, + }), + }, + Ok(_) => Err(RecoverError::Assertion), + } + } + + /// Performs phase 2 of recovery on a particular realm. + #[instrument(level = "trace", skip_all, err(level = "trace", Debug))] + async fn recover2_on_realm( + &self, + realm: &Realm, + configuration: &CheckedConfiguration, + version: &RegistrationVersion, + oprf_blinded_input: &oprf::BlindedInput, + ) -> Result< + ( + OprfVerifyingKey, + Share, + UnlockKeyCommitment, + u16, + ), + RecoverError, + > { + let recover2_request = self.make_request( + realm, + SecretsRequest::Recover2(Recover2Request { + version: version.to_owned(), + oprf_blinded_input: oprf_blinded_input.to_owned(), + }), + ); + + let ( + oprf_signed_public_key, + oprf_blinded_result, + oprf_proof, + unlock_key_commitment, + guesses_remaining, + ) = match recover2_request.await { + Err(RequestError::UpgradeRequired) => return Err(RecoverError::UpgradeRequired), + Err(RequestError::Transient) => return Err(RecoverError::Transient), + Err(RequestError::Assertion) => return Err(RecoverError::Assertion), + Err(RequestError::InvalidAuth) => return Err(RecoverError::InvalidAuth), + Err(RequestError::RateLimitExceeded) => return Err(RecoverError::RateLimitExceeded), + + Ok(SecretsResponse::Recover2(rr)) => match rr { + Recover2Response::Ok { + oprf_signed_public_key, + oprf_blinded_result, + oprf_proof, + unlock_key_commitment, + num_guesses, + guess_count, + } => ( + oprf_signed_public_key, + oprf_blinded_result, + oprf_proof, + unlock_key_commitment, + // Saturating subtraction: if a realm reports + // `guess_count > num_guesses`, this must not underflow + // (panic in debug, wrap to a huge value in release). + num_guesses.saturating_sub(guess_count), + ), + + Recover2Response::VersionMismatch => { + return Err(RecoverError::Assertion); + } + + Recover2Response::NotRegistered => { + return Err(RecoverError::NotRegistered); + } + + Recover2Response::NoGuesses => { + return Err(RecoverError::InvalidPin { + guesses_remaining: 0, + }); + } + }, + + Ok(_) => return Err(RecoverError::Assertion), + }; + + oprf_signed_public_key + .verify(&realm.id) + .map_err(|_| RecoverError::Assertion)?; + + oprf::verify_proof( + oprf_blinded_input, + &oprf_blinded_result, + &oprf_signed_public_key.public_key, + &oprf_proof, + ) + .map_err(|_| RecoverError::Assertion)?; + + let oprf_blinded_result_share = Share { + index: configuration + .share_index(&realm.id) + .ok_or(RecoverError::Assertion)?, + secret: oprf_blinded_result.to_point(), + }; + + Ok(( + oprf_signed_public_key.verifying_key, + oprf_blinded_result_share, + unlock_key_commitment, + guesses_remaining, + )) + } + + /// Performs phase 3 of recovery on a particular realm. + #[instrument(level = "trace", skip_all)] + async fn recover3_on_realm( + &self, + realm: &Realm, + configuration: &CheckedConfiguration, + version: &RegistrationVersion, + unlock_key_tag: UnlockKeyTag, + ) -> Result< + ( + Share, + EncryptedUserSecret, + EncryptedUserSecretCommitment, + Realm, + ), + RecoverError, + > { + let recover3_request = self.make_request( + realm, + SecretsRequest::Recover3(Recover3Request { + version: version.to_owned(), + unlock_key_tag, + }), + ); + + match recover3_request.await { + Err(RequestError::UpgradeRequired) => Err(RecoverError::UpgradeRequired), + Err(RequestError::Transient) => Err(RecoverError::Transient), + Err(RequestError::Assertion) => Err(RecoverError::Assertion), + Err(RequestError::InvalidAuth) => Err(RecoverError::InvalidAuth), + Err(RequestError::RateLimitExceeded) => Err(RecoverError::RateLimitExceeded), + + Ok(SecretsResponse::Recover3(rr)) => match rr { + Recover3Response::Ok { + encryption_key_scalar_share, + encrypted_secret, + encrypted_secret_commitment, + } => { + let secret_share = Share { + index: configuration + .share_index(&realm.id) + .ok_or(RecoverError::Assertion)?, + secret: encryption_key_scalar_share.to_scalar(), + }; + Ok(( + secret_share, + encrypted_secret, + encrypted_secret_commitment, + realm.to_owned(), + )) + } + Recover3Response::NotRegistered => Err(RecoverError::NotRegistered), + Recover3Response::NoGuesses => Err(RecoverError::InvalidPin { + guesses_remaining: 0, + }), + Recover3Response::BadUnlockKeyTag { guesses_remaining } => { + Err(RecoverError::InvalidPin { guesses_remaining }) + } + Recover3Response::VersionMismatch => Err(RecoverError::Assertion), + }, + Ok(_) => Err(RecoverError::Assertion), + } + } +}