use std::collections::BTreeSet;
use std::fmt;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(transparent)]
pub struct SourceId(pub String);
impl SourceId {
pub fn new(s: impl Into<String>) -> Self {
Self(s.into())
}
}
impl fmt::Display for SourceId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(&self.0)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Trust {
Trusted,
Untrusted,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Sensitivity {
Public,
Internal,
Confidential,
Secret,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Label {
pub provenance: BTreeSet<SourceId>,
pub trust: Trust,
pub sensitivity: Sensitivity,
}
impl Label {
#[must_use]
pub fn trusted() -> Self {
Self {
provenance: BTreeSet::new(),
trust: Trust::Trusted,
sensitivity: Sensitivity::Public,
}
}
#[must_use]
pub fn untrusted(source: SourceId) -> Self {
Self {
provenance: BTreeSet::from([source]),
trust: Trust::Untrusted,
sensitivity: Sensitivity::Internal,
}
}
#[must_use]
pub fn with_sensitivity(mut self, s: Sensitivity) -> Self {
self.sensitivity = s;
self
}
#[must_use]
pub fn join(&self, other: &Self) -> Self {
Self {
provenance: self.provenance.union(&other.provenance).cloned().collect(),
trust: self.trust.max(other.trust),
sensitivity: self.sensitivity.max(other.sensitivity),
}
}
#[must_use]
pub fn is_untrusted(&self) -> bool {
self.trust == Trust::Untrusted
}
}
impl Default for Label {
fn default() -> Self {
Self::trusted()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Tainted<T> {
value: T,
label: Label,
}
impl<T> Tainted<T> {
pub fn trusted(value: T) -> Self {
Self {
value,
label: Label::trusted(),
}
}
pub fn from_source(value: T, source: SourceId) -> Self {
Self {
value,
label: Label::untrusted(source),
}
}
pub fn with_label(value: T, label: Label) -> Self {
Self { value, label }
}
pub fn label(&self) -> &Label {
&self.label
}
pub fn into_unlabelled(self) -> T {
self.value
}
pub fn peek(&self) -> &T {
&self.value
}
pub fn map<U>(self, f: impl FnOnce(T) -> U) -> Tainted<U> {
Tainted {
value: f(self.value),
label: self.label,
}
}
pub fn zip<U>(self, other: Tainted<U>) -> Tainted<(T, U)> {
let label = self.label.join(&other.label);
Tainted {
value: (self.value, other.value),
label,
}
}
pub(crate) fn into_parts(self) -> (T, Label) {
(self.value, self.label)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn src(s: &str) -> SourceId {
SourceId::new(s)
}
#[test]
fn trust_join_degrades() {
let t = Label::trusted();
let u = Label::untrusted(src("mcp://tool"));
assert_eq!(t.join(&u).trust, Trust::Untrusted);
assert_eq!(u.join(&t).trust, Trust::Untrusted, "join is commutative");
}
#[test]
fn sensitivity_join_escalates() {
let a = Label::trusted().with_sensitivity(Sensitivity::Public);
let b = Label::trusted().with_sensitivity(Sensitivity::Secret);
assert_eq!(a.join(&b).sensitivity, Sensitivity::Secret);
}
#[test]
fn provenance_accumulates() {
let a = Label::untrusted(src("a"));
let b = Label::untrusted(src("b"));
let j = a.join(&b);
assert_eq!(j.provenance.len(), 2);
}
#[test]
fn join_is_idempotent_and_associative() {
let a = Label::untrusted(src("a"));
let b = Label::trusted().with_sensitivity(Sensitivity::Confidential);
let c = Label::untrusted(src("c")).with_sensitivity(Sensitivity::Secret);
assert_eq!(a.join(&a), a, "idempotent");
assert_eq!(a.join(&b).join(&c), a.join(&b.join(&c)), "associative");
}
#[test]
fn zip_propagates_untrust_to_derived_values() {
let trusted = Tainted::trusted(1);
let untrusted = Tainted::from_source(2, src("mcp://tool"));
let combined = trusted.zip(untrusted).map(|(a, b)| a + b);
assert!(combined.label().is_untrusted());
assert_eq!(*combined.peek(), 3);
}
#[test]
fn map_preserves_label() {
let t = Tainted::from_source("x", src("doc"));
let mapped = t.map(str::to_uppercase);
assert!(mapped.label().is_untrusted());
assert_eq!(mapped.peek(), "X");
}
}