use std::collections::HashMap;
use std::io;
use std::sync::{LazyLock, Mutex};
use constant_time_eq::constant_time_eq_32;
use iroh::endpoint::{Connection, RecvStream};
use spake2::{Ed25519Group, Identity, Password, Spake2};
use zeroize::Zeroizing;
const NO_PASS: u8 = 0;
const PAKE_REQUIRED: u8 = 1;
const VERDICT_OK: u8 = 1;
const VERDICT_REJECT: u8 = 0;
const MAX_PAKE_MSG: usize = 64;
const PAKE_IDENTITY: &[u8] = b"koh-pake-v1";
const SERVER_CONFIRM_LABEL: &[u8] = b"server";
const CLIENT_CONFIRM_LABEL: &[u8] = b"client";
#[expect(
clippy::expect_used,
reason = "Params::new only errors on out-of-range values; these are compile-time constants"
)]
fn kdf_params() -> argon2::Params {
argon2::Params::new(64 * 1024, 3, 1, Some(32)).expect("static Argon2id params are in range")
}
fn kdf_salt() -> [u8; 16] {
let mut salt = [0u8; 16];
salt.copy_from_slice(&blake3::hash(b"koh-pass-kdf-v1").as_bytes()[..16]);
salt
}
#[expect(
clippy::expect_used,
reason = "hash_password_into only errors on invalid params/output-len, fixed valid here"
)]
fn derive_psk(passphrase: &str) -> Zeroizing<[u8; 32]> {
let argon = argon2::Argon2::new(
argon2::Algorithm::Argon2id,
argon2::Version::V0x13,
kdf_params(),
);
let mut psk = Zeroizing::new([0u8; 32]);
argon
.hash_password_into(passphrase.as_bytes(), &kdf_salt(), psk.as_mut_slice())
.expect("Argon2id derivation with valid static params and a 32-byte output cannot fail");
psk
}
type PskCache = HashMap<[u8; 32], Zeroizing<[u8; 32]>>;
static PSK_CACHE: LazyLock<Mutex<PskCache>> = LazyLock::new(|| Mutex::new(HashMap::new()));
#[expect(
clippy::expect_used,
reason = "a poisoned cache mutex is a panic-elsewhere bug, not peer-influenced input"
)]
fn cached_psk(passphrase: &str) -> Zeroizing<[u8; 32]> {
let key = *blake3::hash(passphrase.as_bytes()).as_bytes();
{
let cache = PSK_CACHE.lock().expect("PSK cache mutex poisoned");
if let Some(psk) = cache.get(&key) {
return psk.clone();
}
}
let psk = derive_psk(passphrase);
PSK_CACHE
.lock()
.expect("PSK cache mutex poisoned")
.insert(key, psk.clone());
psk
}
pub fn prewarm_psk(passphrase: &str) {
if !passphrase.is_empty() {
let _ = cached_psk(passphrase);
}
}
fn confirm_tag(shared_key: &[u8], label: &[u8], server_msg: &[u8], client_msg: &[u8]) -> [u8; 32] {
let mac_key = blake3::derive_key("koh-pake-v1 key confirmation", shared_key);
let mut h = blake3::Hasher::new_keyed(&mac_key);
for field in [label, server_msg, client_msg] {
h.update(&(field.len() as u64).to_le_bytes());
h.update(field);
}
*h.finalize().as_bytes()
}
#[derive(Debug, thiserror::Error)]
pub enum AuthError {
#[error("auth stream error: {0}")]
Stream(#[from] io::Error),
#[error("passphrase authentication failed")]
ChallengeFailed,
}
fn write_msg(out: &mut Vec<u8>, msg: &[u8]) -> Result<(), AuthError> {
let len = u8::try_from(msg.len())
.ok()
.filter(|&l| l != 0 && usize::from(l) <= MAX_PAKE_MSG)
.ok_or(AuthError::ChallengeFailed)?;
out.push(len);
out.extend_from_slice(msg);
Ok(())
}
async fn read_msg(recv: &mut RecvStream) -> Result<Vec<u8>, AuthError> {
let mut len = [0u8; 1];
recv.read_exact(&mut len).await.map_err(io::Error::other)?;
let n = usize::from(len[0]);
if n == 0 || n > MAX_PAKE_MSG {
return Err(AuthError::ChallengeFailed);
}
let mut buf = vec![0u8; n];
recv.read_exact(&mut buf).await.map_err(io::Error::other)?;
Ok(buf)
}
pub async fn handshake_server(
conn: &Connection,
passphrase: Option<&str>,
) -> Result<(), AuthError> {
let (mut send, mut recv) = conn.open_bi().await.map_err(io::Error::other)?;
let Some(pass) = passphrase else {
send.write_all(&[NO_PASS]).await.map_err(io::Error::other)?;
let _ = send.finish();
return Ok(());
};
let psk = cached_psk(pass);
let (state, server_msg) = Spake2::<Ed25519Group>::start_symmetric(
&Password::new(psk.as_slice()),
&Identity::new(PAKE_IDENTITY),
);
let mut out = vec![PAKE_REQUIRED];
write_msg(&mut out, &server_msg)?;
send.write_all(&out).await.map_err(io::Error::other)?;
let client_msg = read_msg(&mut recv).await?;
let shared = Zeroizing::new(
state
.finish(&client_msg)
.map_err(|_| AuthError::ChallengeFailed)?,
);
let server_conf = confirm_tag(&shared, SERVER_CONFIRM_LABEL, &server_msg, &client_msg);
let expect_client_conf = confirm_tag(&shared, CLIENT_CONFIRM_LABEL, &server_msg, &client_msg);
send.write_all(&server_conf)
.await
.map_err(io::Error::other)?;
let mut client_conf = [0u8; 32];
recv.read_exact(&mut client_conf)
.await
.map_err(io::Error::other)?;
let ok = constant_time_eq_32(&client_conf, &expect_client_conf);
let verdict = if ok { VERDICT_OK } else { VERDICT_REJECT };
send.write_all(&[verdict]).await.map_err(io::Error::other)?;
let _ = send.finish();
if ok {
Ok(())
} else {
Err(AuthError::ChallengeFailed)
}
}
pub async fn handshake_client(
conn: &Connection,
passphrase: Option<&str>,
) -> Result<(), AuthError> {
let (mut send, mut recv) = conn.accept_bi().await.map_err(io::Error::other)?;
let mut tag = [0u8; 1];
recv.read_exact(&mut tag).await.map_err(io::Error::other)?;
match tag[0] {
PAKE_REQUIRED => {
let server_msg = read_msg(&mut recv).await?;
let psk = cached_psk(passphrase.unwrap_or(""));
let (state, client_msg) = Spake2::<Ed25519Group>::start_symmetric(
&Password::new(psk.as_slice()),
&Identity::new(PAKE_IDENTITY),
);
let mut out = Vec::new();
write_msg(&mut out, &client_msg)?;
send.write_all(&out).await.map_err(io::Error::other)?;
let shared = Zeroizing::new(
state
.finish(&server_msg)
.map_err(|_| AuthError::ChallengeFailed)?,
);
let expect_server_conf =
confirm_tag(&shared, SERVER_CONFIRM_LABEL, &server_msg, &client_msg);
let mut server_conf = [0u8; 32];
recv.read_exact(&mut server_conf)
.await
.map_err(io::Error::other)?;
let server_ok = constant_time_eq_32(&server_conf, &expect_server_conf);
let client_conf = confirm_tag(&shared, CLIENT_CONFIRM_LABEL, &server_msg, &client_msg);
send.write_all(&client_conf)
.await
.map_err(io::Error::other)?;
let _ = send.finish();
let mut verdict = [0u8; 1];
recv.read_exact(&mut verdict)
.await
.map_err(io::Error::other)?;
if !server_ok || verdict[0] != VERDICT_OK {
return Err(AuthError::ChallengeFailed);
}
}
NO_PASS => {
if passphrase.is_some_and(|p| !p.is_empty()) {
return Err(AuthError::ChallengeFailed);
}
}
_ => return Err(AuthError::ChallengeFailed),
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn auth_error_variants_are_constructible_and_reachable() {
let rejected = AuthError::ChallengeFailed;
assert_eq!(rejected.to_string(), "passphrase authentication failed");
let io_err = io::Error::new(io::ErrorKind::UnexpectedEof, "stream closed");
let stream: AuthError = io_err.into();
assert!(matches!(stream, AuthError::Stream(_)));
assert!(stream.to_string().contains("auth stream error"));
let absorbed: anyhow::Error = AuthError::ChallengeFailed.into();
assert!(absorbed.to_string().contains("authentication failed"));
}
#[test]
fn derive_psk_is_deterministic_and_both_peers_agree() {
let server_k = derive_psk("correct horse battery staple");
let client_k = derive_psk("correct horse battery staple");
assert_eq!(
*server_k, *client_k,
"both peers must derive the same value"
);
assert_eq!(*cached_psk("correct horse battery staple"), *server_k);
assert_ne!(
*server_k,
*derive_psk("wrong horse"),
"distinct passphrases -> distinct maps"
);
assert_eq!(kdf_salt(), kdf_salt());
}
fn pake_mutual_confirms(server_pass: &str, client_pass: &str) -> bool {
let s_psk = cached_psk(server_pass);
let c_psk = cached_psk(client_pass);
let (s_state, s_msg) = Spake2::<Ed25519Group>::start_symmetric(
&Password::new(s_psk.as_slice()),
&Identity::new(PAKE_IDENTITY),
);
let (c_state, c_msg) = Spake2::<Ed25519Group>::start_symmetric(
&Password::new(c_psk.as_slice()),
&Identity::new(PAKE_IDENTITY),
);
let s_shared = s_state.finish(&c_msg).expect("server finish");
let c_shared = c_state.finish(&s_msg).expect("client finish");
let client_verifies_server = constant_time_eq_32(
&confirm_tag(&s_shared, SERVER_CONFIRM_LABEL, &s_msg, &c_msg),
&confirm_tag(&c_shared, SERVER_CONFIRM_LABEL, &s_msg, &c_msg),
);
let server_verifies_client = constant_time_eq_32(
&confirm_tag(&c_shared, CLIENT_CONFIRM_LABEL, &s_msg, &c_msg),
&confirm_tag(&s_shared, CLIENT_CONFIRM_LABEL, &s_msg, &c_msg),
);
client_verifies_server && server_verifies_client
}
#[test]
fn matching_passphrases_mutually_confirm_and_mismatched_do_not() {
assert!(
pake_mutual_confirms("hunter2", "hunter2"),
"matching passphrases must mutually confirm"
);
assert!(
!pake_mutual_confirms("hunter2", "nope"),
"a wrong passphrase must fail confirmation"
);
assert!(
!pake_mutual_confirms("", "secret"),
"empty vs set passphrase must fail confirmation"
);
}
#[test]
fn confirm_tag_is_direction_separated_and_transcript_bound() {
let key = [7u8; 32];
let (a, b) = (b"AAAA".as_slice(), b"BBBB".as_slice());
let srv = confirm_tag(&key, SERVER_CONFIRM_LABEL, a, b);
let cli = confirm_tag(&key, CLIENT_CONFIRM_LABEL, a, b);
assert_ne!(srv, cli, "direction labels must separate the tags");
let other_transcript = confirm_tag(&key, SERVER_CONFIRM_LABEL, b, a);
assert_ne!(
srv, other_transcript,
"the tag must bind the (ordered) transcript"
);
}
#[test]
fn framing_round_trips_and_rejects_bad_lengths() {
let mut buf = Vec::new();
write_msg(&mut buf, &[1, 2, 3]).expect("small message frames");
assert_eq!(buf, vec![3, 1, 2, 3], "length prefix then bytes");
assert!(
write_msg(&mut Vec::new(), &[0u8; MAX_PAKE_MSG + 1]).is_err(),
"an over-cap message must be refused, not truncated"
);
assert!(
write_msg(&mut Vec::new(), &[]).is_err(),
"a zero-length message must be refused"
);
}
}