use p256::elliptic_curve::group::ff::PrimeField; use ring::rand::{SecureRandom, SystemRandom};
use crate::error::{Error, Result};
use crate::pase::kdf::{derive_l, derive_w0_w1, validate_params};
use crate::pase::messages::{
Pake1, Pake2, Pake3, PbkdfParamRequest, PbkdfParamResponse, PbkdfParamsInner,
};
use crate::pase::spake2plus::{
compute_ca, compute_cb, compute_y, compute_z_v_verifier, 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 {
AwaitingFirstMessage {
w0: p256::Scalar,
l: [u8; 65],
params: PasePbkdfParams,
y_scalar: p256::Scalar,
responder_random: [u8; 32],
responder_session_id: u16,
},
ReadyToSendPbkdfResponse {
w0: p256::Scalar,
l: [u8; 65],
params: PasePbkdfParams,
y_scalar: p256::Scalar,
request_bytes: Vec<u8>,
responder_random: [u8; 32],
initiator_random: [u8; 32],
responder_session_id: u16,
},
AwaitingPake1 {
w0: p256::Scalar,
l: [u8; 65],
y_scalar: p256::Scalar,
transcript_context: [u8; 32],
},
ReadyToSendPake2 {
y_bytes: [u8; 65],
cb: [u8; 32],
ca_expected: [u8; 32],
session_keys: PaseSessionKeys,
},
Complete { session_keys: PaseSessionKeys },
Poisoned,
}
pub struct PaseVerifier {
state: State,
}
impl PaseVerifier {
pub fn new(
w0: [u8; 32],
l: [u8; 65],
params: PasePbkdfParams,
responder_session_id: u16,
) -> Result<Self> {
validate_params(params.iterations, ¶ms.salt)?;
let rng = SystemRandom::new();
Self::new_using_rng(w0, l, params, responder_session_id, &rng)
}
pub(crate) fn new_using_rng(
w0_bytes: [u8; 32],
l: [u8; 65],
params: PasePbkdfParams,
responder_session_id: u16,
rng: &dyn SecureRandom,
) -> Result<Self> {
validate_params(params.iterations, ¶ms.salt)?;
let w0_opt: Option<p256::Scalar> =
p256::Scalar::from_repr(p256::FieldBytes::from(w0_bytes)).into();
let w0 = w0_opt.ok_or(Error::InvalidScalar)?;
if bool::from(p256::elliptic_curve::group::ff::Field::is_zero(&w0)) {
return Err(Error::InvalidScalar);
}
let y_scalar = sample_scalar(rng)?;
let mut responder_random = [0u8; 32];
rng.fill(&mut responder_random)
.map_err(|_| Error::PinDerivationFailed)?;
Ok(Self {
state: State::AwaitingFirstMessage {
w0,
l,
params,
y_scalar,
responder_random,
responder_session_id,
},
})
}
pub fn new_from_pin(
pin: u32,
params: PasePbkdfParams,
responder_session_id: u16,
) -> Result<Self> {
let rng = SystemRandom::new();
Self::new_from_pin_using_rng(pin, params, responder_session_id, &rng)
}
pub(crate) fn new_from_pin_using_rng(
pin: u32,
params: PasePbkdfParams,
responder_session_id: u16,
rng: &dyn SecureRandom,
) -> Result<Self> {
let (w0_scalar, w1_scalar) = derive_w0_w1(pin, ¶ms.salt, params.iterations)?;
let l = derive_l(&w1_scalar);
let w0_be: p256::FieldBytes = w0_scalar.to_bytes();
let mut w0_arr = [0u8; 32];
w0_arr.copy_from_slice(&w0_be);
Self::new_using_rng(w0_arr, l, params, responder_session_id, rng)
}
pub(crate) fn new_with_scalar(
w0_bytes: [u8; 32],
l: [u8; 65],
params: PasePbkdfParams,
y_scalar_bytes: [u8; 32],
) -> Result<Self> {
Self::new_with_scalar_and_random(w0_bytes, l, params, y_scalar_bytes, [0u8; 32], 0)
}
pub(crate) fn new_with_scalar_and_random(
w0_bytes: [u8; 32],
l: [u8; 65],
params: PasePbkdfParams,
y_scalar_bytes: [u8; 32],
responder_random: [u8; 32],
responder_session_id: u16,
) -> Result<Self> {
use p256::elliptic_curve::group::ff::Field;
validate_params(params.iterations, ¶ms.salt)?;
let w0_opt: Option<p256::Scalar> =
p256::Scalar::from_repr(p256::FieldBytes::from(w0_bytes)).into();
let w0 = w0_opt.ok_or(Error::InvalidScalar)?;
if bool::from(w0.is_zero()) {
return Err(Error::InvalidScalar);
}
let y_opt: Option<p256::Scalar> =
p256::Scalar::from_repr(p256::FieldBytes::from(y_scalar_bytes)).into();
let y_scalar = y_opt.ok_or(Error::InvalidScalar)?;
if bool::from(y_scalar.is_zero()) {
return Err(Error::InvalidScalar);
}
Ok(Self {
state: State::AwaitingFirstMessage {
w0,
l,
params,
y_scalar,
responder_random,
responder_session_id,
},
})
}
pub(crate) fn new_from_pin_with_scalar(
pin: u32,
params: PasePbkdfParams,
y_scalar_bytes: [u8; 32],
) -> Result<Self> {
let (w0_scalar, w1_scalar) = derive_w0_w1(pin, ¶ms.salt, params.iterations)?;
let l = derive_l(&w1_scalar);
let w0_be: p256::FieldBytes = w0_scalar.to_bytes();
let mut w0_arr = [0u8; 32];
w0_arr.copy_from_slice(&w0_be);
Self::new_with_scalar(w0_arr, l, params, y_scalar_bytes)
}
pub fn expected_inbound(&self) -> Option<PaseMessageKind> {
match &self.state {
State::AwaitingFirstMessage { .. } => {
Some(PaseMessageKind::PbkdfParamRequest)
}
State::AwaitingPake1 { .. } => Some(PaseMessageKind::Pake1),
State::ReadyToSendPake2 { .. } => Some(PaseMessageKind::Pake3),
_ => None,
}
}
pub fn handle_pbkdf_request(&mut self, bytes: &[u8]) -> Result<()> {
let prev = std::mem::replace(&mut self.state, State::Poisoned);
match prev {
State::AwaitingFirstMessage {
w0,
l,
params,
y_scalar,
responder_random,
responder_session_id,
} => {
let req = PbkdfParamRequest::decode(bytes)?;
self.state = State::ReadyToSendPbkdfResponse {
w0,
l,
params,
y_scalar,
request_bytes: bytes.to_vec(),
responder_random,
initiator_random: req.initiator_random,
responder_session_id,
};
Ok(())
}
other => {
self.state = other;
Err(Error::UnexpectedMessage {
expected: PaseMessageKind::PbkdfParamRequest,
got: PaseMessageKind::PbkdfParamRequest,
})
}
}
}
pub fn handle_pake1(&mut self, bytes: &[u8]) -> Result<()> {
let prev = std::mem::replace(&mut self.state, State::Poisoned);
match prev {
State::AwaitingFirstMessage {
w0,
l,
y_scalar,
..
} => {
let transcript_context = hash_context(&[]);
self.state = State::Poisoned; self.compute_pake2(w0, l, y_scalar, transcript_context, bytes)
}
State::AwaitingPake1 {
w0,
l,
y_scalar,
transcript_context,
} => {
self.state = State::Poisoned; self.compute_pake2(w0, l, y_scalar, transcript_context, bytes)
}
other => {
self.state = other;
Err(Error::UnexpectedMessage {
expected: PaseMessageKind::Pake1,
got: PaseMessageKind::Pake1,
})
}
}
}
pub fn handle_pake3(&mut self, bytes: &[u8]) -> Result<()> {
let prev = std::mem::replace(&mut self.state, State::Poisoned);
match prev {
State::ReadyToSendPake2 {
y_bytes: _,
cb: _,
ca_expected,
session_keys,
} => {
let pake3 = Pake3::decode(bytes)?;
verify_tag(&ca_expected, &pake3.verifier)?;
self.state = State::Complete { session_keys };
Ok(())
}
other => {
self.state = other;
Err(Error::UnexpectedMessage {
expected: PaseMessageKind::Pake3,
got: PaseMessageKind::Pake3,
})
}
}
}
pub fn next_message(&mut self) -> Result<Vec<u8>> {
let prev = std::mem::replace(&mut self.state, State::Poisoned);
match prev {
State::ReadyToSendPbkdfResponse {
w0,
l,
params,
y_scalar,
request_bytes,
responder_random,
initiator_random,
responder_session_id,
} => {
let resp = PbkdfParamResponse {
initiator_random,
responder_random,
responder_session_id,
pbkdf_parameters: Some(PbkdfParamsInner {
iterations: params.iterations,
salt: params.salt.clone(),
}),
responder_session_params: None,
};
let resp_bytes = resp.encode()?;
let transcript_context = hash_context(&[&request_bytes, &resp_bytes]);
self.state = State::AwaitingPake1 {
w0,
l,
y_scalar,
transcript_context,
};
Ok(resp_bytes)
}
State::ReadyToSendPake2 {
y_bytes,
cb,
ca_expected,
session_keys,
} => {
let pake2_bytes = Pake2 {
y: y_bytes,
verifier: cb,
}
.encode()?;
self.state = State::ReadyToSendPake2 {
y_bytes,
cb,
ca_expected,
session_keys,
};
Ok(pake2_bytes)
}
other => {
self.state = other;
Err(Error::UnexpectedMessage {
expected: PaseMessageKind::PbkdfParamResponse,
got: PaseMessageKind::PbkdfParamResponse,
})
}
}
}
pub fn finish(self) -> Result<PaseSessionKeys> {
match self.state {
State::Complete { session_keys } => Ok(session_keys),
_ => Err(Error::HandshakeIncomplete),
}
}
}
impl PaseVerifier {
#[allow(clippy::similar_names)]
fn compute_pake2(
&mut self,
w0: p256::Scalar,
l: [u8; 65],
y_scalar: p256::Scalar,
transcript_context: [u8; 32],
pake1_bytes: &[u8],
) -> Result<()> {
let pake1 = Pake1::decode(pake1_bytes)?;
let y_bytes = compute_y(&y_scalar, &w0);
let (z_bytes, v_bytes) = compute_z_v_verifier(&y_scalar, &w0, &l, &pake1.x)?;
let t_t = transcript_hash(
&transcript_context,
&pake1.x,
&y_bytes,
&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 = compute_cb(&kcb, &pake1.x);
let ca_expected = compute_ca(&kca, &y_bytes);
let session_keys_blob = derive_session_keys(&ke)?;
let session_keys = build_session_keys(ke, &session_keys_blob);
ke.zeroize();
self.state = State::ReadyToSendPake2 {
y_bytes,
cb,
ca_expected,
session_keys,
};
Ok(())
}
}
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;
fn test_params() -> PasePbkdfParams {
PasePbkdfParams {
iterations: 1_000,
salt: vec![0x42u8; 16],
}
}
const TEST_PIN: u32 = 20_202_021;
#[test]
fn new_from_pin_accepts_valid_params() {
let _ = PaseVerifier::new_from_pin(TEST_PIN, test_params(), 0x0033).unwrap();
}
#[test]
fn new_from_pin_rejects_low_iterations() {
let params = PasePbkdfParams {
iterations: 999,
salt: vec![0u8; 16],
};
assert!(matches!(
PaseVerifier::new_from_pin(TEST_PIN, params, 0x0033),
Err(Error::PbkdfIterationsTooLow(999))
));
}
#[test]
fn new_from_pin_rejects_short_salt() {
let params = PasePbkdfParams {
iterations: 1_000,
salt: vec![0u8; 15],
};
assert!(matches!(
PaseVerifier::new_from_pin(TEST_PIN, params, 0x0033),
Err(Error::PbkdfSaltLengthInvalid(15))
));
}
#[test]
fn new_raw_rejects_invalid_w0_scalar() {
let w0_zero = [0u8; 32];
let l = [0x04u8; 65];
let params = test_params();
assert!(matches!(
PaseVerifier::new(w0_zero, l, params, 0x0033),
Err(Error::InvalidScalar)
));
}
#[test]
fn expected_inbound_after_construction_is_pbkdf_request() {
let v = PaseVerifier::new_from_pin(TEST_PIN, test_params(), 0x0033).unwrap();
assert_eq!(
v.expected_inbound(),
Some(PaseMessageKind::PbkdfParamRequest)
);
}
#[test]
fn handle_pbkdf_request_advances_state() {
let mut v = PaseVerifier::new_from_pin(TEST_PIN, test_params(), 0x0033).unwrap();
let req = PbkdfParamRequest {
initiator_random: [0x11u8; 32],
initiator_session_id: 0,
passcode_id: 0,
has_pbkdf_parameters: false,
initiator_session_params: None,
};
let req_bytes = req.encode().unwrap();
v.handle_pbkdf_request(&req_bytes).unwrap();
assert_eq!(v.expected_inbound(), None);
}
#[test]
fn next_message_after_pbkdf_request_emits_response() {
let mut v = PaseVerifier::new_from_pin(TEST_PIN, test_params(), 0x0033).unwrap();
let req = PbkdfParamRequest {
initiator_random: [0x11u8; 32],
initiator_session_id: 0,
passcode_id: 0,
has_pbkdf_parameters: false,
initiator_session_params: None,
};
v.handle_pbkdf_request(&req.encode().unwrap()).unwrap();
let resp_bytes = v.next_message().unwrap();
assert_eq!(resp_bytes[0], 0x15, "first byte must be 0x15 (anon struct)");
assert_eq!(v.expected_inbound(), Some(PaseMessageKind::Pake1));
let decoded = PbkdfParamResponse::decode(&resp_bytes).unwrap();
assert!(decoded.pbkdf_parameters.is_some());
let inner = decoded.pbkdf_parameters.unwrap();
assert_eq!(inner.iterations, 1_000);
assert_eq!(inner.salt, vec![0x42u8; 16]);
}
#[test]
fn out_of_order_handle_pake3_returns_unexpected_message() {
let mut v = PaseVerifier::new_from_pin(TEST_PIN, test_params(), 0x0033).unwrap();
let dummy_pake3 = Pake3 {
verifier: [0x00u8; 32],
};
let pake3_bytes = dummy_pake3.encode().unwrap();
assert!(matches!(
v.handle_pake3(&pake3_bytes),
Err(Error::UnexpectedMessage { .. })
));
}
#[test]
fn finish_before_complete_returns_handshake_incomplete() {
let v = PaseVerifier::new_from_pin(TEST_PIN, test_params(), 0x0033).unwrap();
assert!(matches!(v.finish(), Err(Error::HandshakeIncomplete)));
}
#[test]
fn handle_pake3_rejects_wrong_ca_tag() {
use crate::pase::kdf::derive_w0_w1;
use crate::pase::spake2plus::{compute_x, sample_scalar};
use ring::rand::SystemRandom;
let rng = SystemRandom::new();
let params = test_params();
let mut v = PaseVerifier::new_from_pin(TEST_PIN, params.clone(), 0x0033).unwrap();
let (w0_scalar, _w1_scalar) =
derive_w0_w1(TEST_PIN, ¶ms.salt, params.iterations).unwrap();
let x_scalar = sample_scalar(&rng).unwrap();
let x_bytes = compute_x(&x_scalar, &w0_scalar);
let pake1_bytes = Pake1 { x: x_bytes }.encode().unwrap();
v.handle_pake1(&pake1_bytes).unwrap();
let _pake2_bytes = v.next_message().unwrap();
let wrong_pake3 = Pake3 {
verifier: [0x00u8; 32],
};
let wrong_pake3_bytes = wrong_pake3.encode().unwrap();
assert!(matches!(
v.handle_pake3(&wrong_pake3_bytes),
Err(Error::ConfirmationTagMismatch)
));
}
#[test]
fn verifier_advertises_responder_session_id() {
let params = test_params();
let mut verifier = PaseVerifier::new_from_pin(TEST_PIN, params, 0x0033).unwrap();
let mut prover =
crate::pase::prover::PaseProver::new_with_negotiation(TEST_PIN, 0x0001).unwrap();
let req = prover.start().unwrap();
verifier.handle_pbkdf_request(&req).unwrap();
let resp = verifier.next_message().unwrap();
let decoded = crate::pase::messages::PbkdfParamResponse::decode(&resp).unwrap();
assert_eq!(decoded.responder_session_id, 0x0033);
}
#[test]
fn handle_pake1_as_first_message_succeeds() {
use crate::pase::kdf::derive_w0_w1;
use crate::pase::spake2plus::{compute_x, sample_scalar};
use ring::rand::SystemRandom;
let rng = SystemRandom::new();
let params = test_params();
let mut v = PaseVerifier::new_from_pin(TEST_PIN, params.clone(), 0x0033).unwrap();
let (w0_scalar, _) = derive_w0_w1(TEST_PIN, ¶ms.salt, params.iterations).unwrap();
let x_scalar = sample_scalar(&rng).unwrap();
let x_bytes = compute_x(&x_scalar, &w0_scalar);
let pake1_bytes = Pake1 { x: x_bytes }.encode().unwrap();
v.handle_pake1(&pake1_bytes).unwrap();
assert_eq!(v.expected_inbound(), Some(PaseMessageKind::Pake3));
let pake2_bytes = v.next_message().unwrap();
assert_eq!(pake2_bytes[0], 0x15, "Pake2 must be anon TLV structure");
let decoded = Pake2::decode(&pake2_bytes).unwrap();
assert_eq!(
decoded.y[0], 0x04,
"Y must have SEC1 uncompressed prefix 0x04"
);
assert_eq!(decoded.verifier.len(), 32);
}
}