use std::collections::HashMap;
use std::fmt;
use std::sync::Arc;
use ferrin_message::Message;
use ferrin_spec::BoxFuture;
use ferrin_spec::JsonValue;
use ferrin_spec::ToolName;
use ferrin_tool::Tool;
use ferrin_tool::ToolContext;
use serde::Deserialize;
use serde::Serialize;
use super::ParsedToolCall;
pub(crate) mod collect;
pub(crate) mod signature;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "status", rename_all = "kebab-case")]
#[non_exhaustive]
pub enum ApprovalStatus {
NotApplicable,
Approved {
#[serde(default, skip_serializing_if = "Option::is_none")]
reason: Option<String>,
},
Denied {
#[serde(default, skip_serializing_if = "Option::is_none")]
reason: Option<String>,
},
UserApproval {
#[serde(default, skip_serializing_if = "Option::is_none")]
reason: Option<String>,
},
}
impl ApprovalStatus {
#[must_use]
pub fn approved() -> Self {
Self::Approved { reason: None }
}
#[must_use]
pub fn denied() -> Self {
Self::Denied { reason: None }
}
#[must_use]
pub fn user_approval() -> Self {
Self::UserApproval { reason: None }
}
#[must_use]
pub fn with_reason(self, reason: impl Into<String>) -> Self {
let reason = Some(reason.into());
match self {
Self::NotApplicable => Self::NotApplicable,
Self::Approved { .. } => Self::Approved { reason },
Self::Denied { .. } => Self::Denied { reason },
Self::UserApproval { .. } => Self::UserApproval { reason },
}
}
#[must_use]
pub fn reason(&self) -> Option<&str> {
match self {
Self::NotApplicable => None,
Self::Approved { reason } | Self::Denied { reason } | Self::UserApproval { reason } => {
reason.as_deref()
}
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct ApprovalContext<'a> {
pub messages: &'a [Message],
pub tools_context: Option<&'a JsonValue>,
pub runtime_context: Option<&'a JsonValue>,
}
pub trait ApprovalPolicy: Send + Sync {
fn resolve<'a>(
&'a self,
call: &'a ParsedToolCall,
ctx: ApprovalContext<'a>,
) -> BoxFuture<'a, Option<ApprovalStatus>>;
}
impl<P: ApprovalPolicy + ?Sized> ApprovalPolicy for Arc<P> {
fn resolve<'a>(
&'a self,
call: &'a ParsedToolCall,
ctx: ApprovalContext<'a>,
) -> BoxFuture<'a, Option<ApprovalStatus>> {
self.as_ref().resolve(call, ctx)
}
}
impl ApprovalPolicy for ApprovalStatus {
fn resolve<'a>(
&'a self,
_call: &'a ParsedToolCall,
_ctx: ApprovalContext<'a>,
) -> BoxFuture<'a, Option<ApprovalStatus>> {
Box::pin(async move { Some(self.clone()) })
}
}
impl ApprovalPolicy for HashMap<ToolName, ApprovalStatus> {
fn resolve<'a>(
&'a self,
call: &'a ParsedToolCall,
_ctx: ApprovalContext<'a>,
) -> BoxFuture<'a, Option<ApprovalStatus>> {
let status = self.get(&call.tool_name).cloned();
Box::pin(async move { status })
}
}
pub struct ApprovalPolicyFn<F>(F);
pub fn approval_policy<F>(f: F) -> ApprovalPolicyFn<F>
where
F: Fn(&ParsedToolCall, &ApprovalContext<'_>) -> Option<ApprovalStatus> + Send + Sync,
{
ApprovalPolicyFn(f)
}
impl<F> fmt::Debug for ApprovalPolicyFn<F> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("ApprovalPolicyFn(..)")
}
}
impl<F> ApprovalPolicy for ApprovalPolicyFn<F>
where
F: Fn(&ParsedToolCall, &ApprovalContext<'_>) -> Option<ApprovalStatus> + Send + Sync,
{
fn resolve<'a>(
&'a self,
call: &'a ParsedToolCall,
ctx: ApprovalContext<'a>,
) -> BoxFuture<'a, Option<ApprovalStatus>> {
let status = (self.0)(call, &ctx).unwrap_or(ApprovalStatus::NotApplicable);
Box::pin(async move { Some(status) })
}
}
pub(crate) async fn resolve_approval(
call: &ParsedToolCall,
tool: Option<&Tool>,
policy: Option<&dyn ApprovalPolicy>,
ctx: ApprovalContext<'_>,
tool_ctx: impl FnOnce() -> ToolContext,
) -> ApprovalStatus {
if let Some(policy) = policy
&& let Some(status) = policy.resolve(call, ctx).await
{
return status;
}
let Some(tool) = tool else {
return ApprovalStatus::NotApplicable;
};
if !tool.needs_approval().is_declared() {
return ApprovalStatus::NotApplicable;
}
if tool
.needs_approval()
.resolve(call.input.clone(), tool_ctx())
.await
{
ApprovalStatus::UserApproval { reason: None }
} else {
ApprovalStatus::NotApplicable
}
}