use std::fs;
use std::path::{Component, Path, PathBuf};
use anyhow::{bail, Context, Result};
use serde::Deserialize;
pub const POLICY_FILE: &str = "agent.lock";
pub const CURRENT_VERSION: u32 = 1;
pub const TEMPLATE: &str = "\
version: 1
protect:
- file
- folder
";
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct Raw {
version: u32,
#[serde(default)]
protect: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct Policy {
pub root: PathBuf,
pub file: PathBuf,
pub version: u32,
pub patterns: Vec<String>,
}
impl Policy {
pub fn find_root(start: &Path) -> Option<PathBuf> {
let start = absolute(start).ok()?;
start
.ancestors()
.find(|dir| dir.join(POLICY_FILE).is_file())
.map(Path::to_path_buf)
}
pub fn load(start: &Path) -> Result<Policy> {
let root = Policy::find_root(start).with_context(|| {
format!(
"no {POLICY_FILE} found in {} or any parent directory (run `ralon init`)",
start.display()
)
})?;
let file = root.join(POLICY_FILE);
let text = fs::read_to_string(&file)
.with_context(|| format!("failed to read {}", file.display()))?;
Policy::parse(root, file, &text)
}
pub fn parse(root: PathBuf, file: PathBuf, text: &str) -> Result<Policy> {
let raw: Raw = serde_yaml_ng::from_str(text)
.with_context(|| format!("failed to parse {}", file.display()))?;
if raw.version != CURRENT_VERSION {
bail!(
"{}: unsupported version {} (this build understands version {})",
file.display(),
raw.version,
CURRENT_VERSION
);
}
let mut patterns = vec![POLICY_FILE.to_string()];
for raw_pattern in &raw.protect {
let pattern = normalize_pattern(raw_pattern)
.with_context(|| format!("{}: invalid pattern", file.display()))?;
if !patterns.contains(&pattern) {
patterns.push(pattern);
}
}
Ok(Policy {
root,
file,
version: raw.version,
patterns,
})
}
pub fn declared_patterns(&self) -> &[String] {
&self.patterns[1..]
}
}
fn normalize_pattern(raw: &str) -> Result<String> {
let trimmed = raw.trim();
if trimmed.is_empty() {
bail!("empty pattern");
}
if let Some(rest) = trimmed.strip_prefix('!') {
bail!("negation is not supported in version 1: `!{rest}`");
}
if trimmed.starts_with('~') {
bail!("`~` is not expanded: `{trimmed}`");
}
let unified = trimmed.replace('\\', "/");
let relative = unified
.strip_prefix("./")
.unwrap_or(&unified)
.trim_start_matches('/')
.trim_end_matches('/');
if relative.is_empty() {
bail!("`{trimmed}` does not name anything inside the project");
}
if relative.split('/').any(|part| part == "..") {
bail!("`..` may not be used to escape the project root: `{trimmed}`");
}
if relative.contains(':') {
bail!("absolute paths are not allowed, patterns are relative to agent.lock: `{trimmed}`");
}
Ok(relative.to_string())
}
pub fn absolute(path: &Path) -> Result<PathBuf> {
let absolute = std::path::absolute(path)
.with_context(|| format!("failed to resolve {}", path.display()))?;
let mut normalized = PathBuf::new();
for component in absolute.components() {
match component {
Component::CurDir => {}
Component::ParentDir => {
normalized.pop();
}
other => normalized.push(other.as_os_str()),
}
}
Ok(normalized)
}
#[cfg(test)]
mod tests {
use super::*;
fn parse(text: &str) -> Result<Policy> {
Policy::parse(PathBuf::from("/p"), PathBuf::from("/p/agent.lock"), text)
}
#[test]
fn parses_minimal_policy() {
let policy = parse("version: 1\nprotect:\n - src/index.tsx\n - config/**\n").unwrap();
assert_eq!(policy.version, 1);
assert_eq!(
policy.patterns,
["agent.lock", "src/index.tsx", "config/**"]
);
assert_eq!(policy.declared_patterns(), ["src/index.tsx", "config/**"]);
}
#[test]
fn protect_defaults_to_empty_but_policy_file_is_always_protected() {
let policy = parse("version: 1\n").unwrap();
assert_eq!(policy.patterns, ["agent.lock"]);
}
#[test]
fn duplicate_of_the_implicit_pattern_is_not_repeated() {
let policy = parse("version: 1\nprotect:\n - ./agent.lock\n").unwrap();
assert_eq!(policy.patterns, ["agent.lock"]);
}
#[test]
fn normalizes_separators_and_affixes() {
let policy = parse("version: 1\nprotect:\n - /src\\auth.ts\n - config/\n").unwrap();
assert_eq!(policy.patterns, ["agent.lock", "src/auth.ts", "config"]);
}
#[test]
fn rejects_unknown_versions() {
let err = parse("version: 2\nprotect: []\n").unwrap_err().to_string();
assert!(err.contains("unsupported version 2"), "{err}");
}
#[test]
fn rejects_unknown_keys() {
assert!(parse("version: 1\nallow:\n - src\n").is_err());
}
#[test]
fn rejects_escaping_patterns() {
for pattern in [
"../outside",
"src/../../etc",
"~/.ssh",
"!src/a.ts",
"C:/win",
] {
let text = format!("version: 1\nprotect:\n - \"{pattern}\"\n");
assert!(parse(&text).is_err(), "accepted {pattern}");
}
}
}