use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
use std::collections::{BTreeMap, HashMap};
use std::fmt;
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema, PartialEq)]
pub enum DataClassification {
#[serde(rename = "public")]
Public,
#[serde(rename = "internal")]
Internal,
#[serde(rename = "confidential")]
Confidential,
#[serde(rename = "regulated")]
Regulated,
}
impl fmt::Display for DataClassification {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
DataClassification::Public => write!(f, "public"),
DataClassification::Internal => write!(f, "internal"),
DataClassification::Confidential => write!(f, "confidential"),
DataClassification::Regulated => write!(f, "regulated"),
}
}
}
fn default_classification() -> DataClassification {
DataClassification::Internal
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct PolicySet {
pub version: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub metadata: Option<PolicyMetadata>,
pub defaults: PolicyDefaults,
pub rules: Vec<PolicyRule>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub classifications: BTreeMap<String, DataClassification>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub credential_scopes: Vec<CredentialScope>,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct CredentialScope {
pub profile: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub allowed_methods: Vec<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub allowed_tags: Vec<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub requires_approval: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct PolicyMetadata {
pub name: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub author: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub created: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub modified: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct PolicyDefaults {
#[serde(default = "default_allow_methods")]
pub allow_methods: Vec<String>,
#[serde(default = "default_true")]
pub deny_external_refs: bool,
#[serde(default = "default_true")]
pub require_auth: bool,
#[serde(default = "default_audit_level")]
pub audit_level: String,
#[serde(default = "default_classification")]
pub default_classification: DataClassification,
#[serde(default)]
pub read_only: bool,
#[serde(default)]
pub max_calls_per_session: u32,
#[serde(default)]
pub max_calls_per_minute: u32,
#[serde(default)]
pub warn_on_confidential_to_llm: bool,
#[serde(default)]
pub block_regulated_to_llm: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct PolicyRule {
pub name: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
pub pattern: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub conditions: Option<Vec<PolicyCondition>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub allow: Option<PolicyAction>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub deny: Option<PolicyAction>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub audit: Option<AuditConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub explain: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct PolicyCondition {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub auth_profile: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub time_window: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub source_ip: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub environment: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct PolicyAction {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub operations: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub methods: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub all: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tags: Option<Vec<String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize, JsonSchema)]
pub struct AuditConfig {
#[serde(default = "default_audit_level")]
pub level: String,
#[serde(default)]
pub include_body: bool,
#[serde(default)]
pub include_response: bool,
}
fn default_allow_methods() -> Vec<String> {
vec!["GET".to_string(), "HEAD".to_string(), "OPTIONS".to_string()]
}
fn default_true() -> bool {
true
}
fn default_audit_level() -> String {
"basic".to_string()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_policy_defaults() {
let defaults = PolicyDefaults {
allow_methods: default_allow_methods(),
deny_external_refs: true,
require_auth: true,
audit_level: "basic".to_string(),
default_classification: DataClassification::Internal,
read_only: false,
max_calls_per_session: 0,
max_calls_per_minute: 0,
warn_on_confidential_to_llm: false,
block_regulated_to_llm: false,
};
assert_eq!(defaults.allow_methods, vec!["GET", "HEAD", "OPTIONS"]);
assert!(defaults.deny_external_refs);
assert!(defaults.require_auth);
}
#[test]
fn test_data_flow_defaults() {
let yaml = r#"
version: "1.0"
defaults:
allow_methods: ["GET"]
require_auth: false
audit_level: "basic"
rules: []
"#;
let policy: PolicySet = serde_yaml::from_str(yaml).unwrap();
assert_eq!(policy.defaults.warn_on_confidential_to_llm, false);
assert_eq!(policy.defaults.block_regulated_to_llm, false);
}
#[test]
fn test_policy_rule_serialization() {
let rule = PolicyRule {
name: "test-rule".to_string(),
description: Some("Test rule".to_string()),
pattern: "*.example.com/*".to_string(),
conditions: None,
allow: Some(PolicyAction {
methods: Some(vec!["GET".to_string()]),
operations: None,
all: None,
tags: None,
}),
deny: None,
audit: None,
explain: None,
};
let json = serde_json::to_string_pretty(&rule).unwrap();
assert!(json.contains("test-rule"));
assert!(json.contains("*.example.com/*"));
}
#[test]
fn test_backward_compat_old_policy_yaml() {
let yaml = r#"
version: "1.0"
defaults:
allow_methods: ["GET"]
deny_external_refs: true
require_auth: true
audit_level: "basic"
rules: []
"#;
let policy: PolicySet = serde_yaml::from_str(yaml).unwrap();
assert_eq!(policy.defaults.read_only, false);
assert_eq!(policy.defaults.max_calls_per_session, 0);
assert_eq!(policy.defaults.max_calls_per_minute, 0);
assert_eq!(policy.defaults.allow_methods, vec!["GET"]);
assert!(policy.defaults.require_auth);
}
#[test]
fn test_new_policy_fields_yaml() {
let yaml = r#"
version: "1.0"
defaults:
allow_methods: ["GET"]
require_auth: false
audit_level: "basic"
read_only: true
max_calls_per_session: 100
max_calls_per_minute: 10
rules: []
"#;
let policy: PolicySet = serde_yaml::from_str(yaml).unwrap();
assert!(policy.defaults.read_only);
assert_eq!(policy.defaults.max_calls_per_session, 100);
assert_eq!(policy.defaults.max_calls_per_minute, 10);
}
#[test]
fn test_classification_serde() {
let yaml = r#"
version: "1.0"
defaults:
allow_methods: ["GET"]
require_auth: false
audit_level: "basic"
default_classification: confidential
rules: []
classifications:
"getPortfolio": confidential
"getStockPrice": public
"getThesis": internal
"*admin*": regulated
"#;
let policy: PolicySet = serde_yaml::from_str(yaml).unwrap();
assert_eq!(
policy.defaults.default_classification,
DataClassification::Confidential
);
assert_eq!(
policy.classifications.get("getPortfolio"),
Some(&DataClassification::Confidential)
);
assert_eq!(
policy.classifications.get("getStockPrice"),
Some(&DataClassification::Public)
);
assert_eq!(
policy.classifications.get("getThesis"),
Some(&DataClassification::Internal)
);
assert_eq!(
policy.classifications.get("*admin*"),
Some(&DataClassification::Regulated)
);
let serialized = serde_yaml::to_string(&policy).unwrap();
let roundtripped: PolicySet = serde_yaml::from_str(&serialized).unwrap();
assert_eq!(
roundtripped.classifications.get("getPortfolio"),
Some(&DataClassification::Confidential)
);
assert_eq!(
roundtripped.defaults.default_classification,
DataClassification::Confidential
);
}
#[test]
fn test_backward_compat_no_classifications() {
let yaml = r#"
version: "1.0"
defaults:
allow_methods: ["GET"]
deny_external_refs: true
require_auth: true
audit_level: "basic"
rules: []
"#;
let policy: PolicySet = serde_yaml::from_str(yaml).unwrap();
assert!(policy.classifications.is_empty());
assert_eq!(
policy.defaults.default_classification,
DataClassification::Internal
);
assert!(policy.credential_scopes.is_empty());
}
#[test]
fn test_credential_scope_serde() {
let yaml = r#"
version: "1.0"
defaults:
allow_methods: ["GET"]
require_auth: false
audit_level: "basic"
rules: []
credential_scopes:
- profile: "agent-readonly"
allowed_methods: ["GET"]
allowed_tags: ["portfolio", "stocks"]
- profile: "agent-full"
allowed_methods: ["GET", "POST"]
requires_approval: ["DELETE", "PATCH"]
"#;
let policy: PolicySet = serde_yaml::from_str(yaml).unwrap();
assert_eq!(policy.credential_scopes.len(), 2);
let readonly = &policy.credential_scopes[0];
assert_eq!(readonly.profile, "agent-readonly");
assert_eq!(readonly.allowed_methods, vec!["GET"]);
assert_eq!(readonly.allowed_tags, vec!["portfolio", "stocks"]);
assert!(readonly.requires_approval.is_empty());
let full = &policy.credential_scopes[1];
assert_eq!(full.profile, "agent-full");
assert_eq!(full.allowed_methods, vec!["GET", "POST"]);
assert!(full.allowed_tags.is_empty());
assert_eq!(full.requires_approval, vec!["DELETE", "PATCH"]);
let serialized = serde_yaml::to_string(&policy).unwrap();
let roundtripped: PolicySet = serde_yaml::from_str(&serialized).unwrap();
assert_eq!(roundtripped.credential_scopes.len(), 2);
assert_eq!(roundtripped.credential_scopes[0].profile, "agent-readonly");
assert_eq!(
roundtripped.credential_scopes[1].requires_approval,
vec!["DELETE", "PATCH"]
);
}
}