use ring::rand::{SecureRandom, SystemRandom};
use crate::error::{Error, Result};
use crate::pase::kdf::{derive_w0_w1, validate_params};
use crate::pase::messages::{Pake1, Pake2, Pake3, PbkdfParamRequest, PbkdfParamResponse};
use crate::pase::spake2plus::{
compute_ca, compute_cb, compute_x, compute_z_v_prover, derive_confirmation_keys,
derive_session_keys, hash_context, ka_ke_from_transcript, sample_scalar, transcript_hash,
verify_tag,
};
use crate::pase::{PaseMessageKind, PasePbkdfParams, PaseSessionKeys};
use zeroize::Zeroize;
#[derive(Debug)]
enum State {
AwaitingStartNegotiation {
pin: u32,
x_scalar: p256::Scalar,
initiator_random: [u8; 32],
initiator_session_id: u16,
},
AwaitingStartKnownParams {
pin: u32,
params: PasePbkdfParams,
x_scalar: p256::Scalar,
},
AwaitingPbkdfResponse {
pin: u32,
x_scalar: p256::Scalar,
sent_request_bytes: Vec<u8>,
},
ReadyToSendPake1 {
pin: u32,
params: PasePbkdfParams,
x_scalar: p256::Scalar,
transcript_context: [u8; 32],
},
AwaitingPake2 {
w0: p256::Scalar,
w1: p256::Scalar,
x_scalar: p256::Scalar,
x_bytes: [u8; 65],
transcript_context: [u8; 32],
},
ReadyToSendPake3 {
ca: [u8; 32],
session_keys: PaseSessionKeys,
},
Complete { session_keys: PaseSessionKeys },
Poisoned,
}
pub struct PaseProver {
state: State,
responder_session_id: Option<u16>,
}
impl PaseProver {
pub fn new_with_negotiation(pin: u32, initiator_session_id: u16) -> Result<Self> {
let rng = SystemRandom::new();
Self::new_with_negotiation_using_rng(pin, initiator_session_id, &rng)
}
pub(crate) fn new_with_negotiation_using_rng(
pin: u32,
initiator_session_id: u16,
rng: &dyn SecureRandom,
) -> Result<Self> {
let x_scalar = sample_scalar(rng)?;
let mut initiator_random = [0u8; 32];
rng.fill(&mut initiator_random)
.map_err(|_| Error::PinDerivationFailed)?;
Ok(Self {
state: State::AwaitingStartNegotiation {
pin,
x_scalar,
initiator_random,
initiator_session_id,
},
responder_session_id: None,
})
}
pub(crate) fn new_with_negotiation_with_scalar(
pin: u32,
x_scalar_bytes: [u8; 32],
initiator_random: [u8; 32],
) -> Result<Self> {
Self::new_with_negotiation_with_scalar_and_session_id(
pin,
x_scalar_bytes,
initiator_random,
0,
)
}
pub(crate) fn new_with_negotiation_with_scalar_and_session_id(
pin: u32,
x_scalar_bytes: [u8; 32],
initiator_random: [u8; 32],
initiator_session_id: u16,
) -> Result<Self> {
use p256::elliptic_curve::group::ff::{Field, PrimeField};
let x_scalar_opt: Option<p256::Scalar> =
p256::Scalar::from_repr(p256::FieldBytes::from(x_scalar_bytes)).into();
let x_scalar = x_scalar_opt.ok_or(Error::InvalidScalar)?;
if bool::from(x_scalar.is_zero()) {
return Err(Error::InvalidScalar);
}
Ok(Self {
state: State::AwaitingStartNegotiation {
pin,
x_scalar,
initiator_random,
initiator_session_id,
},
responder_session_id: None,
})
}
pub fn new_with_known_params(
pin: u32,
params: PasePbkdfParams,
initiator_session_id: u16,
) -> Result<Self> {
validate_params(params.iterations, ¶ms.salt)?;
let rng = SystemRandom::new();
Self::new_with_known_params_using_rng(pin, params, initiator_session_id, &rng)
}
pub(crate) fn new_with_known_params_using_rng(
pin: u32,
params: PasePbkdfParams,
_initiator_session_id: u16,
rng: &dyn SecureRandom,
) -> Result<Self> {
validate_params(params.iterations, ¶ms.salt)?;
let x_scalar = sample_scalar(rng)?;
Ok(Self {
state: State::AwaitingStartKnownParams {
pin,
params,
x_scalar,
},
responder_session_id: None,
})
}
pub(crate) fn new_with_known_params_with_scalar(
pin: u32,
params: PasePbkdfParams,
x_scalar_bytes: [u8; 32],
) -> Result<Self> {
use p256::elliptic_curve::group::ff::{Field, PrimeField};
validate_params(params.iterations, ¶ms.salt)?;
let x_scalar_opt: Option<p256::Scalar> =
p256::Scalar::from_repr(p256::FieldBytes::from(x_scalar_bytes)).into();
let x_scalar = x_scalar_opt.ok_or(Error::InvalidScalar)?;
if bool::from(x_scalar.is_zero()) {
return Err(Error::InvalidScalar);
}
Ok(Self {
state: State::AwaitingStartKnownParams {
pin,
params,
x_scalar,
},
responder_session_id: None,
})
}
pub fn expected_inbound(&self) -> Option<PaseMessageKind> {
match &self.state {
State::AwaitingPbkdfResponse { .. } => Some(PaseMessageKind::PbkdfParamResponse),
State::AwaitingPake2 { .. } => Some(PaseMessageKind::Pake2),
_ => None,
}
}
#[must_use]
pub fn responder_session_id(&self) -> Option<u16> {
self.responder_session_id
}
pub fn start(&mut self) -> Result<Vec<u8>> {
let prev = std::mem::replace(&mut self.state, State::Poisoned);
match prev {
State::AwaitingStartNegotiation {
pin,
x_scalar,
initiator_random,
initiator_session_id,
} => {
let req = PbkdfParamRequest {
initiator_random,
initiator_session_id,
passcode_id: 0,
has_pbkdf_parameters: false,
initiator_session_params: None,
};
let bytes = req.encode()?;
self.state = State::AwaitingPbkdfResponse {
pin,
x_scalar,
sent_request_bytes: bytes.clone(),
};
Ok(bytes)
}
State::AwaitingStartKnownParams {
pin,
params,
x_scalar,
} => {
let transcript_context = hash_context(&[]);
let (w0, w1) = derive_w0_w1(pin, ¶ms.salt, params.iterations)?;
let x_bytes = compute_x(&x_scalar, &w0);
let pake1_bytes = Pake1 { x: x_bytes }.encode()?;
self.state = State::AwaitingPake2 {
w0,
w1,
x_scalar,
x_bytes,
transcript_context,
};
Ok(pake1_bytes)
}
other => {
self.state = other;
Err(Error::UnexpectedMessage {
expected: PaseMessageKind::PbkdfParamRequest,
got: PaseMessageKind::PbkdfParamRequest,
})
}
}
}
pub fn handle_pbkdf_response(&mut self, bytes: &[u8]) -> Result<()> {
let prev = std::mem::replace(&mut self.state, State::Poisoned);
match prev {
State::AwaitingPbkdfResponse {
pin,
x_scalar,
sent_request_bytes,
} => {
let resp = PbkdfParamResponse::decode(bytes)?;
self.responder_session_id = Some(resp.responder_session_id);
let params_inner = resp.pbkdf_parameters.ok_or(Error::InvalidParameter)?;
let params = PasePbkdfParams {
iterations: params_inner.iterations,
salt: params_inner.salt,
};
validate_params(params.iterations, ¶ms.salt)?;
let transcript_context = hash_context(&[&sent_request_bytes, bytes]);
self.state = State::ReadyToSendPake1 {
pin,
params,
x_scalar,
transcript_context,
};
Ok(())
}
other => {
self.state = other;
Err(Error::UnexpectedMessage {
expected: PaseMessageKind::PbkdfParamResponse,
got: PaseMessageKind::PbkdfParamResponse,
})
}
}
}
#[allow(clippy::similar_names)]
pub fn handle_pake2(&mut self, bytes: &[u8]) -> Result<()> {
let prev = std::mem::replace(&mut self.state, State::Poisoned);
match prev {
State::AwaitingPake2 {
w0,
w1,
x_scalar,
x_bytes,
transcript_context,
} => {
let pake2 = Pake2::decode(bytes)?;
let (z_bytes, v_bytes) = compute_z_v_prover(&x_scalar, &w0, &w1, &pake2.y)?;
let t_t = transcript_hash(
&transcript_context,
&x_bytes,
&pake2.y,
&z_bytes,
&v_bytes,
&w0,
);
let (mut ka, mut ke) = ka_ke_from_transcript(&t_t);
let (kca, kcb) = derive_confirmation_keys(&ka)?;
ka.zeroize();
let cb_expected = compute_cb(&kcb, &x_bytes);
verify_tag(&cb_expected, &pake2.verifier)?;
let ca = compute_ca(&kca, &pake2.y);
let session_keys_blob = derive_session_keys(&ke)?;
let session_keys = build_session_keys(ke, &session_keys_blob);
ke.zeroize();
self.state = State::ReadyToSendPake3 { ca, session_keys };
Ok(())
}
other => {
self.state = other;
Err(Error::UnexpectedMessage {
expected: PaseMessageKind::Pake2,
got: PaseMessageKind::Pake2,
})
}
}
}
pub fn next_message(&mut self) -> Result<Vec<u8>> {
let prev = std::mem::replace(&mut self.state, State::Poisoned);
match prev {
State::ReadyToSendPake1 {
pin,
params,
x_scalar,
transcript_context,
} => {
let (w0, w1) = derive_w0_w1(pin, ¶ms.salt, params.iterations)?;
let x_bytes = compute_x(&x_scalar, &w0);
let pake1_bytes = Pake1 { x: x_bytes }.encode()?;
self.state = State::AwaitingPake2 {
w0,
w1,
x_scalar,
x_bytes,
transcript_context,
};
Ok(pake1_bytes)
}
State::ReadyToSendPake3 { ca, session_keys } => {
let pake3_bytes = Pake3 { verifier: ca }.encode()?;
self.state = State::Complete { session_keys };
Ok(pake3_bytes)
}
other => {
self.state = other;
Err(Error::UnexpectedMessage {
expected: PaseMessageKind::Pake1,
got: PaseMessageKind::Pake1,
})
}
}
}
pub fn finish(self) -> Result<PaseSessionKeys> {
match self.state {
State::Complete { session_keys } => Ok(session_keys),
_ => Err(Error::HandshakeIncomplete),
}
}
}
fn build_session_keys(ke: [u8; 16], blob_48: &[u8; 48]) -> PaseSessionKeys {
let mut i2r_key = [0u8; 16];
let mut r2i_key = [0u8; 16];
let mut attestation_key = [0u8; 16];
i2r_key.copy_from_slice(&blob_48[0..16]);
r2i_key.copy_from_slice(&blob_48[16..32]);
attestation_key.copy_from_slice(&blob_48[32..48]);
PaseSessionKeys {
ke,
i2r_key,
r2i_key,
attestation_key,
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)] mod tests {
use super::*;
use crate::pase::messages::{PbkdfParamResponse, PbkdfParamsInner};
#[test]
fn new_with_negotiation_accepts_any_pin() {
let _ = PaseProver::new_with_negotiation(20_202_021, 0x0001).unwrap();
let _ = PaseProver::new_with_negotiation(0, 0x0001).unwrap();
let _ = PaseProver::new_with_negotiation(u32::MAX, 0x0001).unwrap();
}
#[test]
fn new_with_known_params_rejects_low_iterations() {
let params = PasePbkdfParams {
iterations: 999,
salt: vec![0u8; 16],
};
assert!(matches!(
PaseProver::new_with_known_params(20_202_021, params, 0x0001),
Err(Error::PbkdfIterationsTooLow(999))
));
}
#[test]
fn new_with_known_params_rejects_short_salt() {
let params = PasePbkdfParams {
iterations: 1_000,
salt: vec![0u8; 15],
};
assert!(matches!(
PaseProver::new_with_known_params(20_202_021, params, 0x0001),
Err(Error::PbkdfSaltLengthInvalid(15))
));
}
#[test]
fn new_with_known_params_accepts_valid_params() {
let params = PasePbkdfParams {
iterations: 1_000,
salt: vec![0x42u8; 16],
};
let _ = PaseProver::new_with_known_params(20_202_021, params, 0x0001).unwrap();
}
#[test]
fn start_negotiation_emits_tlv_structure() {
let mut prover = PaseProver::new_with_negotiation(20_202_021, 0x0001).unwrap();
let bytes = prover.start().unwrap();
assert_eq!(
bytes[0], 0x15,
"first byte must be anonymous structure tag 0x15"
);
assert!(!bytes.is_empty());
}
#[test]
fn expected_inbound_after_start_negotiation_is_pbkdf_response() {
let mut prover = PaseProver::new_with_negotiation(20_202_021, 0x0001).unwrap();
let _ = prover.start().unwrap();
assert_eq!(
prover.expected_inbound(),
Some(PaseMessageKind::PbkdfParamResponse)
);
}
#[test]
fn handle_pbkdf_response_advances_to_ready_to_send_pake1() {
let mut prover = PaseProver::new_with_negotiation(20_202_021, 0x0001).unwrap();
let _req_bytes = prover.start().unwrap();
let resp = PbkdfParamResponse {
initiator_random: [0x42u8; 32],
responder_random: [0x11u8; 32],
responder_session_id: 1,
pbkdf_parameters: Some(PbkdfParamsInner {
iterations: 1_000,
salt: vec![0xABu8; 16],
}),
responder_session_params: None,
};
let resp_bytes = resp.encode().unwrap();
prover.handle_pbkdf_response(&resp_bytes).unwrap();
assert_eq!(prover.expected_inbound(), None);
}
#[test]
fn handle_pbkdf_response_rejects_missing_pbkdf_params() {
let mut prover = PaseProver::new_with_negotiation(20_202_021, 0x0001).unwrap();
let _ = prover.start().unwrap();
let resp = PbkdfParamResponse {
initiator_random: [0x42u8; 32],
responder_random: [0x11u8; 32],
responder_session_id: 1,
pbkdf_parameters: None, responder_session_params: None,
};
let resp_bytes = resp.encode().unwrap();
assert!(matches!(
prover.handle_pbkdf_response(&resp_bytes),
Err(Error::InvalidParameter)
));
}
#[test]
fn next_message_after_pbkdf_response_emits_pake1() {
let mut prover = PaseProver::new_with_negotiation(20_202_021, 0x0001).unwrap();
let _ = prover.start().unwrap();
let resp = PbkdfParamResponse {
initiator_random: [0x42u8; 32],
responder_random: [0x11u8; 32],
responder_session_id: 1,
pbkdf_parameters: Some(PbkdfParamsInner {
iterations: 1_000,
salt: vec![0xABu8; 16],
}),
responder_session_params: None,
};
prover
.handle_pbkdf_response(&resp.encode().unwrap())
.unwrap();
let pake1_bytes = prover.next_message().unwrap();
assert_eq!(pake1_bytes[0], 0x15);
assert_eq!(prover.expected_inbound(), Some(PaseMessageKind::Pake2));
}
#[test]
fn finish_before_complete_returns_handshake_incomplete() {
let prover = PaseProver::new_with_negotiation(20_202_021, 0x0001).unwrap();
assert!(matches!(prover.finish(), Err(Error::HandshakeIncomplete)));
}
#[test]
fn out_of_order_handle_pake2_returns_unexpected_message() {
let mut prover = PaseProver::new_with_negotiation(20_202_021, 0x0001).unwrap();
let dummy_pake2 = Pake2 {
y: [0x04u8; 65],
verifier: [0x00u8; 32],
};
let pake2_bytes = dummy_pake2.encode().unwrap();
assert!(matches!(
prover.handle_pake2(&pake2_bytes),
Err(Error::UnexpectedMessage { .. })
));
}
#[test]
fn prover_advertises_initiator_session_id_and_starts_unknowing_responder_id() {
let mut prover = PaseProver::new_with_negotiation(20_202_021, 0x0011).unwrap();
assert_eq!(prover.responder_session_id(), None);
let req = prover.start().unwrap();
let decoded_req = crate::pase::messages::PbkdfParamRequest::decode(&req).unwrap();
assert_eq!(decoded_req.initiator_session_id, 0x0011);
assert_eq!(prover.responder_session_id(), None);
}
}