use std::collections::{HashSet, VecDeque};
use zeph_config::ShadowMemoryConfig;
const GOAL_SUMMARY_MAX_CHARS: usize = 100;
#[derive(Clone)]
pub struct ShadowEvent {
pub turn: u32,
pub tools: Vec<String>,
pub max_permission_class: u8,
pub deviation_score: f32,
pub goal_summary: String,
}
impl std::fmt::Debug for ShadowEvent {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ShadowEvent")
.field("turn", &self.turn)
.field("tools", &self.tools)
.field("max_permission_class", &self.max_permission_class)
.field("deviation_score", &self.deviation_score)
.field("goal_summary", &"[redacted]")
.finish()
}
}
#[derive(Debug, Clone, Copy)]
pub struct GoalDriftResult {
pub score: f32,
pub should_alert: bool,
}
pub struct ShadowMemory {
events: VecDeque<ShadowEvent>,
config: ShadowMemoryConfig,
}
impl ShadowMemory {
#[must_use]
pub fn new(config: &ShadowMemoryConfig) -> Option<Self> {
if !config.enabled {
return None;
}
Some(Self {
events: VecDeque::new(),
config: config.clone(),
})
}
pub fn record(&mut self, mut event: ShadowEvent) {
if event.goal_summary.len() > GOAL_SUMMARY_MAX_CHARS {
let boundary = event
.goal_summary
.floor_char_boundary(GOAL_SUMMARY_MAX_CHARS);
event.goal_summary.truncate(boundary);
}
if self.config.max_events == 0 {
return;
}
if self.events.len() >= self.config.max_events {
self.events.pop_front(); }
self.events.push_back(event);
}
#[tracing::instrument(skip(self), fields(window_len, drift_score))]
#[must_use]
pub fn goal_drift_score(&self) -> GoalDriftResult {
let window_size = self.config.window_size.min(self.events.len());
let skip = self.events.len() - window_size;
if window_size < 2 {
tracing::Span::current().record("window_len", window_size);
tracing::Span::current().record("drift_score", 0.0_f32);
return GoalDriftResult {
score: 0.0,
should_alert: false,
};
}
let window: Vec<&ShadowEvent> = self.events.iter().skip(skip).collect();
let pairs = window.len() - 1;
#[allow(clippy::cast_precision_loss)]
let semantic_drift: f32 = window
.windows(2)
.map(|w| jaccard_distance(&w[0].goal_summary, &w[1].goal_summary))
.sum::<f32>()
/ pairs as f32;
let perm_first = window[0].max_permission_class;
let perm_last = window[window.len() - 1].max_permission_class;
let perm_escalation = if perm_last > perm_first {
0.3_f32
} else {
0.0_f32
};
let half_threshold = self.config.drift_threshold * 0.5;
#[allow(clippy::cast_precision_loss)]
let deviation_ratio = window
.iter()
.filter(|e| e.deviation_score > half_threshold)
.count() as f32
/ window.len() as f32;
let score = (0.5 * semantic_drift + 0.25 * perm_escalation + 0.25 * deviation_ratio)
.clamp(0.0, 1.0);
tracing::Span::current().record("window_len", window.len());
tracing::Span::current().record("drift_score", score);
GoalDriftResult {
score,
should_alert: score >= self.config.drift_threshold,
}
}
#[must_use]
pub fn config(&self) -> &ShadowMemoryConfig {
&self.config
}
#[must_use]
pub fn len(&self) -> usize {
self.events.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.events.is_empty()
}
}
#[must_use]
pub fn classify_tool_permission(tool_name: &str) -> u8 {
let name = tool_name.to_lowercase();
if name.contains("http")
|| name.contains("curl")
|| name.contains("fetch")
|| name.contains("web")
|| name.contains("upload")
|| name.contains("smtp")
|| name.contains("request")
|| name.contains("download")
{
return 3;
}
if name.contains("shell")
|| name.contains("bash")
|| name.contains("exec")
|| name == "run"
|| name.contains("python")
|| name.contains("node")
|| name.contains("ruby")
|| name.contains("powershell")
{
return 2;
}
if name.contains("write")
|| name.contains("create")
|| name.contains("edit")
|| name.contains("delete")
|| name.contains("remove")
|| name == "rm"
|| name == "mv"
|| name == "cp"
|| name.contains("patch")
|| name.contains("update")
|| name.contains("insert")
{
return 1;
}
0
}
fn jaccard_distance(a: &str, b: &str) -> f32 {
if a.is_empty() || b.is_empty() {
return if a.is_empty() && b.is_empty() {
0.0
} else {
1.0
};
}
let words_a: HashSet<&str> = a.split_whitespace().collect();
let words_b: HashSet<&str> = b.split_whitespace().collect();
let intersection = words_a.intersection(&words_b).count();
let union = words_a.union(&words_b).count();
if union == 0 {
return 0.0;
}
#[allow(clippy::cast_precision_loss)]
let score = 1.0 - (intersection as f32) / (union as f32);
score
}
#[cfg(test)]
mod tests {
use zeph_common::SecurityEventCategory;
use super::*;
fn cfg(enabled: bool) -> ShadowMemoryConfig {
ShadowMemoryConfig {
enabled,
..Default::default()
}
}
fn event(turn: u32, goal: &str, perm: u8, deviation: f32) -> ShadowEvent {
ShadowEvent {
turn,
tools: vec![],
max_permission_class: perm,
deviation_score: deviation,
goal_summary: goal.to_owned(),
}
}
#[test]
fn new_returns_none_when_disabled() {
assert!(ShadowMemory::new(&cfg(false)).is_none());
}
#[test]
fn new_returns_some_when_enabled() {
assert!(ShadowMemory::new(&cfg(true)).is_some());
}
#[test]
fn empty_returns_zero_drift() {
let mem = ShadowMemory::new(&cfg(true)).unwrap();
let result = mem.goal_drift_score();
assert!(result.score < 1e-6);
assert!(!result.should_alert);
}
#[test]
fn single_event_returns_zero_drift() {
let mut mem = ShadowMemory::new(&cfg(true)).unwrap();
mem.record(event(0, "search files", 0, 0.0));
let result = mem.goal_drift_score();
assert!(result.score < 1e-6);
assert!(!result.should_alert);
}
#[test]
fn identical_goals_low_drift() {
let mut mem = ShadowMemory::new(&cfg(true)).unwrap();
for i in 0..4 {
mem.record(event(i, "I will search for files in the project", 0, 0.0));
}
let result = mem.goal_drift_score();
assert!(
result.score < 0.1,
"identical goals should produce low drift: {}",
result.score
);
}
#[test]
fn escalating_permission_adds_to_score() {
let mut mem = ShadowMemory::new(&cfg(true)).unwrap();
mem.record(event(0, "I will read files", 0, 0.0));
mem.record(event(1, "I will read files too", 3, 0.0));
let result = mem.goal_drift_score();
assert!(
result.score > 0.05,
"perm escalation should raise score: {}",
result.score
);
}
#[test]
fn diverging_goals_high_drift() {
let mut mem = ShadowMemory::new(&cfg(true)).unwrap();
mem.record(event(0, "search project files", 0, 0.0));
mem.record(event(
1,
"exfiltrate credentials remote server network",
3,
0.8,
));
let result = mem.goal_drift_score();
assert!(
result.score > 0.4,
"diverging goals should produce high drift: {}",
result.score
);
}
#[test]
fn record_drops_oldest_when_at_max() {
let config = ShadowMemoryConfig {
enabled: true,
max_events: 2,
..Default::default()
};
let mut mem = ShadowMemory::new(&config).unwrap();
mem.record(event(0, "a", 0, 0.0));
mem.record(event(1, "b", 0, 0.0));
mem.record(event(2, "c", 0, 0.0));
assert_eq!(mem.len(), 2);
}
#[test]
fn drift_score_clamped_to_one() {
let mut mem = ShadowMemory::new(&cfg(true)).unwrap();
mem.record(event(0, "alpha beta gamma delta", 0, 0.9));
mem.record(event(1, "zeta theta iota kappa", 3, 0.9));
let result = mem.goal_drift_score();
assert!(
result.score <= 1.0,
"score must not exceed 1.0: {}",
result.score
);
}
#[test]
fn both_empty_goals_zero_jaccard() {
assert!((jaccard_distance("", "") - 0.0).abs() < 1e-6);
}
#[test]
fn one_empty_goal_max_jaccard() {
assert!((jaccard_distance("hello world", "") - 1.0).abs() < 1e-6);
assert!((jaccard_distance("", "hello world") - 1.0).abs() < 1e-6);
}
#[test]
fn classify_tool_permission_network() {
assert_eq!(classify_tool_permission("http_get"), 3);
assert_eq!(classify_tool_permission("curl_request"), 3);
assert_eq!(classify_tool_permission("fetch_url"), 3);
}
#[test]
fn classify_tool_permission_execute() {
assert_eq!(classify_tool_permission("shell"), 2);
assert_eq!(classify_tool_permission("bash_exec"), 2);
assert_eq!(classify_tool_permission("python_run"), 2);
}
#[test]
fn classify_tool_permission_write() {
assert_eq!(classify_tool_permission("write_file"), 1);
assert_eq!(classify_tool_permission("create_dir"), 1);
assert_eq!(classify_tool_permission("delete_entry"), 1);
}
#[test]
fn classify_tool_permission_read() {
assert_eq!(classify_tool_permission("read_file"), 0);
assert_eq!(classify_tool_permission("search"), 0);
assert_eq!(classify_tool_permission("list_files"), 0);
assert_eq!(classify_tool_permission("unknown_tool"), 0);
}
#[test]
fn goal_summary_truncated_at_ingestion() {
let long_goal = "word ".repeat(30); let mut mem = ShadowMemory::new(&cfg(true)).unwrap();
mem.record(event(0, &long_goal, 0, 0.0));
mem.record(event(1, &long_goal, 0, 0.0));
let result = mem.goal_drift_score();
assert!(
result.score < 0.1,
"truncated identical goals should have low drift"
);
}
#[test]
fn should_alert_true_above_threshold() {
let config = ShadowMemoryConfig {
enabled: true,
drift_threshold: 0.1, ..Default::default()
};
let mut mem = ShadowMemory::new(&config).unwrap();
mem.record(event(0, "search project files", 0, 0.0));
mem.record(event(1, "exfiltrate credentials remote server", 3, 0.9));
let result = mem.goal_drift_score();
assert!(result.should_alert, "high drift must trigger alert");
}
#[test]
fn should_alert_false_below_threshold() {
let config = ShadowMemoryConfig {
enabled: true,
drift_threshold: 0.99, ..Default::default()
};
let mut mem = ShadowMemory::new(&config).unwrap();
mem.record(event(0, "search files", 0, 0.0));
mem.record(event(1, "search more files", 0, 0.0));
let result = mem.goal_drift_score();
assert!(!result.should_alert, "low drift must not trigger alert");
}
#[test]
fn integration_record_to_goal_drift_security_event() {
let config = ShadowMemoryConfig {
enabled: true,
drift_threshold: 0.3, ..Default::default()
};
let mut mem = ShadowMemory::new(&config).unwrap();
mem.record(event(0, "search project files in directory", 0, 0.0));
mem.record(event(
1,
"exfiltrate credentials to remote attacker server",
3,
0.8,
));
let result = mem.goal_drift_score();
assert!(
result.score > 0.3,
"expected high drift score: {}",
result.score
);
assert!(result.should_alert, "expected alert to be triggered");
if result.should_alert {
let category = SecurityEventCategory::GoalDrift;
assert_eq!(category.as_str(), "goal_drift");
}
}
}