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