use crate::error::PolicyError;
use crate::ids::{RunId, SessionId, ToolCallId};
use crate::tool::{Capability, PreparedToolCall};
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use tokio_util::sync::CancellationToken;
#[derive(Clone)]
pub struct PolicyRequest {
pub call: PreparedToolCall,
pub session_id: SessionId,
pub run_id: RunId,
pub turn: u32,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "decision")]
pub enum Decision {
Allow,
Deny { reason: String },
Ask,
}
#[async_trait]
pub trait Policy: Send + Sync {
async fn evaluate(&self, request: PolicyRequest) -> Result<Decision, PolicyError>;
}
pub struct AllowAllPolicy;
#[async_trait]
impl Policy for AllowAllPolicy {
async fn evaluate(&self, _request: PolicyRequest) -> Result<Decision, PolicyError> {
Ok(Decision::Allow)
}
}
pub struct DenyAllPolicy {
pub reason: String,
}
#[async_trait]
impl Policy for DenyAllPolicy {
async fn evaluate(&self, _request: PolicyRequest) -> Result<Decision, PolicyError> {
Ok(Decision::Deny {
reason: self.reason.clone(),
})
}
}
pub struct ApprovalRequest {
pub call_id: ToolCallId,
pub name: String,
pub capabilities: Vec<Capability>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ApprovalDecision {
Approved,
Denied { reason: String },
}
#[async_trait]
pub trait ApprovalHandler: Send + Sync {
async fn wait(&self, request: ApprovalRequest, cancel: CancellationToken) -> ApprovalDecision;
}
#[async_trait]
pub trait SteerSource: Send + Sync {
async fn pending(&self) -> Vec<String>;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decision_roundtrips() {
for decision in [
Decision::Allow,
Decision::Deny {
reason: "机密".into(),
},
Decision::Ask,
] {
let json = serde_json::to_string(&decision).unwrap();
let back: Decision = serde_json::from_str(&json).unwrap();
assert_eq!(back, decision);
}
}
}