Skip to main content

tea_tools/
effect.rs

1use std::str::FromStr;
2
3use serde::{Deserialize, Deserializer, Serialize, Serializer};
4use thiserror::Error;
5
6/// Declared capability or external effect of a tool invocation.
7#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
8pub enum ToolEffect {
9    /// Read filesystem content.
10    FsRead,
11    /// Create or modify filesystem content.
12    FsWrite,
13    /// Delete filesystem content.
14    FsDelete,
15    /// Spawn or control an operating-system process.
16    ProcessSpawn,
17    /// Perform a network request.
18    NetworkRequest,
19    /// Read credentials or secret material.
20    CredentialRead,
21    /// Read clipboard content.
22    ClipboardRead,
23    /// Require direct user interaction.
24    UserInteraction,
25    /// Mutate state in an external system.
26    ExternalMutation,
27    /// Future namespaced effect unknown to this runtime version.
28    Unknown(String),
29}
30
31impl ToolEffect {
32    /// Returns the canonical dotted effect name.
33    #[must_use]
34    pub fn as_str(&self) -> &str {
35        match self {
36            Self::FsRead => "fs.read",
37            Self::FsWrite => "fs.write",
38            Self::FsDelete => "fs.delete",
39            Self::ProcessSpawn => "process.spawn",
40            Self::NetworkRequest => "network.request",
41            Self::CredentialRead => "credential.read",
42            Self::ClipboardRead => "clipboard.read",
43            Self::UserInteraction => "user.interaction",
44            Self::ExternalMutation => "external.mutation",
45            Self::Unknown(value) => value,
46        }
47    }
48
49    /// Returns whether this runtime does not understand the effect semantics.
50    #[must_use]
51    pub const fn is_unknown(&self) -> bool {
52        matches!(self, Self::Unknown(_))
53    }
54
55    pub(crate) const fn is_read_only(&self) -> bool {
56        matches!(
57            self,
58            Self::FsRead | Self::CredentialRead | Self::ClipboardRead
59        )
60    }
61}
62
63impl Serialize for ToolEffect {
64    fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
65    where
66        S: Serializer,
67    {
68        serializer.serialize_str(self.as_str())
69    }
70}
71
72impl<'de> Deserialize<'de> for ToolEffect {
73    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
74    where
75        D: Deserializer<'de>,
76    {
77        String::deserialize(deserializer)?
78            .parse()
79            .map_err(serde::de::Error::custom)
80    }
81}
82
83impl FromStr for ToolEffect {
84    type Err = ToolEffectParseError;
85
86    fn from_str(value: &str) -> Result<Self, Self::Err> {
87        let known = match value {
88            "fs.read" => Some(Self::FsRead),
89            "fs.write" => Some(Self::FsWrite),
90            "fs.delete" => Some(Self::FsDelete),
91            "process.spawn" => Some(Self::ProcessSpawn),
92            "network.request" => Some(Self::NetworkRequest),
93            "credential.read" => Some(Self::CredentialRead),
94            "clipboard.read" => Some(Self::ClipboardRead),
95            "user.interaction" => Some(Self::UserInteraction),
96            "external.mutation" => Some(Self::ExternalMutation),
97            _ => None,
98        };
99        if let Some(known) = known {
100            return Ok(known);
101        }
102        if value.len() > 256 || !valid_namespaced_effect(value) {
103            return Err(ToolEffectParseError);
104        }
105        Ok(Self::Unknown(value.to_owned()))
106    }
107}
108
109/// Error returned when parsing a tool effect.
110#[derive(Debug, Clone, Copy, PartialEq, Eq, Error)]
111#[error("tool effect must be known or use a lowercase namespaced dotted name")]
112pub struct ToolEffectParseError;
113
114fn valid_namespaced_effect(value: &str) -> bool {
115    let segments = value.split('.').collect::<Vec<_>>();
116    segments.len() >= 3
117        && segments.iter().all(|segment| {
118            let mut bytes = segment.bytes();
119            bytes.next().is_some_and(|byte| byte.is_ascii_lowercase())
120                && bytes.all(|byte| {
121                    byte.is_ascii_lowercase()
122                        || byte.is_ascii_digit()
123                        || matches!(byte, b'-' | b'_')
124                })
125        })
126}