use ed25519_dalek::{Signature, VerifyingKey};
use std::collections::BTreeSet;
use std::time::{SystemTime, UNIX_EPOCH};
use crate::ServerError;
use crate::config::PassConfig;
pub const PASS_VERSION_V1: u8 = 1;
const SIGNATURE_LEN: usize = 64;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PassPrincipal {
pub participant: Vec<u8>,
pub public_key: [u8; 32],
pub conversations: BTreeSet<u64>,
pub live: String,
pub may_enroll: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WirePassV1 {
pub participant: Vec<u8>,
pub public_key: [u8; 32],
pub conversations: Vec<u64>,
pub live: String,
pub may_enroll: bool,
pub issued_at: u64,
pub expires_at: u64,
pub signature: [u8; SIGNATURE_LEN],
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PassCheckError {
Malformed,
Signature,
Expired,
}
impl PassCheckError {
#[must_use]
pub const fn check_name(self) -> &'static str {
match self {
Self::Malformed => "malformed",
Self::Signature => "signature",
Self::Expired => "expired",
}
}
}
#[derive(Clone, Debug)]
pub(crate) struct PassVerifier {
key: VerifyingKey,
skew: u64,
}
impl PassVerifier {
pub(crate) fn from_config(config: &PassConfig) -> Result<Self, ServerError> {
let key = hex::decode(&config.registry_verifying_key).map_err(|error| {
ServerError::ConfigValidation {
message: format!(
"auth.pass.registry_verifying_key: expected 64 hexadecimal digits: {error}"
),
}
})?;
Self::new(&key, config.maximum_clock_skew_seconds).map_err(|_| {
ServerError::ConfigValidation {
message: "auth.pass.registry_verifying_key: expected a valid 32-byte Ed25519 verifying key".to_owned(),
}
})
}
pub(crate) fn new(key: &[u8], skew: u64) -> Result<Self, PassCheckError> {
let key: [u8; 32] = key.try_into().map_err(|_| PassCheckError::Malformed)?;
let key = VerifyingKey::from_bytes(&key).map_err(|_| PassCheckError::Malformed)?;
Ok(Self { key, skew })
}
pub(crate) fn verify(&self, bytes: &[u8]) -> Result<PassPrincipal, PassCheckError> {
WirePassV1::verify(bytes, &self.key, self.skew, SystemTime::now())
}
}
impl WirePassV1 {
pub fn canonical_unsigned_bytes(&self) -> Result<Vec<u8>, PassCheckError> {
if self.issued_at > self.expires_at || self.conversations.windows(2).any(|p| p[0] >= p[1]) {
return Err(PassCheckError::Malformed);
}
let mut out = vec![PASS_VERSION_V1];
put_bytes(&mut out, &self.participant)?;
out.extend_from_slice(&self.public_key);
put_len(&mut out, self.conversations.len())?;
for id in &self.conversations {
out.extend_from_slice(&id.to_be_bytes());
}
put_bytes(&mut out, self.live.as_bytes())?;
out.push(u8::from(self.may_enroll));
out.extend_from_slice(&self.issued_at.to_be_bytes());
out.extend_from_slice(&self.expires_at.to_be_bytes());
Ok(out)
}
pub fn canonical_bytes(&self) -> Result<Vec<u8>, PassCheckError> {
let mut out = self.canonical_unsigned_bytes()?;
out.extend_from_slice(&self.signature);
Ok(out)
}
pub fn parse(bytes: &[u8]) -> Result<Self, PassCheckError> {
let unsigned_len = bytes
.len()
.checked_sub(SIGNATURE_LEN)
.ok_or(PassCheckError::Malformed)?;
let (unsigned, sig) = bytes.split_at(unsigned_len);
let mut c = Cursor(unsigned);
if c.u8()? != PASS_VERSION_V1 {
return Err(PassCheckError::Malformed);
}
let participant = c.bytes()?.to_vec();
let public_key = c.array()?;
let count = c.len()?;
let mut conversations = Vec::with_capacity(count);
for _ in 0..count {
conversations.push(c.u64()?);
}
let live = std::str::from_utf8(c.bytes()?)
.map_err(|_| PassCheckError::Malformed)?
.to_owned();
let may_enroll = match c.u8()? {
0 => false,
1 => true,
_ => return Err(PassCheckError::Malformed),
};
let issued_at = c.u64()?;
let expires_at = c.u64()?;
if !c.0.is_empty()
|| issued_at > expires_at
|| conversations.windows(2).any(|p| p[0] >= p[1])
{
return Err(PassCheckError::Malformed);
}
Ok(Self {
participant,
public_key,
conversations,
live,
may_enroll,
issued_at,
expires_at,
signature: sig.try_into().map_err(|_| PassCheckError::Malformed)?,
})
}
pub fn verify(
bytes: &[u8],
key: &VerifyingKey,
skew: u64,
now: SystemTime,
) -> Result<PassPrincipal, PassCheckError> {
let pass = Self::parse(bytes)?;
key.verify_strict(
&pass.canonical_unsigned_bytes()?,
&Signature::from_bytes(&pass.signature),
)
.map_err(|_| PassCheckError::Signature)?;
let now = now
.duration_since(UNIX_EPOCH)
.map_err(|_| PassCheckError::Expired)?
.as_secs();
if now < pass.issued_at.saturating_sub(skew) || now > pass.expires_at.saturating_add(skew) {
return Err(PassCheckError::Expired);
}
Ok(PassPrincipal {
participant: pass.participant,
public_key: pass.public_key,
conversations: pass.conversations.into_iter().collect(),
live: pass.live,
may_enroll: pass.may_enroll,
})
}
}
fn put_len(out: &mut Vec<u8>, len: usize) -> Result<(), PassCheckError> {
out.extend_from_slice(
&u32::try_from(len)
.map_err(|_| PassCheckError::Malformed)?
.to_be_bytes(),
);
Ok(())
}
fn put_bytes(out: &mut Vec<u8>, bytes: &[u8]) -> Result<(), PassCheckError> {
put_len(out, bytes.len())?;
out.extend_from_slice(bytes);
Ok(())
}
struct Cursor<'a>(&'a [u8]);
impl<'a> Cursor<'a> {
const fn take(&mut self, n: usize) -> Result<&'a [u8], PassCheckError> {
if self.0.len() < n {
return Err(PassCheckError::Malformed);
}
let (a, b) = self.0.split_at(n);
self.0 = b;
Ok(a)
}
fn u8(&mut self) -> Result<u8, PassCheckError> {
Ok(self.take(1)?[0])
}
fn len(&mut self) -> Result<usize, PassCheckError> {
let a: [u8; 4] = self
.take(4)?
.try_into()
.map_err(|_| PassCheckError::Malformed)?;
usize::try_from(u32::from_be_bytes(a)).map_err(|_| PassCheckError::Malformed)
}
fn u64(&mut self) -> Result<u64, PassCheckError> {
Ok(u64::from_be_bytes(
self.take(8)?
.try_into()
.map_err(|_| PassCheckError::Malformed)?,
))
}
fn bytes(&mut self) -> Result<&'a [u8], PassCheckError> {
let n = self.len()?;
self.take(n)
}
fn array<const N: usize>(&mut self) -> Result<[u8; N], PassCheckError> {
self.take(N)?
.try_into()
.map_err(|_| PassCheckError::Malformed)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn shared_version_one_vector_is_byte_exact_and_verifies()
-> Result<(), Box<dyn std::error::Error>> {
let vector: serde_json::Value =
serde_json::from_str(include_str!("../test-vectors/wire-pass-v1.json"))?;
let unsigned = hex::decode(
vector["unsigned_pass_hex"]
.as_str()
.ok_or("unsigned bytes")?,
)?;
let signature: [u8; 64] =
hex::decode(vector["signature_hex"].as_str().ok_or("signature")?)?
.try_into()
.map_err(|_| "signature length")?;
let bytes = hex::decode(vector["pass_hex"].as_str().ok_or("pass bytes")?)?;
let registry_key: [u8; 32] = hex::decode(
vector["registry_verifying_key_hex"]
.as_str()
.ok_or("verifying key")?,
)?
.try_into()
.map_err(|_| "key length")?;
let pass = WirePassV1::parse(&bytes).map_err(|error| format!("parse failed: {error:?}"))?;
assert_eq!(
pass.canonical_unsigned_bytes()
.map_err(|error| format!("encode failed: {error:?}"))?,
unsigned
);
assert_eq!(pass.signature, signature);
let key = VerifyingKey::from_bytes(®istry_key)?;
let principal = WirePassV1::verify(
&bytes,
&key,
0,
UNIX_EPOCH + std::time::Duration::from_secs(pass.issued_at),
)
.map_err(|error| format!("verify failed: {error:?}"))?;
assert_eq!(principal.participant, b"participant-42");
assert_eq!(
principal.conversations.into_iter().collect::<Vec<_>>(),
vec![7, 42, 9001]
);
assert_eq!(principal.live, "workspace/acme/");
assert!(principal.may_enroll);
Ok(())
}
}