Skip to main content

a3s_code_core/security/
mod.rs

1//! Security Module
2//!
3//! Provides a trait-based security interface for A3S Code sessions.
4//! External consumers implement `SecurityProvider` to plug in their own
5//! security logic (sanitization, taint tracking, injection detection, etc.).
6
7pub mod config;
8pub mod default;
9mod event_sanitizer;
10
11pub use config::{RedactionStrategy, SecurityConfig, SensitivityLevel};
12pub use default::{DefaultSecurityConfig, DefaultSecurityProvider, SensitivePattern};
13pub(crate) use event_sanitizer::AgentEventStreamSanitizer;
14
15use crate::hooks::HookEngine;
16
17/// Sanitize every data-bearing field of an agent event while preserving the
18/// identifiers and discriminants used to correlate the event stream.
19///
20/// Runtime boundaries call this immediately before events are observed,
21/// persisted, or exposed to an SDK. Implementations of [`SecurityProvider`]
22/// should therefore make [`SecurityProvider::sanitize_output`] idempotent.
23pub fn sanitize_agent_event(
24    provider: &dyn SecurityProvider,
25    event: &crate::agent::AgentEvent,
26) -> crate::agent::AgentEvent {
27    use crate::agent::AgentEvent;
28
29    fn text(provider: &dyn SecurityProvider, value: &str) -> String {
30        provider.sanitize_output(value)
31    }
32
33    fn optional_text(provider: &dyn SecurityProvider, value: &Option<String>) -> Option<String> {
34        value.as_deref().map(|value| text(provider, value))
35    }
36
37    fn strings(provider: &dyn SecurityProvider, values: &[String]) -> Vec<String> {
38        values.iter().map(|value| text(provider, value)).collect()
39    }
40
41    fn json(provider: &dyn SecurityProvider, value: &serde_json::Value) -> serde_json::Value {
42        match value {
43            serde_json::Value::String(value) => serde_json::Value::String(text(provider, value)),
44            serde_json::Value::Array(values) => {
45                serde_json::Value::Array(values.iter().map(|value| json(provider, value)).collect())
46            }
47            serde_json::Value::Object(values) => serde_json::Value::Object(
48                values
49                    .iter()
50                    .map(|(key, value)| (key.clone(), json(provider, value)))
51                    .collect(),
52            ),
53            value => value.clone(),
54        }
55    }
56
57    fn optional_json(
58        provider: &dyn SecurityProvider,
59        value: &Option<serde_json::Value>,
60    ) -> Option<serde_json::Value> {
61        value.as_ref().map(|value| json(provider, value))
62    }
63
64    fn tool_error_kind(
65        provider: &dyn SecurityProvider,
66        value: &Option<crate::tools::ToolErrorKind>,
67    ) -> Option<crate::tools::ToolErrorKind> {
68        use crate::tools::ToolErrorKind;
69
70        value.as_ref().map(|kind| match kind {
71            ToolErrorKind::VersionConflict {
72                path,
73                expected,
74                actual,
75            } => ToolErrorKind::VersionConflict {
76                path: text(provider, path),
77                expected: text(provider, expected),
78                actual: optional_text(provider, actual),
79            },
80            ToolErrorKind::RemoteGitConflict { code, message } => {
81                ToolErrorKind::RemoteGitConflict {
82                    code: text(provider, code),
83                    message: text(provider, message),
84                }
85            }
86            ToolErrorKind::NotFound { path } => ToolErrorKind::NotFound {
87                path: text(provider, path),
88            },
89            ToolErrorKind::InvalidArgument { message } => ToolErrorKind::InvalidArgument {
90                message: text(provider, message),
91            },
92            ToolErrorKind::Unsupported { message } => ToolErrorKind::Unsupported {
93                message: text(provider, message),
94            },
95            ToolErrorKind::Timeout { op, duration_ms } => ToolErrorKind::Timeout {
96                op: text(provider, op),
97                duration_ms: *duration_ms,
98            },
99            ToolErrorKind::Transport { op } => ToolErrorKind::Transport {
100                op: text(provider, op),
101            },
102            ToolErrorKind::Cancelled { op } => ToolErrorKind::Cancelled {
103                op: text(provider, op),
104            },
105            ToolErrorKind::PartialFailure { failed, total } => ToolErrorKind::PartialFailure {
106                failed: *failed,
107                total: *total,
108            },
109            ToolErrorKind::RateLimited { retry_after_ms } => ToolErrorKind::RateLimited {
110                retry_after_ms: *retry_after_ms,
111            },
112        })
113    }
114
115    fn task(
116        provider: &dyn SecurityProvider,
117        task: &crate::planning::Task,
118    ) -> crate::planning::Task {
119        let mut task = task.clone();
120        task.content = text(provider, &task.content);
121        task.success_criteria = optional_text(provider, &task.success_criteria);
122        task
123    }
124
125    fn plan(
126        provider: &dyn SecurityProvider,
127        plan: &crate::planning::ExecutionPlan,
128    ) -> crate::planning::ExecutionPlan {
129        let mut plan = plan.clone();
130        plan.goal = text(provider, &plan.goal);
131        plan.steps = plan.steps.iter().map(|item| task(provider, item)).collect();
132        plan
133    }
134
135    fn goal(
136        provider: &dyn SecurityProvider,
137        goal: &crate::planning::AgentGoal,
138    ) -> crate::planning::AgentGoal {
139        let mut goal = goal.clone();
140        goal.description = text(provider, &goal.description);
141        goal.success_criteria = strings(provider, &goal.success_criteria);
142        goal
143    }
144
145    fn verification_summary(
146        provider: &dyn SecurityProvider,
147        summary: &crate::verification::VerificationSummary,
148    ) -> crate::verification::VerificationSummary {
149        let mut summary = summary.clone();
150        summary.pending_subjects = strings(provider, &summary.pending_subjects);
151        summary.failed_subjects = strings(provider, &summary.failed_subjects);
152        summary
153    }
154
155    fn response_meta(
156        provider: &dyn SecurityProvider,
157        meta: &Option<crate::llm::LlmResponseMeta>,
158    ) -> Option<crate::llm::LlmResponseMeta> {
159        meta.as_ref().map(|meta| {
160            let mut meta = meta.clone();
161            meta.request_url = optional_text(provider, &meta.request_url);
162            meta
163        })
164    }
165
166    match event {
167        AgentEvent::Start { prompt } => AgentEvent::Start {
168            prompt: text(provider, prompt),
169        },
170        AgentEvent::AgentModeChanged {
171            mode,
172            agent,
173            description,
174        } => AgentEvent::AgentModeChanged {
175            mode: mode.clone(),
176            agent: agent.clone(),
177            description: text(provider, description),
178        },
179        AgentEvent::TextDelta { text: value } => AgentEvent::TextDelta {
180            text: text(provider, value),
181        },
182        AgentEvent::ReasoningDelta { text: value } => AgentEvent::ReasoningDelta {
183            text: text(provider, value),
184        },
185        AgentEvent::ToolInputDelta { id, delta } => AgentEvent::ToolInputDelta {
186            id: id.clone(),
187            delta: text(provider, delta),
188        },
189        AgentEvent::ToolExecutionStart { id, name, args } => AgentEvent::ToolExecutionStart {
190            id: id.clone(),
191            name: name.clone(),
192            args: json(provider, args),
193        },
194        AgentEvent::ToolEnd {
195            id,
196            name,
197            args,
198            output,
199            exit_code,
200            metadata,
201            error_kind,
202        } => AgentEvent::ToolEnd {
203            id: id.clone(),
204            name: name.clone(),
205            args: optional_json(provider, args),
206            output: text(provider, output),
207            exit_code: *exit_code,
208            metadata: optional_json(provider, metadata),
209            error_kind: tool_error_kind(provider, error_kind),
210        },
211        AgentEvent::ToolOutputDelta { id, name, delta } => AgentEvent::ToolOutputDelta {
212            id: id.clone(),
213            name: name.clone(),
214            delta: text(provider, delta),
215        },
216        AgentEvent::End {
217            text: value,
218            usage,
219            verification_summary: summary,
220            meta,
221        } => AgentEvent::End {
222            text: text(provider, value),
223            usage: usage.clone(),
224            verification_summary: Box::new(verification_summary(provider, summary)),
225            meta: response_meta(provider, meta),
226        },
227        AgentEvent::Error { message } => AgentEvent::Error {
228            message: text(provider, message),
229        },
230        AgentEvent::ConfirmationRequired {
231            tool_id,
232            tool_name,
233            args,
234            timeout_ms,
235        } => AgentEvent::ConfirmationRequired {
236            tool_id: tool_id.clone(),
237            tool_name: tool_name.clone(),
238            args: json(provider, args),
239            timeout_ms: *timeout_ms,
240        },
241        AgentEvent::ConfirmationReceived {
242            tool_id,
243            approved,
244            reason,
245        } => AgentEvent::ConfirmationReceived {
246            tool_id: tool_id.clone(),
247            approved: *approved,
248            reason: optional_text(provider, reason),
249        },
250        AgentEvent::ConfirmationTimeout {
251            tool_id,
252            action_taken,
253        } => AgentEvent::ConfirmationTimeout {
254            tool_id: tool_id.clone(),
255            action_taken: action_taken.clone(),
256        },
257        AgentEvent::ExternalTaskPending {
258            task_id,
259            session_id,
260            lane,
261            command_type,
262            payload,
263            timeout_ms,
264        } => AgentEvent::ExternalTaskPending {
265            task_id: task_id.clone(),
266            session_id: session_id.clone(),
267            lane: *lane,
268            command_type: command_type.clone(),
269            payload: json(provider, payload),
270            timeout_ms: *timeout_ms,
271        },
272        AgentEvent::PermissionDenied {
273            tool_id,
274            tool_name,
275            args,
276            reason,
277        } => AgentEvent::PermissionDenied {
278            tool_id: tool_id.clone(),
279            tool_name: tool_name.clone(),
280            args: json(provider, args),
281            reason: text(provider, reason),
282        },
283        AgentEvent::CommandDeadLettered {
284            command_id,
285            command_type,
286            lane,
287            error,
288            attempts,
289        } => AgentEvent::CommandDeadLettered {
290            command_id: command_id.clone(),
291            command_type: command_type.clone(),
292            lane: lane.clone(),
293            error: text(provider, error),
294            attempts: *attempts,
295        },
296        AgentEvent::QueueAlert {
297            level,
298            alert_type,
299            message,
300        } => AgentEvent::QueueAlert {
301            level: level.clone(),
302            alert_type: alert_type.clone(),
303            message: text(provider, message),
304        },
305        AgentEvent::TaskUpdated { session_id, tasks } => AgentEvent::TaskUpdated {
306            session_id: session_id.clone(),
307            tasks: tasks.iter().map(|item| task(provider, item)).collect(),
308        },
309        AgentEvent::MemoryStored {
310            memory_id,
311            memory_type,
312            importance,
313            tags,
314        } => AgentEvent::MemoryStored {
315            memory_id: memory_id.clone(),
316            memory_type: memory_type.clone(),
317            importance: *importance,
318            tags: strings(provider, tags),
319        },
320        AgentEvent::MemoryRecalled {
321            memory_id,
322            content,
323            relevance,
324        } => AgentEvent::MemoryRecalled {
325            memory_id: memory_id.clone(),
326            content: text(provider, content),
327            relevance: *relevance,
328        },
329        AgentEvent::MemoriesSearched {
330            query,
331            tags,
332            result_count,
333        } => AgentEvent::MemoriesSearched {
334            query: optional_text(provider, query),
335            tags: strings(provider, tags),
336            result_count: *result_count,
337        },
338        AgentEvent::SubagentStart {
339            task_id,
340            session_id,
341            parent_session_id,
342            agent,
343            description,
344            started_ms,
345        } => AgentEvent::SubagentStart {
346            task_id: task_id.clone(),
347            session_id: session_id.clone(),
348            parent_session_id: parent_session_id.clone(),
349            agent: agent.clone(),
350            description: text(provider, description),
351            started_ms: *started_ms,
352        },
353        AgentEvent::SubagentProgress {
354            task_id,
355            session_id,
356            status,
357            metadata,
358        } => AgentEvent::SubagentProgress {
359            task_id: task_id.clone(),
360            session_id: session_id.clone(),
361            status: text(provider, status),
362            metadata: json(provider, metadata),
363        },
364        AgentEvent::SubagentEnd {
365            task_id,
366            session_id,
367            agent,
368            output,
369            success,
370            finished_ms,
371        } => AgentEvent::SubagentEnd {
372            task_id: task_id.clone(),
373            session_id: session_id.clone(),
374            agent: agent.clone(),
375            output: text(provider, output),
376            success: *success,
377            finished_ms: *finished_ms,
378        },
379        AgentEvent::PlanningStart { prompt } => AgentEvent::PlanningStart {
380            prompt: text(provider, prompt),
381        },
382        AgentEvent::PlanningEnd {
383            plan: value,
384            estimated_steps,
385        } => AgentEvent::PlanningEnd {
386            plan: plan(provider, value),
387            estimated_steps: *estimated_steps,
388        },
389        AgentEvent::StepStart {
390            step_id,
391            description,
392            step_number,
393            total_steps,
394        } => AgentEvent::StepStart {
395            step_id: step_id.clone(),
396            description: text(provider, description),
397            step_number: *step_number,
398            total_steps: *total_steps,
399        },
400        AgentEvent::GoalExtracted { goal: value } => AgentEvent::GoalExtracted {
401            goal: goal(provider, value),
402        },
403        AgentEvent::GoalProgress {
404            goal,
405            progress,
406            completed_steps,
407            total_steps,
408        } => AgentEvent::GoalProgress {
409            goal: text(provider, goal),
410            progress: *progress,
411            completed_steps: *completed_steps,
412            total_steps: *total_steps,
413        },
414        AgentEvent::GoalAchieved {
415            goal,
416            total_steps,
417            duration_ms,
418        } => AgentEvent::GoalAchieved {
419            goal: text(provider, goal),
420            total_steps: *total_steps,
421            duration_ms: *duration_ms,
422        },
423        AgentEvent::PersistenceFailed {
424            session_id,
425            operation,
426            error,
427        } => AgentEvent::PersistenceFailed {
428            session_id: session_id.clone(),
429            operation: operation.clone(),
430            error: text(provider, error),
431        },
432        AgentEvent::BudgetThresholdHit {
433            resource,
434            kind,
435            consumed,
436            limit,
437            message,
438        } => AgentEvent::BudgetThresholdHit {
439            resource: resource.clone(),
440            kind: kind.clone(),
441            consumed: *consumed,
442            limit: *limit,
443            message: optional_text(provider, message),
444        },
445        AgentEvent::PassivationRequested {
446            reason,
447            deadline_ms,
448        } => AgentEvent::PassivationRequested {
449            reason: text(provider, reason),
450            deadline_ms: *deadline_ms,
451        },
452        _ => event.clone(),
453    }
454}
455
456/// Trait for pluggable security providers.
457///
458/// Implement this trait to provide custom security logic for sessions.
459/// The default `NoOpSecurityProvider` passes everything through unchanged.
460pub trait SecurityProvider: Send + Sync {
461    /// Classify and register sensitive data found in input text
462    fn taint_input(&self, _text: &str) {}
463
464    /// Sanitize output text by redacting sensitive data.
465    /// Returns the sanitized text.
466    fn sanitize_output(&self, text: &str) -> String {
467        text.to_string()
468    }
469
470    /// Securely wipe all session security state
471    fn wipe(&self) {}
472
473    /// Register security hooks with the given engine
474    fn register_hooks(&self, _hook_engine: &HookEngine) {}
475
476    /// Unregister all hooks from the engine
477    fn teardown(&self, _hook_engine: &HookEngine) {}
478}
479
480/// No-op security provider (default when security is disabled)
481pub struct NoOpSecurityProvider;
482
483impl SecurityProvider for NoOpSecurityProvider {}
484
485#[cfg(test)]
486mod tests {
487    use super::*;
488
489    #[test]
490    fn test_noop_provider_passthrough() {
491        let provider = NoOpSecurityProvider;
492        provider.taint_input("SSN: 123-45-6789");
493        let output = provider.sanitize_output("SSN: 123-45-6789");
494        assert_eq!(output, "SSN: 123-45-6789");
495    }
496
497    #[test]
498    fn test_noop_provider_wipe() {
499        let provider = NoOpSecurityProvider;
500        provider.wipe(); // Should not panic
501    }
502
503    #[test]
504    fn test_noop_provider_hooks() {
505        let engine = HookEngine::new();
506        let provider = NoOpSecurityProvider;
507        provider.register_hooks(&engine);
508        provider.teardown(&engine);
509        assert_eq!(engine.hook_count(), 0);
510    }
511
512    #[test]
513    fn agent_event_sanitization_redacts_streams_arguments_and_outputs() {
514        let provider = DefaultSecurityProvider::new();
515        let secret = "user@example.com";
516        let events = [
517            crate::agent::AgentEvent::TextDelta {
518                text: secret.to_string(),
519            },
520            crate::agent::AgentEvent::ReasoningDelta {
521                text: secret.to_string(),
522            },
523            crate::agent::AgentEvent::ToolInputDelta {
524                id: Some("tool-1".to_string()),
525                delta: format!(r#"{{"token":"{secret}"}}"#),
526            },
527            crate::agent::AgentEvent::ToolExecutionStart {
528                id: "tool-1".to_string(),
529                name: "bash".to_string(),
530                args: serde_json::json!({"command": format!("echo {secret}")}),
531            },
532            crate::agent::AgentEvent::ToolOutputDelta {
533                id: "tool-1".to_string(),
534                name: "bash".to_string(),
535                delta: secret.to_string(),
536            },
537            crate::agent::AgentEvent::ToolEnd {
538                id: "tool-1".to_string(),
539                name: "bash".to_string(),
540                args: Some(serde_json::json!({"command": format!("echo {secret}")})),
541                output: secret.to_string(),
542                exit_code: 0,
543                metadata: Some(serde_json::json!({"contact": secret})),
544                error_kind: None,
545            },
546        ];
547
548        for event in &events {
549            let sanitized = sanitize_agent_event(&provider, event);
550            let json = serde_json::to_string(&sanitized).unwrap();
551            assert!(!json.contains(secret), "unsanitized event: {json}");
552            assert!(json.contains("REDACTED:EMAIL"));
553        }
554
555        let sanitized = sanitize_agent_event(&provider, &events[3]);
556        assert!(matches!(
557            sanitized,
558            crate::agent::AgentEvent::ToolExecutionStart { id, name, .. }
559                if id == "tool-1" && name == "bash"
560        ));
561    }
562
563    #[test]
564    fn agent_event_sanitization_redacts_every_tool_error_kind_string() {
565        use crate::agent::AgentEvent;
566        use crate::tools::ToolErrorKind;
567
568        let provider = DefaultSecurityProvider::new();
569        let secret = "user@example.com";
570        let redacted = "[REDACTED:EMAIL]";
571        let cases = [
572            (
573                ToolErrorKind::VersionConflict {
574                    path: format!("path/{secret}"),
575                    expected: format!("expected: {secret}"),
576                    actual: Some(format!("actual: {secret}")),
577                },
578                ToolErrorKind::VersionConflict {
579                    path: format!("path/{redacted}"),
580                    expected: format!("expected: {redacted}"),
581                    actual: Some(format!("actual: {redacted}")),
582                },
583            ),
584            (
585                ToolErrorKind::RemoteGitConflict {
586                    code: format!("code: {secret}"),
587                    message: format!("message: {secret}"),
588                },
589                ToolErrorKind::RemoteGitConflict {
590                    code: format!("code: {redacted}"),
591                    message: format!("message: {redacted}"),
592                },
593            ),
594            (
595                ToolErrorKind::NotFound {
596                    path: format!("path/{secret}"),
597                },
598                ToolErrorKind::NotFound {
599                    path: format!("path/{redacted}"),
600                },
601            ),
602            (
603                ToolErrorKind::InvalidArgument {
604                    message: format!("message: {secret}"),
605                },
606                ToolErrorKind::InvalidArgument {
607                    message: format!("message: {redacted}"),
608                },
609            ),
610            (
611                ToolErrorKind::Unsupported {
612                    message: format!("message: {secret}"),
613                },
614                ToolErrorKind::Unsupported {
615                    message: format!("message: {redacted}"),
616                },
617            ),
618            (
619                ToolErrorKind::Timeout {
620                    op: format!("operation: {secret}"),
621                    duration_ms: 42,
622                },
623                ToolErrorKind::Timeout {
624                    op: format!("operation: {redacted}"),
625                    duration_ms: 42,
626                },
627            ),
628            (
629                ToolErrorKind::Transport {
630                    op: format!("operation: {secret}"),
631                },
632                ToolErrorKind::Transport {
633                    op: format!("operation: {redacted}"),
634                },
635            ),
636        ];
637
638        for (error_kind, expected) in cases {
639            let event = AgentEvent::ToolEnd {
640                id: "tool-1".to_string(),
641                name: "test".to_string(),
642                args: None,
643                output: String::new(),
644                exit_code: 1,
645                metadata: None,
646                error_kind: Some(error_kind),
647            };
648
649            let sanitized = sanitize_agent_event(&provider, &event);
650            let AgentEvent::ToolEnd {
651                id,
652                name,
653                exit_code,
654                error_kind,
655                ..
656            } = sanitized
657            else {
658                panic!("sanitization changed the event variant");
659            };
660
661            assert_eq!(id, "tool-1");
662            assert_eq!(name, "test");
663            assert_eq!(exit_code, 1);
664            assert_eq!(error_kind, Some(expected));
665        }
666    }
667}