use std::collections::HashSet;
use std::fmt;
use std::path::Path;
use crate::crypto::{CryptoPolicyDefault, hexdigest};
use crate::error::{Error, Result};
use crate::etree::{ParseOps, TextNode};
#[derive(Clone, Debug, Eq, PartialEq, Hash, PartialOrd, Ord)]
pub struct WordId(String);
impl WordId {
pub fn new(s: impl Into<String>) -> Self {
WordId(s.into())
}
pub fn as_str(&self) -> &str {
&self.0
}
}
impl fmt::Display for WordId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, Hash, PartialOrd, Ord)]
pub struct KeyFp([u8; 32]);
impl KeyFp {
pub fn from_pem(pem: &str) -> Result<Self> {
let policy = CryptoPolicyDefault {};
let hash = hexdigest("sha3-256", pem.as_bytes(), &policy)?;
let mut arr = [0u8; 32];
let bytes = hex::decode(&hash)?;
if bytes.len() != 32 {
return Err(Error::msg(format!(
"expected 32-byte SHA3-256, got {} bytes",
bytes.len()
)));
}
arr.copy_from_slice(&bytes);
Ok(KeyFp(arr))
}
pub fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
pub fn from_bytes(arr: [u8; 32]) -> Self {
KeyFp(arr)
}
pub fn to_hex(&self) -> String {
hex::encode(self.0)
}
}
impl fmt::Display for KeyFp {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.to_hex())
}
}
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub enum Capability {
Viewer,
Reader,
Decryptor(WordId),
Signer(KeyFp),
Verifier(KeyFp),
}
impl Capability {
pub fn tier(&self) -> &'static str {
match self {
Capability::Viewer => "viewer",
Capability::Reader => "reader",
Capability::Decryptor(_) => "decryptor",
Capability::Signer(_) => "signer",
Capability::Verifier(_) => "verifier",
}
}
}
impl fmt::Display for Capability {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Capability::Viewer => f.write_str("viewer"),
Capability::Reader => f.write_str("reader"),
Capability::Decryptor(w) => write!(f, "decryptor({})", w),
Capability::Signer(k) => write!(f, "signer({})", k),
Capability::Verifier(k) => write!(f, "verifier({})", k),
}
}
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct CapabilitySet(HashSet<Capability>);
impl CapabilitySet {
pub fn empty() -> Self {
CapabilitySet(HashSet::new())
}
pub fn viewing() -> Self {
let mut s = Self::empty();
s.0.insert(Capability::Viewer);
s
}
pub fn from_paops(paops: &ParseOps) -> Self {
let mut s = Self::viewing();
if Path::new(&paops.io.casdir).is_dir() {
s.0.insert(Capability::Reader);
}
for word in paops.passwords.keys() {
s.0.insert(Capability::Decryptor(WordId::new(word)));
}
s
}
pub fn with_signer(mut self, fp: KeyFp) -> Self {
self.0.insert(Capability::Signer(fp));
self
}
pub fn with_verifier(mut self, fp: KeyFp) -> Self {
self.0.insert(Capability::Verifier(fp));
self
}
pub fn insert(&mut self, c: Capability) {
self.0.insert(c);
}
pub fn union(mut self, other: &Self) -> Self {
for c in &other.0 {
self.0.insert(c.clone());
}
self
}
pub fn contains(&self, c: &Capability) -> bool {
self.0.contains(c)
}
pub fn satisfies(&self, required: &Self) -> bool {
for c in &required.0 {
if !self.0.contains(c) {
return false;
}
}
true
}
pub fn iter_sorted(&self) -> Vec<&Capability> {
let mut v: Vec<&Capability> = self.0.iter().collect();
v.sort_by_key(|c| c.to_string());
v
}
pub fn len(&self) -> usize {
self.0.len()
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
}
impl fmt::Display for CapabilitySet {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let sorted = self.iter_sorted();
let parts: Vec<String> = sorted.iter().map(|c| c.to_string()).collect();
f.write_str(&parts.join(", "))
}
}
pub fn required_for(node: &TextNode) -> CapabilitySet {
let mut s = CapabilitySet::viewing();
match node {
TextNode::Plain(_) => {}
TextNode::Data(_) => {}
TextNode::Stored { .. } => {
s.insert(Capability::Reader);
}
TextNode::BeginEnd { .. } => {}
TextNode::Encrypted { keyw, .. } => {
s.insert(Capability::Decryptor(WordId::new(keyw)));
}
TextNode::Chain { .. } => {}
TextNode::Include { .. } => {
s.insert(Capability::Reader);
}
TextNode::Conflict { ours, theirs, .. } => {
for child in ours.iter().chain(theirs.iter()) {
let child_req = required_for(child);
for c in child_req.iter_sorted() {
s.insert(c.clone());
}
}
}
}
s
}
#[cfg(test)]
mod tests {
use super::*;
use crate::etree::ParseOps;
fn empty_paops() -> ParseOps {
ParseOps::new(Box::new(crate::crypto::CryptoPolicyDefault {})).unwrap()
}
#[test]
fn viewer_is_always_present() {
let paops = empty_paops();
let caps = CapabilitySet::from_paops(&paops);
assert!(caps.contains(&Capability::Viewer));
}
#[test]
fn reader_added_when_casdir_exists() {
let tmp = tempfile::tempdir().unwrap();
let mut paops = empty_paops();
paops.io.casdir = tmp.path().to_path_buf();
let caps = CapabilitySet::from_paops(&paops);
assert!(caps.contains(&Capability::Reader));
}
#[test]
fn reader_not_added_when_casdir_missing() {
let mut paops = empty_paops();
paops.io.casdir = "/nonexistent/path/that/does/not/exist".into();
let caps = CapabilitySet::from_paops(&paops);
assert!(!caps.contains(&Capability::Reader));
}
#[test]
fn decryptor_added_per_password() {
let mut paops = empty_paops();
paops.passwords.insert("Agent_007".into(), "secret".into());
paops.passwords.insert("GEHEIM".into(), "secret".into());
let caps = CapabilitySet::from_paops(&paops);
assert!(caps.contains(&Capability::Decryptor(WordId::new("Agent_007"))));
assert!(caps.contains(&Capability::Decryptor(WordId::new("GEHEIM"))));
assert!(!caps.contains(&Capability::Decryptor(WordId::new("OTHER"))));
}
#[test]
fn decryptor_for_one_word_does_not_satisfy_another() {
let mut paops = empty_paops();
paops.passwords.insert("Agent_007".into(), "secret".into());
let held = CapabilitySet::from_paops(&paops);
let mut required = CapabilitySet::viewing();
required.insert(Capability::Decryptor(WordId::new("Agent_007")));
assert!(held.satisfies(&required));
let mut required_other = CapabilitySet::viewing();
required_other.insert(Capability::Decryptor(WordId::new("GEHEIM")));
assert!(!held.satisfies(&required_other));
}
#[test]
fn union_grows_capability_set() {
let mut a = CapabilitySet::viewing();
a.insert(Capability::Decryptor(WordId::new("Agent_007")));
let mut b = CapabilitySet::viewing();
b.insert(Capability::Decryptor(WordId::new("GEHEIM")));
let u = a.union(&b);
assert!(u.contains(&Capability::Decryptor(WordId::new("Agent_007"))));
assert!(u.contains(&Capability::Decryptor(WordId::new("GEHEIM"))));
}
#[test]
fn keyfp_from_pem_is_stable() {
let pem = "-----BEGIN PUBLIC KEY-----\nMCowBQYDK2VwAyEA9k0KvS0XUdLK8+nJSlQYlkqY8CyJa0nTK0YJGSk=\n-----END PUBLIC KEY-----\n";
let fp1 = KeyFp::from_pem(pem).unwrap();
let fp2 = KeyFp::from_pem(pem).unwrap();
assert_eq!(fp1.to_hex(), fp2.to_hex());
assert_eq!(fp1.to_hex().len(), 64); }
#[test]
fn required_for_node_mappings_are_mece() {
let plain_caps = required_for(&TextNode::Plain("hello".into()));
assert_eq!(plain_caps.len(), 1);
assert!(plain_caps.contains(&Capability::Viewer));
let be_caps = required_for(&TextNode::BeginEnd {
keyw: "X".into(),
txt: Vec::new(),
});
assert_eq!(be_caps.len(), 1);
let stored_caps = required_for(&TextNode::Stored {
keyw: "ct".into(),
cas: "abc".into(),
});
assert!(stored_caps.contains(&Capability::Viewer));
assert!(stored_caps.contains(&Capability::Reader));
let enc_caps = required_for(&TextNode::Encrypted {
keyw: "Agent_007".into(),
txt: Vec::new(),
extfields: std::collections::BTreeMap::new(),
});
assert!(enc_caps.contains(&Capability::Viewer));
assert!(enc_caps.contains(&Capability::Decryptor(WordId::new("Agent_007"))));
}
#[test]
fn display_is_stable_and_sorted() {
let mut s = CapabilitySet::viewing();
s.insert(Capability::Decryptor(WordId::new("Z")));
s.insert(Capability::Decryptor(WordId::new("A")));
let display = s.to_string();
let a_idx = display.find("A").unwrap();
let z_idx = display.find("Z").unwrap();
assert!(a_idx < z_idx);
}
}