use std::io::Write;
use std::path::PathBuf;
use tirith_core::context_detect::{self, ContextDetectFailure, Provider, ProviderContext};
use tirith_core::policy::{self as policy_mod, Policy};
const ALLOWED_CRITICALITIES: &[&str] = &[
"critical",
"production",
"prod",
"live",
"p0",
"p1",
"p2",
"staging",
"dev",
"test",
];
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LabelScope {
User,
Repo,
}
impl LabelScope {
fn as_str(self) -> &'static str {
match self {
Self::User => "user",
Self::Repo => "repo",
}
}
pub fn parse(s: &str) -> Option<Self> {
match s.trim().to_lowercase().as_str() {
"user" => Some(Self::User),
"repo" | "project" | "workspace" => Some(Self::Repo),
_ => None,
}
}
}
pub fn status(json: bool) -> i32 {
let mut policy = Policy::discover_partial(None);
policy.load_context_labels(None);
let detection = context_detect::detect_all();
if json {
return emit_status_json(&detection, &policy);
}
if detection.contexts.is_empty() && detection.failures.is_empty() {
eprintln!("tirith context status: no cloud / k8s context detected");
eprintln!(" (configure ~/.kube/config, AWS_PROFILE, gcloud or az to populate)");
return 0;
}
eprintln!("tirith context status:");
for provider in [
Provider::Kube,
Provider::Aws,
Provider::Gcp,
Provider::Azure,
] {
match (
detection.contexts.get(&provider),
detection.failures.get(&provider),
) {
(Some(ctx), _) => {
let label = policy
.context_labels
.get(&ctx.label_key())
.map(String::as_str)
.unwrap_or("(unlabeled)");
let context = super::sanitize_for_human_output(&ctx.context, false);
let label = super::sanitize_for_human_output(label, false);
eprintln!(" {:<6} {} [label: {label}]", provider.as_str(), context,);
}
(None, Some(failure)) => {
let failure = super::sanitize_for_human_output(&failure.to_string(), false);
eprintln!(" {:<6} <error: {failure}>", provider.as_str());
}
(None, None) => {
eprintln!(" {:<6} (not configured)", provider.as_str());
}
}
}
eprintln!(
" guard: {} label-file (user): {}",
if policy.context_guard_enabled {
"ON"
} else {
"OFF"
},
policy_mod::user_context_labels_path()
.map(|p| p.display().to_string())
.unwrap_or_else(|| "<unknown>".into()),
);
0
}
fn emit_status_json(detection: &context_detect::DetectionResult, policy: &Policy) -> i32 {
#[derive(serde::Serialize)]
struct ProviderEntry {
provider: &'static str,
context: Option<String>,
label: Option<String>,
error: Option<String>,
}
#[derive(serde::Serialize)]
struct Out {
schema_version: u32,
guard_enabled: bool,
user_label_file: Option<String>,
repo_label_file: Option<String>,
providers: Vec<ProviderEntry>,
}
let mut providers = Vec::new();
for provider in [
Provider::Kube,
Provider::Aws,
Provider::Gcp,
Provider::Azure,
] {
let (context, label, error) = match (
detection.contexts.get(&provider),
detection.failures.get(&provider),
) {
(Some(ctx), _) => (
Some(ctx.context.clone()),
policy.context_labels.get(&ctx.label_key()).cloned(),
None,
),
(None, Some(f)) => (None, None, Some(f.to_string())),
(None, None) => (None, None, None),
};
providers.push(ProviderEntry {
provider: provider.as_str(),
context,
label,
error,
});
}
let out = Out {
schema_version: 1,
guard_enabled: policy.context_guard_enabled,
user_label_file: policy_mod::user_context_labels_path().map(|p| p.display().to_string()),
repo_label_file: policy_mod::repo_context_labels_path(None)
.map(|p| p.display().to_string()),
providers,
};
let mut stdout = std::io::stdout().lock();
if serde_json::to_writer_pretty(&mut stdout, &out).is_err() || writeln!(stdout).is_err() {
eprintln!("tirith context status: failed to write JSON output");
return 1;
}
0
}
pub fn guard(action: &str, json: bool) -> i32 {
let enable = match action {
"on" | "enable" | "true" => true,
"off" | "disable" | "false" => false,
"status" => return guard_status(json),
other => {
eprintln!("tirith context guard: unknown action '{other}' (expected on|off|status)");
return 2;
}
};
let target_path = match resolve_policy_path_for_guard() {
Ok(p) => p,
Err(code) => return code,
};
if let Err(e) = update_policy_guard_key(&target_path, enable) {
eprintln!(
"tirith context guard: failed to update {}: {e}",
target_path.display()
);
return 1;
}
if json {
let out = serde_json::json!({
"schema_version": 1,
"guard_enabled": enable,
"policy_path": target_path.display().to_string(),
});
let mut stdout = std::io::stdout().lock();
if serde_json::to_writer_pretty(&mut stdout, &out).is_err() || writeln!(stdout).is_err() {
return 1;
}
} else {
eprintln!(
"tirith context guard: {} (written to {})",
if enable { "ON" } else { "OFF" },
target_path.display(),
);
}
0
}
fn guard_status(json: bool) -> i32 {
let policy = Policy::discover_partial(None);
if json {
let out = serde_json::json!({
"schema_version": 1,
"guard_enabled": policy.context_guard_enabled,
"policy_path": policy.path,
});
let mut stdout = std::io::stdout().lock();
if serde_json::to_writer_pretty(&mut stdout, &out).is_err() || writeln!(stdout).is_err() {
return 1;
}
} else {
eprintln!(
"tirith context guard: {}",
if policy.context_guard_enabled {
"ON"
} else {
"OFF"
}
);
}
0
}
fn resolve_policy_path_for_guard() -> Result<PathBuf, i32> {
if let Some(existing) = policy_mod::discover_local_policy_path(None) {
return Ok(existing);
}
let user = policy_mod::config_dir().ok_or_else(|| {
eprintln!("tirith context guard: could not resolve user config dir");
1
})?;
Ok(user.join("policy.yaml"))
}
pub(super) fn update_policy_guard_key(path: &std::path::Path, enable: bool) -> std::io::Result<()> {
let root = path
.parent()
.filter(|parent| !parent.as_os_str().is_empty())
.ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
"policy path has no containing directory",
)
})?;
let policy = Policy::discover_local_only(root.to_str());
let contained =
super::prepare_config_destination_permitted(root, path, true, &policy, true, true)?;
let existing = read_existing_policy_for_guard(&contained, path)?;
let new_line = format!("context_guard_enabled: {enable}");
let mut out = String::new();
let mut replaced = false;
for line in existing.lines() {
if line.starts_with("context_guard_enabled:") {
out.push_str(&new_line);
out.push('\n');
replaced = true;
} else {
out.push_str(line);
out.push('\n');
}
}
if !replaced {
if !out.is_empty() && !out.ends_with('\n') {
out.push('\n');
}
out.push_str(&new_line);
out.push('\n');
}
verify_context_guard_effective(&out, enable)?;
super::write_prepared_config_file_permitted(
root,
path,
contained,
out.as_bytes(),
true,
&policy,
true,
)
}
fn verify_context_guard_effective(candidate: &str, expected: bool) -> std::io::Result<()> {
let parsed: serde_yaml::Value = serde_yaml::from_str(candidate).map_err(|e| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("resulting policy would not parse as YAML: {e}"),
)
})?;
let effective = parsed
.get("context_guard_enabled")
.and_then(serde_yaml::Value::as_bool);
if effective == Some(expected) {
Ok(())
} else {
Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"resulting policy does not have the requested top-level context_guard_enabled value",
))
}
}
fn read_existing_policy_for_guard(
contained: &tirith_core::util::ContainedAtomicFile,
path: &std::path::Path,
) -> std::io::Result<String> {
use tirith_core::util::OpenRegularError;
const GUARD_POLICY_READ_CAP: u64 = 1024 * 1024;
match contained.read_capped(GUARD_POLICY_READ_CAP) {
Ok(bytes) => String::from_utf8(bytes).map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"refusing to update {}: existing policy is not valid UTF-8",
path.display()
),
)
}),
Err(OpenRegularError::NotFound) => Ok(String::new()),
Err(OpenRegularError::NotRegularFile) => Err(std::io::Error::new(
std::io::ErrorKind::PermissionDenied,
format!(
"refusing to update {}: not a regular file (symlink?)",
path.display()
),
)),
Err(OpenRegularError::TooLarge) => Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("refusing to update {}: exceeds 1 MiB", path.display()),
)),
Err(OpenRegularError::Io(e)) => Err(e),
}
}
pub fn label(label_key: &str, criticality: &str, scope: LabelScope, json: bool) -> i32 {
if !label_key.contains(':') {
eprintln!(
"tirith context label: '{label_key}' is not a valid 'provider:context' key (e.g. kube:prod-us-east)"
);
return 2;
}
let (provider_str, ctx_part) = match label_key.split_once(':') {
Some(parts) => parts,
None => unreachable!("contains ':' checked above"),
};
if Provider::parse(provider_str).is_none() {
eprintln!(
"tirith context label: unknown provider '{provider_str}' (expected one of: kube, aws, gcp, azure)"
);
return 2;
}
if ctx_part.is_empty() {
eprintln!("tirith context label: context name is empty after the colon");
return 2;
}
let criticality_norm = criticality.trim().to_lowercase();
if !ALLOWED_CRITICALITIES.iter().any(|c| *c == criticality_norm) {
eprintln!(
"tirith context label: '{criticality}' is not a known criticality (expected one of: {}; case-insensitive)",
ALLOWED_CRITICALITIES.join(", "),
);
return 2;
}
let target_path = match scope {
LabelScope::User => match policy_mod::user_context_labels_path() {
Some(p) => p,
None => {
eprintln!("tirith context label: could not resolve user config dir");
return 1;
}
},
LabelScope::Repo => match policy_mod::repo_context_labels_path(None) {
Some(p) => p,
None => {
eprintln!("tirith context label: --scope repo requires running inside a git repo");
return 1;
}
},
};
let policy = Policy::discover_local_only(
target_path
.parent()
.and_then(std::path::Path::parent)
.and_then(std::path::Path::to_str),
);
if let Err(e) =
super::write_context_labels_permitted(&target_path, &[(label_key, criticality)], &policy)
{
eprintln!(
"tirith context label: failed to write {}: {e}",
target_path.display()
);
return 1;
}
if json {
let out = serde_json::json!({
"schema_version": 1,
"scope": scope.as_str(),
"path": target_path.display().to_string(),
"label_key": label_key,
"criticality": criticality,
});
let mut stdout = std::io::stdout().lock();
if serde_json::to_writer_pretty(&mut stdout, &out).is_err() || writeln!(stdout).is_err() {
return 1;
}
} else {
eprintln!(
"tirith context label: {label_key} -> {criticality} (scope={}, file={})",
scope.as_str(),
target_path.display(),
);
}
0
}
// Silence unused-import warnings under cfg combinations.
#[allow(dead_code)]
fn _silence_unused(_pc: &ProviderContext, _f: &ContextDetectFailure) {}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[test]
fn label_scope_parse() {
assert_eq!(LabelScope::parse("user"), Some(LabelScope::User));
assert_eq!(LabelScope::parse("USER"), Some(LabelScope::User));
assert_eq!(LabelScope::parse("repo"), Some(LabelScope::Repo));
assert_eq!(LabelScope::parse("workspace"), Some(LabelScope::Repo));
assert_eq!(LabelScope::parse("invalid"), None);
}
#[test]
fn update_policy_guard_key_creates_file() {
let dir = tempdir().unwrap();
let path = dir.path().join("policy.yaml");
update_policy_guard_key(&path, true).unwrap();
let content = std::fs::read_to_string(&path).unwrap();
assert!(content.contains("context_guard_enabled: true"));
}
#[test]
fn update_policy_guard_key_replaces_existing() {
let dir = tempdir().unwrap();
let path = dir.path().join("policy.yaml");
std::fs::write(
&path,
"paranoia: 2\ncontext_guard_enabled: true\nfail_mode: open\n",
)
.unwrap();
update_policy_guard_key(&path, false).unwrap();
let content = std::fs::read_to_string(&path).unwrap();
assert!(content.contains("context_guard_enabled: false"));
assert!(content.contains("paranoia: 2"));
assert!(content.contains("fail_mode: open"));
assert!(!content.contains("context_guard_enabled: true"));
}
#[test]
fn update_policy_guard_key_appends_when_missing() {
let dir = tempdir().unwrap();
let path = dir.path().join("policy.yaml");
std::fs::write(&path, "paranoia: 2\n").unwrap();
update_policy_guard_key(&path, true).unwrap();
let content = std::fs::read_to_string(&path).unwrap();
assert!(content.contains("paranoia: 2"));
assert!(content.contains("context_guard_enabled: true"));
}
#[test]
fn update_policy_guard_key_ignores_indented_lookalike() {
// An indented nested-mapping lookalike must not be rewritten at column
// zero; the real root key is appended and the document stays valid.
let dir = tempdir().unwrap();
let path = dir.path().join("policy.yaml");
std::fs::write(
&path,
"custom_rule:\n context_guard_enabled: false\n other: 1\n",
)
.unwrap();
update_policy_guard_key(&path, true).unwrap();
let content = std::fs::read_to_string(&path).unwrap();
assert!(content.contains(" context_guard_enabled: false"));
let top_level = content
.lines()
.filter(|l| l.starts_with("context_guard_enabled:"))
.count();
assert_eq!(top_level, 1);
let parsed: serde_yaml::Value = serde_yaml::from_str(&content).unwrap();
assert_eq!(
parsed
.get("context_guard_enabled")
.and_then(|v| v.as_bool()),
Some(true)
);
}
#[cfg(unix)]
#[test]
fn update_policy_guard_key_refuses_symlink_target() {
// Regression: repo-0371 — a repository-controlled policy.yaml symlink
// must not turn the guard update into an arbitrary-file rewrite.
let dir = tempdir().unwrap();
let outside = dir.path().join("outside.txt");
std::fs::write(&outside, "do not touch\n").unwrap();
let link = dir.path().join("policy.yaml");
std::os::unix::fs::symlink(&outside, &link).unwrap();
let err = update_policy_guard_key(&link, true).unwrap_err();
assert_eq!(err.kind(), std::io::ErrorKind::InvalidInput);
assert_eq!(std::fs::read_to_string(&outside).unwrap(), "do not touch\n");
}
#[test]
fn update_policy_guard_key_refuses_non_utf8_existing() {
// Regression: repo-0371 — non-UTF-8 content must not be treated as
// empty and clobbered with only the guard key.
let dir = tempdir().unwrap();
let path = dir.path().join("policy.yaml");
std::fs::write(&path, [0xff, 0xfe, 0x00, 0x01]).unwrap();
let err = update_policy_guard_key(&path, true).unwrap_err();
assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
assert_eq!(
std::fs::read(&path).unwrap(),
vec![0xff, 0xfe, 0x00, 0x01],
"target must be left untouched"
);
}
#[test]
fn status_sanitizes_untrusted_fields() {
// Regression: repo-0372 — context names / labels must not carry
// terminal control sequences into human status output.
let s = super::super::sanitize_for_human_output(
"aws:default\u{1b}]52;c;SGFja2Vk\u{7}\u{202e}",
false,
);
assert!(!s.contains('\u{1b}'));
assert!(!s.contains('\u{202e}'));
}
}