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 classified_builtin(tool_name).unwrap_or_else(|| {
477 ToolClassification::new(
478 TaintLevel::Public,
479 ToolDirection::Outbound,
480 TaintLevel::Public,
481 )
482 })
483}
484
485pub fn classified_builtin(tool_name: &str) -> Option<ToolClassification> {
493 let classification = match tool_name {
494 "read_file" | "read_files" => ToolClassification::new(
496 TaintLevel::Internal,
497 ToolDirection::Inbound,
498 TaintLevel::Public,
499 ),
500 "write_file" | "install_self_tool" | "install_global_tool" | "install_tool" => {
505 ToolClassification::new(
506 TaintLevel::Internal,
507 ToolDirection::Internal,
508 TaintLevel::Public,
509 )
510 }
511 "edit_file" | "edit_document" => ToolClassification::new(
514 TaintLevel::Internal,
515 ToolDirection::Internal,
516 TaintLevel::Public,
517 ),
518 "context_write" | "context_append" | "context_read" | "context_delete" | "context_list"
522 | "context_attach" | "context_export" | "todo_add" | "todo_done" | "todo_note" => {
523 ToolClassification::new(
524 TaintLevel::Internal,
525 ToolDirection::Internal,
526 TaintLevel::Public,
527 )
528 }
529 "submit_output" => ToolClassification::new(
539 TaintLevel::Public,
540 ToolDirection::Outbound,
541 TaintLevel::Public,
542 ),
543 "list_dir" => ToolClassification::new(
544 TaintLevel::Internal,
545 ToolDirection::Inbound,
546 TaintLevel::Public,
547 ),
548 "shell" | "bash" => ToolClassification::new(
549 TaintLevel::Public,
550 ToolDirection::Outbound,
551 TaintLevel::Public,
552 ),
553 "web_search" | "web_fetch" | "http_get" | "http_post" | "fetch" => ToolClassification::new(
557 TaintLevel::Public,
558 ToolDirection::Outbound,
559 TaintLevel::Public,
560 ),
561 "current_time" | "locale_info" => ToolClassification::new(
570 TaintLevel::Public,
571 ToolDirection::Inbound,
572 TaintLevel::Public,
573 ),
574 "system_info" | "environment_info" | "which_command" | "runtime_info" => {
575 ToolClassification::new(
576 TaintLevel::Internal,
577 ToolDirection::Inbound,
578 TaintLevel::Public,
579 )
580 }
581 "ask_user_text" | "ask_user_choice" | "ask_user_confirm" | "present_for_review" => {
582 ToolClassification::new(
583 TaintLevel::Internal,
584 ToolDirection::Internal,
585 TaintLevel::Public,
586 )
587 }
588 "spawn_agent" | "check_agent" | "wait_for_agent" | "send_to_agent" | "kill_agent"
590 | "fan_out" => ToolClassification::new(
591 TaintLevel::Internal,
592 ToolDirection::Internal,
593 TaintLevel::Public,
594 ),
595 _ => return None,
596 };
597 Some(classification)
598}
599
600#[cfg(test)]
601mod tests {
602 use super::*;
603
604 #[test]
608 fn taint_rebuilds_from_persisted_entries_at_the_highest_level() {
609 let restored = RegionTaint::from_entry_taints(vec![
610 TaintLevel::Public,
611 TaintLevel::Private,
612 TaintLevel::Internal,
613 ]);
614 assert_eq!(restored.level(), TaintLevel::Private);
615 assert_eq!(restored.entry_taint(1), Some(TaintLevel::Private));
616
617 assert_eq!(
620 RegionTaint::from_entry_taints(Vec::new()).level(),
621 TaintLevel::Public
622 );
623 }
624
625 #[test]
628 fn taint_level_ordering() {
629 assert!(TaintLevel::Public < TaintLevel::Internal);
630 assert!(TaintLevel::Internal < TaintLevel::Private);
631 assert!(TaintLevel::Public < TaintLevel::Private);
632 }
633
634 #[test]
635 fn taint_level_equality() {
636 assert_eq!(TaintLevel::Public, TaintLevel::Public);
637 assert_eq!(TaintLevel::Internal, TaintLevel::Internal);
638 assert_eq!(TaintLevel::Private, TaintLevel::Private);
639 assert_ne!(TaintLevel::Public, TaintLevel::Private);
640 }
641
642 #[test]
643 fn taint_level_max() {
644 assert_eq!(
645 TaintLevel::Public.max(TaintLevel::Internal),
646 TaintLevel::Internal
647 );
648 assert_eq!(
649 TaintLevel::Private.max(TaintLevel::Public),
650 TaintLevel::Private
651 );
652 assert_eq!(
653 TaintLevel::Internal.max(TaintLevel::Internal),
654 TaintLevel::Internal
655 );
656 }
657
658 #[test]
659 fn taint_level_default_is_internal() {
660 assert_eq!(TaintLevel::default(), TaintLevel::Internal);
661 }
662
663 #[test]
664 fn taint_level_display() {
665 assert_eq!(format!("{}", TaintLevel::Public), "public");
666 assert_eq!(format!("{}", TaintLevel::Internal), "internal");
667 assert_eq!(format!("{}", TaintLevel::Private), "private");
668 }
669
670 #[test]
671 fn taint_level_from_str_loose() {
672 assert_eq!(
673 TaintLevel::from_str_loose("public"),
674 Some(TaintLevel::Public)
675 );
676 assert_eq!(
677 TaintLevel::from_str_loose("INTERNAL"),
678 Some(TaintLevel::Internal)
679 );
680 assert_eq!(
681 TaintLevel::from_str_loose("Private"),
682 Some(TaintLevel::Private)
683 );
684 assert_eq!(TaintLevel::from_str_loose("unknown"), None);
685 }
686
687 #[test]
688 fn taint_level_as_str() {
689 assert_eq!(TaintLevel::Public.as_str(), "public");
690 assert_eq!(TaintLevel::Internal.as_str(), "internal");
691 assert_eq!(TaintLevel::Private.as_str(), "private");
692 }
693
694 #[test]
695 fn taint_level_serde_roundtrip() {
696 for level in [
697 TaintLevel::Public,
698 TaintLevel::Internal,
699 TaintLevel::Private,
700 ] {
701 let json = serde_json::to_string(&level).unwrap();
702 let back: TaintLevel = serde_json::from_str(&json).unwrap();
703 assert_eq!(level, back);
704 }
705 }
706
707 #[test]
708 fn taint_level_hash() {
709 use std::collections::HashSet;
710 let mut set = HashSet::new();
711 set.insert(TaintLevel::Public);
712 set.insert(TaintLevel::Internal);
713 set.insert(TaintLevel::Private);
714 set.insert(TaintLevel::Public); assert_eq!(set.len(), 3);
716 }
717
718 #[test]
721 fn tool_direction_from_str_loose() {
722 assert_eq!(
723 ToolDirection::from_str_loose("inbound"),
724 Some(ToolDirection::Inbound)
725 );
726 assert_eq!(
727 ToolDirection::from_str_loose("OUTBOUND"),
728 Some(ToolDirection::Outbound)
729 );
730 assert_eq!(
731 ToolDirection::from_str_loose("Internal"),
732 Some(ToolDirection::Internal)
733 );
734 assert_eq!(ToolDirection::from_str_loose("nope"), None);
735 }
736
737 #[test]
738 fn tool_direction_default_is_internal() {
739 assert_eq!(ToolDirection::default(), ToolDirection::Internal);
740 }
741
742 #[test]
743 fn tool_direction_display() {
744 assert_eq!(format!("{}", ToolDirection::Inbound), "inbound");
745 assert_eq!(format!("{}", ToolDirection::Internal), "internal");
746 assert_eq!(format!("{}", ToolDirection::Outbound), "outbound");
747 }
748
749 #[test]
750 fn tool_direction_serde_roundtrip() {
751 for dir in [
752 ToolDirection::Inbound,
753 ToolDirection::Internal,
754 ToolDirection::Outbound,
755 ] {
756 let json = serde_json::to_string(&dir).unwrap();
757 let back: ToolDirection = serde_json::from_str(&json).unwrap();
758 assert_eq!(dir, back);
759 }
760 }
761
762 #[test]
765 fn tool_classification_default() {
766 let tc = ToolClassification::default();
767 assert_eq!(tc.sensitivity, TaintLevel::Internal);
768 assert_eq!(tc.direction, ToolDirection::Internal);
769 assert_eq!(tc.clearance, TaintLevel::Public);
770 }
771
772 #[test]
773 fn tool_classification_outbound_check() {
774 let tc = ToolClassification::new(
775 TaintLevel::Public,
776 ToolDirection::Outbound,
777 TaintLevel::Internal,
778 );
779 assert!(tc.is_outbound());
780 assert!(tc.check_clearance(TaintLevel::Public));
781 assert!(tc.check_clearance(TaintLevel::Internal));
782 assert!(!tc.check_clearance(TaintLevel::Private));
783 }
784
785 #[test]
786 fn tool_classification_non_outbound_always_passes() {
787 let tc = ToolClassification::new(
788 TaintLevel::Private,
789 ToolDirection::Inbound,
790 TaintLevel::Public, );
792 assert!(!tc.is_outbound());
793 assert!(tc.check_clearance(TaintLevel::Private));
794 }
795
796 #[test]
797 fn tool_classification_serde_roundtrip() {
798 let tc = ToolClassification::new(
799 TaintLevel::Private,
800 ToolDirection::Outbound,
801 TaintLevel::Internal,
802 );
803 let json = serde_json::to_string(&tc).unwrap();
804 let back: ToolClassification = serde_json::from_str(&json).unwrap();
805 assert_eq!(tc, back);
806 }
807
808 #[test]
811 fn region_taint_starts_public() {
812 let rt = RegionTaint::new();
813 assert_eq!(rt.level(), TaintLevel::Public);
814 assert_eq!(rt.entry_count(), 0);
815 }
816
817 #[test]
818 fn region_taint_add_entry_raises_level() {
819 let mut rt = RegionTaint::new();
820 rt.add_entry(TaintLevel::Internal);
821 assert_eq!(rt.level(), TaintLevel::Internal);
822 rt.add_entry(TaintLevel::Private);
823 assert_eq!(rt.level(), TaintLevel::Private);
824 }
825
826 #[test]
827 fn region_taint_add_public_doesnt_lower() {
828 let mut rt = RegionTaint::new();
829 rt.add_entry(TaintLevel::Private);
830 rt.add_entry(TaintLevel::Public);
831 assert_eq!(rt.level(), TaintLevel::Private);
832 }
833
834 #[test]
835 fn region_taint_remove_oldest_recovers() {
836 let mut rt = RegionTaint::new();
837 rt.add_entry(TaintLevel::Private);
838 rt.add_entry(TaintLevel::Public);
839 assert_eq!(rt.level(), TaintLevel::Private);
840
841 rt.remove_oldest(); assert_eq!(rt.level(), TaintLevel::Public);
843 }
844
845 #[test]
846 fn region_taint_remove_oldest_empty() {
847 let mut rt = RegionTaint::new();
848 rt.remove_oldest(); assert_eq!(rt.level(), TaintLevel::Public);
850 }
851
852 #[test]
853 fn region_taint_clear() {
854 let mut rt = RegionTaint::new();
855 rt.add_entry(TaintLevel::Private);
856 rt.add_entry(TaintLevel::Internal);
857 rt.clear();
858 assert_eq!(rt.level(), TaintLevel::Public);
859 assert_eq!(rt.entry_count(), 0);
860 }
861
862 #[test]
863 fn region_taint_recompute() {
864 let mut rt = RegionTaint::new();
865 rt.add_entry(TaintLevel::Private);
866 rt.add_entry(TaintLevel::Internal);
867 rt.add_entry(TaintLevel::Public);
868 assert_eq!(rt.entry_count(), 3);
869
870 rt.remove_oldest();
872 assert_eq!(rt.level(), TaintLevel::Internal);
873 assert_eq!(rt.entry_count(), 2);
874 }
875
876 #[test]
877 fn region_taint_entry_taint() {
878 let mut rt = RegionTaint::new();
879 rt.add_entry(TaintLevel::Public);
880 rt.add_entry(TaintLevel::Private);
881 assert_eq!(rt.entry_taint(0), Some(TaintLevel::Public));
882 assert_eq!(rt.entry_taint(1), Some(TaintLevel::Private));
883 assert_eq!(rt.entry_taint(2), None);
884 }
885
886 #[test]
887 fn region_taint_default() {
888 let rt = RegionTaint::default();
889 assert_eq!(rt.level(), TaintLevel::Public);
890 }
891
892 #[test]
893 fn region_taint_serde_roundtrip() {
894 let mut rt = RegionTaint::new();
895 rt.add_entry(TaintLevel::Internal);
896 rt.add_entry(TaintLevel::Private);
897 let json = serde_json::to_string(&rt).unwrap();
898 let back: RegionTaint = serde_json::from_str(&json).unwrap();
899 assert_eq!(back.level(), TaintLevel::Private);
900 assert_eq!(back.entry_count(), 2);
901 }
902
903 #[test]
906 fn security_config_default() {
907 let sc = SecurityConfig::default();
908 assert!(sc.taint_tracking);
909 }
910
911 #[test]
912 fn security_config_serde_roundtrip() {
913 let sc = SecurityConfig {
914 taint_tracking: false,
915 };
916 let json = serde_json::to_string(&sc).unwrap();
917 let back: SecurityConfig = serde_json::from_str(&json).unwrap();
918 assert!(!back.taint_tracking);
919 }
920
921 #[test]
924 fn gate_decision_allowed() {
925 let d = GateDecision::Allowed;
926 assert!(d.is_allowed());
927 }
928
929 #[test]
930 fn gate_decision_blocked() {
931 let d = GateDecision::Blocked {
932 taint_level: TaintLevel::Private,
933 clearance: TaintLevel::Public,
934 source_regions: vec!["conversation".into()],
935 tool_name: "send_email".into(),
936 };
937 assert!(!d.is_allowed());
938 }
939
940 #[test]
943 fn gate_event_serde_roundtrip() {
944 let event = GateEvent {
945 timestamp: 1234567890,
946 agent_id: "agent-1".into(),
947 tool_name: "send_email".into(),
948 taint_level: TaintLevel::Private,
949 clearance: TaintLevel::Public,
950 allowed: false,
951 decision_source: GateDecisionSource::UserDenied,
952 };
953 let json = serde_json::to_string(&event).unwrap();
954 let back: GateEvent = serde_json::from_str(&json).unwrap();
955 assert_eq!(back.agent_id, "agent-1");
956 assert!(!back.allowed);
957 }
958
959 #[test]
960 fn gate_decision_source_variants() {
961 let sources = vec![
962 GateDecisionSource::AutoAllow,
963 GateDecisionSource::AllowlistRule { rule_index: 0 },
964 GateDecisionSource::ScriptedRule {
965 script_name: "test.rhai".into(),
966 },
967 GateDecisionSource::UserAllowOnce,
968 GateDecisionSource::UserAlwaysAllow,
969 GateDecisionSource::UserDenied,
970 GateDecisionSource::TaintDisabled,
971 ];
972 for src in sources {
973 let json = serde_json::to_string(&src).unwrap();
974 let back: GateDecisionSource = serde_json::from_str(&json).unwrap();
975 assert_eq!(src, back);
976 }
977 }
978
979 #[test]
982 fn builtin_read_file_classification() {
983 let tc = builtin_tool_classification("read_file");
984 assert_eq!(tc.sensitivity, TaintLevel::Internal);
985 assert_eq!(tc.direction, ToolDirection::Inbound);
986 }
987
988 #[test]
989 fn builtin_shell_classification() {
990 let tc = builtin_tool_classification("shell");
991 assert_eq!(tc.sensitivity, TaintLevel::Public);
992 assert_eq!(tc.direction, ToolDirection::Outbound);
993 assert_eq!(tc.clearance, TaintLevel::Public);
994
995 let tc2 = builtin_tool_classification("bash");
997 assert_eq!(tc2.direction, ToolDirection::Outbound);
998 }
999
1000 #[test]
1007 fn network_capable_tools_are_outbound() {
1008 for name in ["web_search", "web_fetch", "http_get", "http_post", "fetch"] {
1009 let tc = builtin_tool_classification(name);
1010 assert_eq!(tc.sensitivity, TaintLevel::Public, "{name}");
1011 assert_eq!(tc.direction, ToolDirection::Outbound, "{name}");
1012 }
1013 }
1014
1015 #[test]
1021 fn environment_tools_are_inbound_and_never_gated() {
1022 for name in [
1023 "current_time",
1024 "system_info",
1025 "locale_info",
1026 "environment_info",
1027 "which_command",
1028 "runtime_info",
1029 ] {
1030 let tc = builtin_tool_classification(name);
1031 assert_eq!(tc.direction, ToolDirection::Inbound, "{name}");
1032 assert_eq!(tc.clearance, TaintLevel::Public, "{name}");
1034 }
1035 }
1036
1037 #[test]
1042 fn environment_tools_are_graded_by_what_they_reveal() {
1043 for name in ["current_time", "locale_info"] {
1044 assert_eq!(
1045 builtin_tool_classification(name).sensitivity,
1046 TaintLevel::Public,
1047 "{name}"
1048 );
1049 }
1050 for name in [
1051 "system_info",
1052 "environment_info",
1053 "which_command",
1054 "runtime_info",
1055 ] {
1056 assert_eq!(
1057 builtin_tool_classification(name).sensitivity,
1058 TaintLevel::Internal,
1059 "{name}"
1060 );
1061 }
1062 }
1063
1064 #[test]
1065 fn builtin_ask_user_classification() {
1066 for name in [
1067 "ask_user_text",
1068 "ask_user_choice",
1069 "ask_user_confirm",
1070 "present_for_review",
1071 ] {
1072 let tc = builtin_tool_classification(name);
1073 assert_eq!(tc.direction, ToolDirection::Internal);
1074 }
1075 }
1076
1077 #[test]
1078 fn builtin_subagent_classification() {
1079 for name in [
1080 "spawn_agent",
1081 "check_agent",
1082 "wait_for_agent",
1083 "send_to_agent",
1084 "kill_agent",
1085 ] {
1086 let tc = builtin_tool_classification(name);
1087 assert_eq!(tc.direction, ToolDirection::Internal);
1088 }
1089 }
1090
1091 #[test]
1092 fn builtin_write_file_classification() {
1093 let tc = builtin_tool_classification("write_file");
1094 assert_eq!(tc.direction, ToolDirection::Internal);
1095 }
1096
1097 #[test]
1101 fn install_tool_is_classified_like_write_file() {
1102 assert_eq!(
1103 classified_builtin("install_tool"),
1104 classified_builtin("write_file")
1105 );
1106 let tc = builtin_tool_classification("install_tool");
1107 assert_eq!(tc.sensitivity, TaintLevel::Internal);
1108 assert_eq!(tc.direction, ToolDirection::Internal);
1109 assert_eq!(tc.clearance, TaintLevel::Public);
1110 }
1111
1112 #[test]
1117 fn unknown_tools_fail_closed_as_outbound() {
1118 let tc = builtin_tool_classification("some_mcp_tool");
1119 assert_eq!(tc.sensitivity, TaintLevel::Public);
1120 assert_eq!(tc.direction, ToolDirection::Outbound);
1121 assert_eq!(tc.clearance, TaintLevel::Public);
1122 }
1123
1124 #[test]
1125 fn builtin_edit_file_classification() {
1126 let tc = builtin_tool_classification("edit_file");
1127 assert_eq!(tc.sensitivity, TaintLevel::Internal);
1128 assert_eq!(tc.direction, ToolDirection::Internal);
1129 assert_eq!(tc.clearance, TaintLevel::Public);
1130 }
1131
1132 #[test]
1133 fn builtin_list_dir_classification() {
1134 let tc = builtin_tool_classification("list_dir");
1135 assert_eq!(tc.sensitivity, TaintLevel::Internal);
1136 assert_eq!(tc.direction, ToolDirection::Inbound);
1137 assert_eq!(tc.clearance, TaintLevel::Public);
1138 }
1139
1140 #[test]
1148 fn the_remaining_builtins_are_classified_like_their_siblings() {
1149 assert_eq!(
1150 classified_builtin("read_files"),
1151 classified_builtin("read_file"),
1152 "read_files"
1153 );
1154 assert_eq!(
1155 classified_builtin("fan_out"),
1156 classified_builtin("spawn_agent"),
1157 "fan_out"
1158 );
1159 assert_eq!(
1160 classified_builtin("edit_document"),
1161 classified_builtin("edit_file"),
1162 "edit_document"
1163 );
1164 for name in [
1165 "context_write",
1166 "context_append",
1167 "context_read",
1168 "context_delete",
1169 "context_list",
1170 "context_attach",
1171 "context_export",
1172 "todo_add",
1173 "todo_done",
1174 "todo_note",
1175 ] {
1176 let tc = classified_builtin(name);
1177 assert!(tc.is_some(), "{name} has no arm");
1178 let tc = tc.unwrap();
1179 assert_eq!(tc.direction, ToolDirection::Internal, "{name}");
1180 assert_eq!(tc.sensitivity, TaintLevel::Internal, "{name}");
1181 assert_eq!(tc.clearance, TaintLevel::Public, "{name}");
1182 }
1183 }
1184
1185 #[test]
1190 fn submit_output_is_an_outbound_channel() {
1191 assert_eq!(
1192 classified_builtin("submit_output"),
1193 classified_builtin("shell"),
1194 "submit_output"
1195 );
1196 }
1197
1198 #[test]
1200 fn a_third_party_name_has_no_arm_of_its_own() {
1201 assert_eq!(classified_builtin("some_mcp_tool"), None);
1202 assert_eq!(
1203 classified_builtin("shell"),
1204 Some(builtin_tool_classification("shell"))
1205 );
1206 }
1207
1208 fn sec(taint: bool) -> SecurityConfig {
1211 SecurityConfig {
1212 taint_tracking: taint,
1213 }
1214 }
1215
1216 #[test]
1217 fn resolve_taint_enabled_inherits_global_when_unset() {
1218 assert!(!resolve_taint_enabled(false, None, None));
1219 assert!(resolve_taint_enabled(true, None, None));
1220 }
1221
1222 #[test]
1223 fn resolve_taint_enabled_agent_may_opt_in_but_not_out() {
1224 assert!(resolve_taint_enabled(false, Some(&sec(true)), None));
1226 assert!(resolve_taint_enabled(true, Some(&sec(false)), None));
1230 }
1231
1232 #[test]
1233 fn resolve_taint_enabled_stage_may_opt_in_but_not_out() {
1234 assert!(resolve_taint_enabled(
1236 false,
1237 Some(&sec(false)),
1238 Some(&sec(true))
1239 ));
1240 assert!(resolve_taint_enabled(
1242 true,
1243 Some(&sec(true)),
1244 Some(&sec(false))
1245 ));
1246 }
1247
1248 #[test]
1249 fn resolve_batch_tool_hint_cascade() {
1250 assert!(resolve_batch_tool_hint(true, None, None));
1252 assert!(!resolve_batch_tool_hint(false, None, None));
1253 assert!(!resolve_batch_tool_hint(true, Some(false), None));
1255 assert!(resolve_batch_tool_hint(false, Some(true), None));
1256 assert!(!resolve_batch_tool_hint(true, Some(true), Some(false)));
1258 assert!(resolve_batch_tool_hint(false, Some(false), Some(true)));
1259 }
1260
1261 #[test]
1262 fn gate_decision_blocked_levels() {
1263 let blocked = GateDecision::Blocked {
1264 taint_level: TaintLevel::Private,
1265 clearance: TaintLevel::Public,
1266 source_regions: vec![],
1267 tool_name: "shell".into(),
1268 };
1269 assert_eq!(
1270 blocked.blocked_levels(),
1271 Some((TaintLevel::Private, TaintLevel::Public))
1272 );
1273 assert_eq!(GateDecision::Allowed.blocked_levels(), None);
1274 }
1275
1276 #[test]
1277 fn resolve_security_prefers_most_specific_but_clamps_taint() {
1278 assert!(resolve_security(true, None, None).taint_tracking);
1280 assert!(!resolve_security(false, None, None).taint_tracking);
1281 assert!(resolve_security(false, Some(&sec(false)), Some(&sec(true))).taint_tracking);
1283 assert!(resolve_security(true, Some(&sec(false)), None).taint_tracking);
1286 }
1287
1288 #[test]
1289 fn test_region_taint_remove_at_recomputes_level() {
1290 let mut rt = RegionTaint::new();
1291 rt.add_entry(TaintLevel::Public);
1292 rt.add_entry(TaintLevel::Private);
1293 rt.add_entry(TaintLevel::Public);
1294 assert_eq!(rt.level(), TaintLevel::Private);
1295
1296 rt.remove_at(1);
1298 assert_eq!(rt.entry_count(), 2);
1299 assert_eq!(rt.level(), TaintLevel::Public);
1300
1301 rt.remove_at(99);
1303 assert_eq!(rt.entry_count(), 2);
1304 }
1305}