1use serde::{Deserialize, Serialize};
10use std::fmt;
11
12#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
17pub enum TaintLevel {
18 Public,
20 #[default]
22 Internal,
23 Private,
25}
26
27impl TaintLevel {
28 fn rank(self) -> u8 {
30 match self {
31 TaintLevel::Public => 0,
32 TaintLevel::Internal => 1,
33 TaintLevel::Private => 2,
34 }
35 }
36
37 pub fn max(self, other: TaintLevel) -> TaintLevel {
39 if self >= other { self } else { other }
40 }
41
42 pub fn from_str_loose(s: &str) -> Option<TaintLevel> {
44 match s.to_lowercase().as_str() {
45 "public" => Some(TaintLevel::Public),
46 "internal" => Some(TaintLevel::Internal),
47 "private" => Some(TaintLevel::Private),
48 _ => None,
49 }
50 }
51
52 pub fn as_str(self) -> &'static str {
54 match self {
55 TaintLevel::Public => "public",
56 TaintLevel::Internal => "internal",
57 TaintLevel::Private => "private",
58 }
59 }
60}
61
62impl PartialOrd for TaintLevel {
63 fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
64 Some(self.cmp(other))
65 }
66}
67
68impl Ord for TaintLevel {
69 fn cmp(&self, other: &Self) -> std::cmp::Ordering {
70 self.rank().cmp(&other.rank())
71 }
72}
73
74impl fmt::Display for TaintLevel {
75 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
76 f.write_str(self.as_str())
77 }
78}
79
80#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
82pub enum ToolDirection {
83 Inbound,
85 #[default]
87 Internal,
88 Outbound,
90}
91
92impl ToolDirection {
93 pub fn from_str_loose(s: &str) -> Option<ToolDirection> {
95 match s.to_lowercase().as_str() {
96 "inbound" => Some(ToolDirection::Inbound),
97 "internal" => Some(ToolDirection::Internal),
98 "outbound" => Some(ToolDirection::Outbound),
99 _ => None,
100 }
101 }
102
103 pub fn as_str(self) -> &'static str {
105 match self {
106 ToolDirection::Inbound => "inbound",
107 ToolDirection::Internal => "internal",
108 ToolDirection::Outbound => "outbound",
109 }
110 }
111}
112
113impl fmt::Display for ToolDirection {
114 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
115 f.write_str(self.as_str())
116 }
117}
118
119#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
125pub struct ToolClassification {
126 pub sensitivity: TaintLevel,
128 pub direction: ToolDirection,
130 pub clearance: TaintLevel,
133}
134
135impl ToolClassification {
136 pub fn new(sensitivity: TaintLevel, direction: ToolDirection, clearance: TaintLevel) -> Self {
138 Self {
139 sensitivity,
140 direction,
141 clearance,
142 }
143 }
144
145 pub fn is_outbound(&self) -> bool {
147 self.direction == ToolDirection::Outbound
148 }
149
150 pub fn check_clearance(&self, taint: TaintLevel) -> bool {
154 if !self.is_outbound() {
155 return true;
156 }
157 taint <= self.clearance
158 }
159}
160
161impl Default for ToolClassification {
162 fn default() -> Self {
163 Self {
164 sensitivity: TaintLevel::Internal,
165 direction: ToolDirection::Internal,
166 clearance: TaintLevel::Public,
167 }
168 }
169}
170
171#[derive(Debug, Clone, Serialize, Deserialize)]
176pub struct RegionTaint {
177 current_level: TaintLevel,
179 entry_taints: Vec<TaintLevel>,
181}
182
183impl RegionTaint {
184 pub fn new() -> Self {
186 Self {
187 current_level: TaintLevel::Public,
188 entry_taints: Vec::new(),
189 }
190 }
191
192 pub fn level(&self) -> TaintLevel {
194 self.current_level
195 }
196
197 pub fn add_entry(&mut self, taint: TaintLevel) {
200 self.entry_taints.push(taint);
201 self.current_level = self.current_level.max(taint);
202 }
203
204 pub fn remove_oldest(&mut self) {
207 if !self.entry_taints.is_empty() {
208 self.entry_taints.remove(0);
209 self.recompute();
210 }
211 }
212
213 pub fn remove_at(&mut self, idx: usize) {
216 if idx < self.entry_taints.len() {
217 self.entry_taints.remove(idx);
218 self.recompute();
219 }
220 }
221
222 pub fn clear(&mut self) {
224 self.entry_taints.clear();
225 self.current_level = TaintLevel::Public;
226 }
227
228 pub fn recompute(&mut self) {
231 self.current_level = self
232 .entry_taints
233 .iter()
234 .copied()
235 .max()
236 .unwrap_or(TaintLevel::Public);
237 }
238
239 pub fn entry_count(&self) -> usize {
241 self.entry_taints.len()
242 }
243
244 pub fn from_entry_taints(entry_taints: Vec<TaintLevel>) -> Self {
251 let current_level = entry_taints
252 .iter()
253 .copied()
254 .max()
255 .unwrap_or(TaintLevel::Public);
256 Self {
257 current_level,
258 entry_taints,
259 }
260 }
261
262 pub fn entry_taint(&self, index: usize) -> Option<TaintLevel> {
268 self.entry_taints.get(index).copied()
269 }
270}
271
272impl Default for RegionTaint {
273 fn default() -> Self {
274 Self::new()
275 }
276}
277
278#[derive(Debug, Clone, Serialize, Deserialize)]
280pub struct SecurityConfig {
281 pub taint_tracking: bool,
283}
284
285impl Default for SecurityConfig {
286 fn default() -> Self {
287 Self {
295 taint_tracking: true,
296 }
297 }
298}
299
300pub fn resolve_taint_enabled(
310 global: bool,
311 agent: Option<&SecurityConfig>,
312 stage: Option<&SecurityConfig>,
313) -> bool {
314 let manifest = stage
315 .map(|s| s.taint_tracking)
316 .or_else(|| agent.map(|a| a.taint_tracking));
317 global || manifest.unwrap_or(false)
318}
319
320pub fn resolve_security(
327 global: bool,
328 agent: Option<&SecurityConfig>,
329 stage: Option<&SecurityConfig>,
330) -> SecurityConfig {
331 let mut resolved = match stage.or(agent) {
332 Some(c) => c.clone(),
333 None => SecurityConfig {
334 taint_tracking: global,
335 },
336 };
337 resolved.taint_tracking = resolve_taint_enabled(global, agent, stage);
338 resolved
339}
340
341fn resolve_hint(global: bool, agent: Option<bool>, stage: Option<bool>) -> bool {
347 stage.or(agent).unwrap_or(global)
348}
349
350pub fn resolve_batch_tool_hint(global: bool, agent: Option<bool>, stage: Option<bool>) -> bool {
354 resolve_hint(global, agent, stage)
355}
356
357pub fn resolve_shell_hint(global: bool, agent: Option<bool>, stage: Option<bool>) -> bool {
364 resolve_hint(global, agent, stage)
365}
366
367#[derive(Debug, Clone, PartialEq, Eq)]
369pub enum GateDecision {
370 Allowed,
372 Blocked {
374 taint_level: TaintLevel,
376 clearance: TaintLevel,
378 source_regions: Vec<String>,
380 tool_name: String,
382 },
383}
384
385impl GateDecision {
386 pub fn is_allowed(&self) -> bool {
388 matches!(self, GateDecision::Allowed)
389 }
390
391 pub fn blocked_levels(&self) -> Option<(TaintLevel, TaintLevel)> {
394 match self {
395 GateDecision::Blocked {
396 taint_level,
397 clearance,
398 ..
399 } => Some((*taint_level, *clearance)),
400 GateDecision::Allowed => None,
401 }
402 }
403}
404
405#[derive(Debug, Clone, Serialize, Deserialize)]
407pub struct GateEvent {
408 pub timestamp: i64,
410 pub agent_id: String,
412 pub tool_name: String,
414 pub taint_level: TaintLevel,
416 pub clearance: TaintLevel,
418 pub allowed: bool,
420 pub decision_source: GateDecisionSource,
422}
423
424#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
426pub enum GateDecisionSource {
427 AutoAllow,
429 AutoBlock,
431 AllowlistRule {
433 rule_index: usize,
436 },
437 ScriptedRule {
439 script_name: String,
441 },
442 UserAllowOnce,
444 UserAlwaysAllow,
446 UserDenied,
448 TaintDisabled,
450 YoloAutoApprove,
454}
455
456pub fn builtin_tool_classification(tool_name: &str) -> ToolClassification {
471 match tool_name {
472 "read_file" => ToolClassification::new(
473 TaintLevel::Internal,
474 ToolDirection::Inbound,
475 TaintLevel::Public,
476 ),
477 "write_file" => ToolClassification::new(
478 TaintLevel::Internal,
479 ToolDirection::Internal,
480 TaintLevel::Public,
481 ),
482 "edit_file" => ToolClassification::new(
483 TaintLevel::Internal,
484 ToolDirection::Internal,
485 TaintLevel::Public,
486 ),
487 "list_dir" => ToolClassification::new(
488 TaintLevel::Internal,
489 ToolDirection::Inbound,
490 TaintLevel::Public,
491 ),
492 "shell" | "bash" => ToolClassification::new(
493 TaintLevel::Public,
494 ToolDirection::Outbound,
495 TaintLevel::Public,
496 ),
497 "web_search" | "web_fetch" | "http_get" | "http_post" | "fetch" => ToolClassification::new(
501 TaintLevel::Public,
502 ToolDirection::Outbound,
503 TaintLevel::Public,
504 ),
505 "current_time" | "locale_info" => ToolClassification::new(
514 TaintLevel::Public,
515 ToolDirection::Inbound,
516 TaintLevel::Public,
517 ),
518 "system_info" | "environment_info" | "which_command" | "runtime_info" => {
519 ToolClassification::new(
520 TaintLevel::Internal,
521 ToolDirection::Inbound,
522 TaintLevel::Public,
523 )
524 }
525 "ask_user_text" | "ask_user_choice" | "ask_user_confirm" | "present_for_review" => {
526 ToolClassification::new(
527 TaintLevel::Internal,
528 ToolDirection::Internal,
529 TaintLevel::Public,
530 )
531 }
532 "spawn_agent" | "check_agent" | "wait_for_agent" | "send_to_agent" | "kill_agent" => {
533 ToolClassification::new(
534 TaintLevel::Internal,
535 ToolDirection::Internal,
536 TaintLevel::Public,
537 )
538 }
539 _ => ToolClassification::new(
545 TaintLevel::Public,
546 ToolDirection::Outbound,
547 TaintLevel::Public,
548 ),
549 }
550}
551
552#[cfg(test)]
553mod tests {
554 use super::*;
555
556 #[test]
560 fn taint_rebuilds_from_persisted_entries_at_the_highest_level() {
561 let restored = RegionTaint::from_entry_taints(vec![
562 TaintLevel::Public,
563 TaintLevel::Private,
564 TaintLevel::Internal,
565 ]);
566 assert_eq!(restored.level(), TaintLevel::Private);
567 assert_eq!(restored.entry_taint(1), Some(TaintLevel::Private));
568
569 assert_eq!(
572 RegionTaint::from_entry_taints(Vec::new()).level(),
573 TaintLevel::Public
574 );
575 }
576
577 #[test]
580 fn taint_level_ordering() {
581 assert!(TaintLevel::Public < TaintLevel::Internal);
582 assert!(TaintLevel::Internal < TaintLevel::Private);
583 assert!(TaintLevel::Public < TaintLevel::Private);
584 }
585
586 #[test]
587 fn taint_level_equality() {
588 assert_eq!(TaintLevel::Public, TaintLevel::Public);
589 assert_eq!(TaintLevel::Internal, TaintLevel::Internal);
590 assert_eq!(TaintLevel::Private, TaintLevel::Private);
591 assert_ne!(TaintLevel::Public, TaintLevel::Private);
592 }
593
594 #[test]
595 fn taint_level_max() {
596 assert_eq!(
597 TaintLevel::Public.max(TaintLevel::Internal),
598 TaintLevel::Internal
599 );
600 assert_eq!(
601 TaintLevel::Private.max(TaintLevel::Public),
602 TaintLevel::Private
603 );
604 assert_eq!(
605 TaintLevel::Internal.max(TaintLevel::Internal),
606 TaintLevel::Internal
607 );
608 }
609
610 #[test]
611 fn taint_level_default_is_internal() {
612 assert_eq!(TaintLevel::default(), TaintLevel::Internal);
613 }
614
615 #[test]
616 fn taint_level_display() {
617 assert_eq!(format!("{}", TaintLevel::Public), "public");
618 assert_eq!(format!("{}", TaintLevel::Internal), "internal");
619 assert_eq!(format!("{}", TaintLevel::Private), "private");
620 }
621
622 #[test]
623 fn taint_level_from_str_loose() {
624 assert_eq!(
625 TaintLevel::from_str_loose("public"),
626 Some(TaintLevel::Public)
627 );
628 assert_eq!(
629 TaintLevel::from_str_loose("INTERNAL"),
630 Some(TaintLevel::Internal)
631 );
632 assert_eq!(
633 TaintLevel::from_str_loose("Private"),
634 Some(TaintLevel::Private)
635 );
636 assert_eq!(TaintLevel::from_str_loose("unknown"), None);
637 }
638
639 #[test]
640 fn taint_level_as_str() {
641 assert_eq!(TaintLevel::Public.as_str(), "public");
642 assert_eq!(TaintLevel::Internal.as_str(), "internal");
643 assert_eq!(TaintLevel::Private.as_str(), "private");
644 }
645
646 #[test]
647 fn taint_level_serde_roundtrip() {
648 for level in [
649 TaintLevel::Public,
650 TaintLevel::Internal,
651 TaintLevel::Private,
652 ] {
653 let json = serde_json::to_string(&level).unwrap();
654 let back: TaintLevel = serde_json::from_str(&json).unwrap();
655 assert_eq!(level, back);
656 }
657 }
658
659 #[test]
660 fn taint_level_hash() {
661 use std::collections::HashSet;
662 let mut set = HashSet::new();
663 set.insert(TaintLevel::Public);
664 set.insert(TaintLevel::Internal);
665 set.insert(TaintLevel::Private);
666 set.insert(TaintLevel::Public); assert_eq!(set.len(), 3);
668 }
669
670 #[test]
673 fn tool_direction_from_str_loose() {
674 assert_eq!(
675 ToolDirection::from_str_loose("inbound"),
676 Some(ToolDirection::Inbound)
677 );
678 assert_eq!(
679 ToolDirection::from_str_loose("OUTBOUND"),
680 Some(ToolDirection::Outbound)
681 );
682 assert_eq!(
683 ToolDirection::from_str_loose("Internal"),
684 Some(ToolDirection::Internal)
685 );
686 assert_eq!(ToolDirection::from_str_loose("nope"), None);
687 }
688
689 #[test]
690 fn tool_direction_default_is_internal() {
691 assert_eq!(ToolDirection::default(), ToolDirection::Internal);
692 }
693
694 #[test]
695 fn tool_direction_display() {
696 assert_eq!(format!("{}", ToolDirection::Inbound), "inbound");
697 assert_eq!(format!("{}", ToolDirection::Internal), "internal");
698 assert_eq!(format!("{}", ToolDirection::Outbound), "outbound");
699 }
700
701 #[test]
702 fn tool_direction_serde_roundtrip() {
703 for dir in [
704 ToolDirection::Inbound,
705 ToolDirection::Internal,
706 ToolDirection::Outbound,
707 ] {
708 let json = serde_json::to_string(&dir).unwrap();
709 let back: ToolDirection = serde_json::from_str(&json).unwrap();
710 assert_eq!(dir, back);
711 }
712 }
713
714 #[test]
717 fn tool_classification_default() {
718 let tc = ToolClassification::default();
719 assert_eq!(tc.sensitivity, TaintLevel::Internal);
720 assert_eq!(tc.direction, ToolDirection::Internal);
721 assert_eq!(tc.clearance, TaintLevel::Public);
722 }
723
724 #[test]
725 fn tool_classification_outbound_check() {
726 let tc = ToolClassification::new(
727 TaintLevel::Public,
728 ToolDirection::Outbound,
729 TaintLevel::Internal,
730 );
731 assert!(tc.is_outbound());
732 assert!(tc.check_clearance(TaintLevel::Public));
733 assert!(tc.check_clearance(TaintLevel::Internal));
734 assert!(!tc.check_clearance(TaintLevel::Private));
735 }
736
737 #[test]
738 fn tool_classification_non_outbound_always_passes() {
739 let tc = ToolClassification::new(
740 TaintLevel::Private,
741 ToolDirection::Inbound,
742 TaintLevel::Public, );
744 assert!(!tc.is_outbound());
745 assert!(tc.check_clearance(TaintLevel::Private));
746 }
747
748 #[test]
749 fn tool_classification_serde_roundtrip() {
750 let tc = ToolClassification::new(
751 TaintLevel::Private,
752 ToolDirection::Outbound,
753 TaintLevel::Internal,
754 );
755 let json = serde_json::to_string(&tc).unwrap();
756 let back: ToolClassification = serde_json::from_str(&json).unwrap();
757 assert_eq!(tc, back);
758 }
759
760 #[test]
763 fn region_taint_starts_public() {
764 let rt = RegionTaint::new();
765 assert_eq!(rt.level(), TaintLevel::Public);
766 assert_eq!(rt.entry_count(), 0);
767 }
768
769 #[test]
770 fn region_taint_add_entry_raises_level() {
771 let mut rt = RegionTaint::new();
772 rt.add_entry(TaintLevel::Internal);
773 assert_eq!(rt.level(), TaintLevel::Internal);
774 rt.add_entry(TaintLevel::Private);
775 assert_eq!(rt.level(), TaintLevel::Private);
776 }
777
778 #[test]
779 fn region_taint_add_public_doesnt_lower() {
780 let mut rt = RegionTaint::new();
781 rt.add_entry(TaintLevel::Private);
782 rt.add_entry(TaintLevel::Public);
783 assert_eq!(rt.level(), TaintLevel::Private);
784 }
785
786 #[test]
787 fn region_taint_remove_oldest_recovers() {
788 let mut rt = RegionTaint::new();
789 rt.add_entry(TaintLevel::Private);
790 rt.add_entry(TaintLevel::Public);
791 assert_eq!(rt.level(), TaintLevel::Private);
792
793 rt.remove_oldest(); assert_eq!(rt.level(), TaintLevel::Public);
795 }
796
797 #[test]
798 fn region_taint_remove_oldest_empty() {
799 let mut rt = RegionTaint::new();
800 rt.remove_oldest(); assert_eq!(rt.level(), TaintLevel::Public);
802 }
803
804 #[test]
805 fn region_taint_clear() {
806 let mut rt = RegionTaint::new();
807 rt.add_entry(TaintLevel::Private);
808 rt.add_entry(TaintLevel::Internal);
809 rt.clear();
810 assert_eq!(rt.level(), TaintLevel::Public);
811 assert_eq!(rt.entry_count(), 0);
812 }
813
814 #[test]
815 fn region_taint_recompute() {
816 let mut rt = RegionTaint::new();
817 rt.add_entry(TaintLevel::Private);
818 rt.add_entry(TaintLevel::Internal);
819 rt.add_entry(TaintLevel::Public);
820 assert_eq!(rt.entry_count(), 3);
821
822 rt.remove_oldest();
824 assert_eq!(rt.level(), TaintLevel::Internal);
825 assert_eq!(rt.entry_count(), 2);
826 }
827
828 #[test]
829 fn region_taint_entry_taint() {
830 let mut rt = RegionTaint::new();
831 rt.add_entry(TaintLevel::Public);
832 rt.add_entry(TaintLevel::Private);
833 assert_eq!(rt.entry_taint(0), Some(TaintLevel::Public));
834 assert_eq!(rt.entry_taint(1), Some(TaintLevel::Private));
835 assert_eq!(rt.entry_taint(2), None);
836 }
837
838 #[test]
839 fn region_taint_default() {
840 let rt = RegionTaint::default();
841 assert_eq!(rt.level(), TaintLevel::Public);
842 }
843
844 #[test]
845 fn region_taint_serde_roundtrip() {
846 let mut rt = RegionTaint::new();
847 rt.add_entry(TaintLevel::Internal);
848 rt.add_entry(TaintLevel::Private);
849 let json = serde_json::to_string(&rt).unwrap();
850 let back: RegionTaint = serde_json::from_str(&json).unwrap();
851 assert_eq!(back.level(), TaintLevel::Private);
852 assert_eq!(back.entry_count(), 2);
853 }
854
855 #[test]
858 fn security_config_default() {
859 let sc = SecurityConfig::default();
860 assert!(sc.taint_tracking);
861 }
862
863 #[test]
864 fn security_config_serde_roundtrip() {
865 let sc = SecurityConfig {
866 taint_tracking: false,
867 };
868 let json = serde_json::to_string(&sc).unwrap();
869 let back: SecurityConfig = serde_json::from_str(&json).unwrap();
870 assert!(!back.taint_tracking);
871 }
872
873 #[test]
876 fn gate_decision_allowed() {
877 let d = GateDecision::Allowed;
878 assert!(d.is_allowed());
879 }
880
881 #[test]
882 fn gate_decision_blocked() {
883 let d = GateDecision::Blocked {
884 taint_level: TaintLevel::Private,
885 clearance: TaintLevel::Public,
886 source_regions: vec!["conversation".into()],
887 tool_name: "send_email".into(),
888 };
889 assert!(!d.is_allowed());
890 }
891
892 #[test]
895 fn gate_event_serde_roundtrip() {
896 let event = GateEvent {
897 timestamp: 1234567890,
898 agent_id: "agent-1".into(),
899 tool_name: "send_email".into(),
900 taint_level: TaintLevel::Private,
901 clearance: TaintLevel::Public,
902 allowed: false,
903 decision_source: GateDecisionSource::UserDenied,
904 };
905 let json = serde_json::to_string(&event).unwrap();
906 let back: GateEvent = serde_json::from_str(&json).unwrap();
907 assert_eq!(back.agent_id, "agent-1");
908 assert!(!back.allowed);
909 }
910
911 #[test]
912 fn gate_decision_source_variants() {
913 let sources = vec![
914 GateDecisionSource::AutoAllow,
915 GateDecisionSource::AllowlistRule { rule_index: 0 },
916 GateDecisionSource::ScriptedRule {
917 script_name: "test.rhai".into(),
918 },
919 GateDecisionSource::UserAllowOnce,
920 GateDecisionSource::UserAlwaysAllow,
921 GateDecisionSource::UserDenied,
922 GateDecisionSource::TaintDisabled,
923 ];
924 for src in sources {
925 let json = serde_json::to_string(&src).unwrap();
926 let back: GateDecisionSource = serde_json::from_str(&json).unwrap();
927 assert_eq!(src, back);
928 }
929 }
930
931 #[test]
934 fn builtin_read_file_classification() {
935 let tc = builtin_tool_classification("read_file");
936 assert_eq!(tc.sensitivity, TaintLevel::Internal);
937 assert_eq!(tc.direction, ToolDirection::Inbound);
938 }
939
940 #[test]
941 fn builtin_shell_classification() {
942 let tc = builtin_tool_classification("shell");
943 assert_eq!(tc.sensitivity, TaintLevel::Public);
944 assert_eq!(tc.direction, ToolDirection::Outbound);
945 assert_eq!(tc.clearance, TaintLevel::Public);
946
947 let tc2 = builtin_tool_classification("bash");
949 assert_eq!(tc2.direction, ToolDirection::Outbound);
950 }
951
952 #[test]
959 fn network_capable_tools_are_outbound() {
960 for name in ["web_search", "web_fetch", "http_get", "http_post", "fetch"] {
961 let tc = builtin_tool_classification(name);
962 assert_eq!(tc.sensitivity, TaintLevel::Public, "{name}");
963 assert_eq!(tc.direction, ToolDirection::Outbound, "{name}");
964 }
965 }
966
967 #[test]
973 fn environment_tools_are_inbound_and_never_gated() {
974 for name in [
975 "current_time",
976 "system_info",
977 "locale_info",
978 "environment_info",
979 "which_command",
980 "runtime_info",
981 ] {
982 let tc = builtin_tool_classification(name);
983 assert_eq!(tc.direction, ToolDirection::Inbound, "{name}");
984 assert_eq!(tc.clearance, TaintLevel::Public, "{name}");
986 }
987 }
988
989 #[test]
994 fn environment_tools_are_graded_by_what_they_reveal() {
995 for name in ["current_time", "locale_info"] {
996 assert_eq!(
997 builtin_tool_classification(name).sensitivity,
998 TaintLevel::Public,
999 "{name}"
1000 );
1001 }
1002 for name in [
1003 "system_info",
1004 "environment_info",
1005 "which_command",
1006 "runtime_info",
1007 ] {
1008 assert_eq!(
1009 builtin_tool_classification(name).sensitivity,
1010 TaintLevel::Internal,
1011 "{name}"
1012 );
1013 }
1014 }
1015
1016 #[test]
1017 fn builtin_ask_user_classification() {
1018 for name in [
1019 "ask_user_text",
1020 "ask_user_choice",
1021 "ask_user_confirm",
1022 "present_for_review",
1023 ] {
1024 let tc = builtin_tool_classification(name);
1025 assert_eq!(tc.direction, ToolDirection::Internal);
1026 }
1027 }
1028
1029 #[test]
1030 fn builtin_subagent_classification() {
1031 for name in [
1032 "spawn_agent",
1033 "check_agent",
1034 "wait_for_agent",
1035 "send_to_agent",
1036 "kill_agent",
1037 ] {
1038 let tc = builtin_tool_classification(name);
1039 assert_eq!(tc.direction, ToolDirection::Internal);
1040 }
1041 }
1042
1043 #[test]
1044 fn builtin_write_file_classification() {
1045 let tc = builtin_tool_classification("write_file");
1046 assert_eq!(tc.direction, ToolDirection::Internal);
1047 }
1048
1049 #[test]
1054 fn unknown_tools_fail_closed_as_outbound() {
1055 let tc = builtin_tool_classification("some_mcp_tool");
1056 assert_eq!(tc.sensitivity, TaintLevel::Public);
1057 assert_eq!(tc.direction, ToolDirection::Outbound);
1058 assert_eq!(tc.clearance, TaintLevel::Public);
1059 }
1060
1061 #[test]
1062 fn builtin_edit_file_classification() {
1063 let tc = builtin_tool_classification("edit_file");
1064 assert_eq!(tc.sensitivity, TaintLevel::Internal);
1065 assert_eq!(tc.direction, ToolDirection::Internal);
1066 assert_eq!(tc.clearance, TaintLevel::Public);
1067 }
1068
1069 #[test]
1070 fn builtin_list_dir_classification() {
1071 let tc = builtin_tool_classification("list_dir");
1072 assert_eq!(tc.sensitivity, TaintLevel::Internal);
1073 assert_eq!(tc.direction, ToolDirection::Inbound);
1074 assert_eq!(tc.clearance, TaintLevel::Public);
1075 }
1076
1077 fn sec(taint: bool) -> SecurityConfig {
1080 SecurityConfig {
1081 taint_tracking: taint,
1082 }
1083 }
1084
1085 #[test]
1086 fn resolve_taint_enabled_inherits_global_when_unset() {
1087 assert!(!resolve_taint_enabled(false, None, None));
1088 assert!(resolve_taint_enabled(true, None, None));
1089 }
1090
1091 #[test]
1092 fn resolve_taint_enabled_agent_may_opt_in_but_not_out() {
1093 assert!(resolve_taint_enabled(false, Some(&sec(true)), None));
1095 assert!(resolve_taint_enabled(true, Some(&sec(false)), None));
1099 }
1100
1101 #[test]
1102 fn resolve_taint_enabled_stage_may_opt_in_but_not_out() {
1103 assert!(resolve_taint_enabled(
1105 false,
1106 Some(&sec(false)),
1107 Some(&sec(true))
1108 ));
1109 assert!(resolve_taint_enabled(
1111 true,
1112 Some(&sec(true)),
1113 Some(&sec(false))
1114 ));
1115 }
1116
1117 #[test]
1118 fn resolve_batch_tool_hint_cascade() {
1119 assert!(resolve_batch_tool_hint(true, None, None));
1121 assert!(!resolve_batch_tool_hint(false, None, None));
1122 assert!(!resolve_batch_tool_hint(true, Some(false), None));
1124 assert!(resolve_batch_tool_hint(false, Some(true), None));
1125 assert!(!resolve_batch_tool_hint(true, Some(true), Some(false)));
1127 assert!(resolve_batch_tool_hint(false, Some(false), Some(true)));
1128 }
1129
1130 #[test]
1131 fn gate_decision_blocked_levels() {
1132 let blocked = GateDecision::Blocked {
1133 taint_level: TaintLevel::Private,
1134 clearance: TaintLevel::Public,
1135 source_regions: vec![],
1136 tool_name: "shell".into(),
1137 };
1138 assert_eq!(
1139 blocked.blocked_levels(),
1140 Some((TaintLevel::Private, TaintLevel::Public))
1141 );
1142 assert_eq!(GateDecision::Allowed.blocked_levels(), None);
1143 }
1144
1145 #[test]
1146 fn resolve_security_prefers_most_specific_but_clamps_taint() {
1147 assert!(resolve_security(true, None, None).taint_tracking);
1149 assert!(!resolve_security(false, None, None).taint_tracking);
1150 assert!(resolve_security(false, Some(&sec(false)), Some(&sec(true))).taint_tracking);
1152 assert!(resolve_security(true, Some(&sec(false)), None).taint_tracking);
1155 }
1156
1157 #[test]
1158 fn test_region_taint_remove_at_recomputes_level() {
1159 let mut rt = RegionTaint::new();
1160 rt.add_entry(TaintLevel::Public);
1161 rt.add_entry(TaintLevel::Private);
1162 rt.add_entry(TaintLevel::Public);
1163 assert_eq!(rt.level(), TaintLevel::Private);
1164
1165 rt.remove_at(1);
1167 assert_eq!(rt.entry_count(), 2);
1168 assert_eq!(rt.level(), TaintLevel::Public);
1169
1170 rt.remove_at(99);
1172 assert_eq!(rt.entry_count(), 2);
1173 }
1174}