use std::io;
use std::path::{Path, PathBuf};
use crate::agent::{AgentProfile, Entitlements};
pub const PINS_DIR: &str = "entitlement-pins";
const PIN_EXT: &str = "json";
const PIN_VERSION: u32 = 1;
#[derive(serde::Serialize, serde::Deserialize)]
struct PinFile {
version: u32,
entitlements: serde_json::Value,
}
pub fn pin_path(mur_home: &Path, agent: &str) -> PathBuf {
mur_home.join(PINS_DIR).join(format!("{agent}.{PIN_EXT}"))
}
#[derive(Debug, PartialEq, Eq)]
pub enum PinCheck {
Match,
Missing,
Mismatch { changed: Vec<String> },
}
#[derive(Debug, thiserror::Error)]
pub enum PinError {
#[error("io error on {path}: {source}")]
Io { path: PathBuf, source: io::Error },
#[error("invalid pin file {path}: {msg}")]
Invalid { path: PathBuf, msg: String },
#[error("invalid profile {path}: {msg}")]
Profile { path: PathBuf, msg: String },
#[error("invalid agent name '{0}'")]
Name(String),
}
fn canonical(ent: &Entitlements) -> serde_json::Value {
serde_json::to_value(ent).expect("Entitlements serializes to JSON")
}
fn checked_name(agent: &str) -> Result<(), PinError> {
crate::agent_name::validate_agent_name(agent).map_err(|_| PinError::Name(agent.to_string()))
}
pub fn write_pin(mur_home: &Path, agent: &str, ent: &Entitlements) -> Result<(), PinError> {
checked_name(agent)?;
let path = pin_path(mur_home, agent);
let io_err = |source| PinError::Io {
path: path.clone(),
source,
};
let dir = mur_home.join(PINS_DIR);
std::fs::create_dir_all(&dir).map_err(io_err)?;
let body = PinFile {
version: PIN_VERSION,
entitlements: canonical(ent),
};
let bytes = crate::jcs::to_jcs_for(&body);
let mut tmp = tempfile::NamedTempFile::new_in(&dir).map_err(io_err)?;
io::Write::write_all(&mut tmp, &bytes).map_err(io_err)?;
tmp.persist(&path).map_err(|e| io_err(e.error))?;
Ok(())
}
pub fn entitlements_from_yaml(yaml: &str, path: &Path) -> Result<Entitlements, PinError> {
serde_yaml_ng::from_str::<AgentProfile>(yaml)
.map(|p| p.entitlements)
.map_err(|e| PinError::Profile {
path: path.to_path_buf(),
msg: e.to_string(),
})
}
pub fn repin_from_profile(mur_home: &Path, agent: &str) -> Result<(), PinError> {
checked_name(agent)?;
let path = mur_home.join("agents").join(agent).join("profile.yaml");
let yaml = std::fs::read_to_string(&path).map_err(|source| PinError::Io {
path: path.clone(),
source,
})?;
let ent = entitlements_from_yaml(&yaml, &path)?;
write_pin(mur_home, agent, &ent)
}
pub fn advance_pin(
mur_home: &Path,
agent: &str,
prior: Option<&Entitlements>,
new: &Entitlements,
) -> Result<bool, PinError> {
let trusted = match prior {
None => true,
Some(p) => matches!(
check(mur_home, agent, p)?,
PinCheck::Match | PinCheck::Missing
),
};
if trusted {
write_pin(mur_home, agent, new)?;
}
Ok(trusted)
}
pub fn remove_pin(mur_home: &Path, agent: &str) -> Result<(), PinError> {
checked_name(agent)?;
let path = pin_path(mur_home, agent);
match std::fs::remove_file(&path) {
Ok(()) => Ok(()),
Err(e) if e.kind() == io::ErrorKind::NotFound => Ok(()),
Err(source) => Err(PinError::Io { path, source }),
}
}
pub fn rename_pin(mur_home: &Path, old: &str, new: &str) -> Result<(), PinError> {
checked_name(old)?;
checked_name(new)?;
let (from, to) = (pin_path(mur_home, old), pin_path(mur_home, new));
match std::fs::rename(&from, &to) {
Ok(()) => Ok(()),
Err(e) if e.kind() == io::ErrorKind::NotFound => Ok(()),
Err(source) => Err(PinError::Io { path: from, source }),
}
}
pub fn check(mur_home: &Path, agent: &str, ent: &Entitlements) -> Result<PinCheck, PinError> {
checked_name(agent)?;
let path = pin_path(mur_home, agent);
let raw = match std::fs::read(&path) {
Ok(b) => b,
Err(e) if e.kind() == io::ErrorKind::NotFound => return Ok(PinCheck::Missing),
Err(source) => return Err(PinError::Io { path, source }),
};
let invalid = |msg: String| PinError::Invalid {
path: path.clone(),
msg,
};
let file: PinFile = serde_json::from_slice(&raw).map_err(|e| invalid(e.to_string()))?;
if file.version != PIN_VERSION {
return Err(invalid(format!("unsupported version {}", file.version)));
}
let pinned: Entitlements =
serde_json::from_value(file.entitlements).map_err(|e| invalid(e.to_string()))?;
let (a, b) = (canonical(&pinned), canonical(ent));
if a == b {
return Ok(PinCheck::Match);
}
Ok(PinCheck::Mismatch {
changed: changed_keys(&a, &b),
})
}
fn changed_keys(a: &serde_json::Value, b: &serde_json::Value) -> Vec<String> {
let (Some(a), Some(b)) = (a.as_object(), b.as_object()) else {
return vec!["entitlements".to_string()];
};
let mut keys: Vec<String> = a
.keys()
.chain(b.keys())
.filter(|k| a.get(*k) != b.get(*k))
.cloned()
.collect();
keys.sort();
keys.dedup();
keys
}
#[cfg(test)]
mod tests;