Skip to main content

rig_core/completion/message/
identity.rs

1//! Identities a message carries: a tool call's one [`CallId`] and a tool's
2//! [`ToolName`].
3
4use std::borrow::Cow;
5
6use serde::{Deserialize, Serialize};
7
8/// A tool name, never empty.
9///
10/// ```
11/// use rig_core::message::ToolName;
12///
13/// let name = ToolName::new("add")?;
14/// assert_eq!(name.as_str(), "add");
15/// assert!(ToolName::new("").is_err());
16/// # Ok::<(), rig_core::message::EmptyToolName>(())
17/// ```
18#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
19#[serde(try_from = "String", into = "String")]
20pub struct ToolName(String);
21
22/// A tool name was empty.
23#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
24#[error("a tool name cannot be empty")]
25pub struct EmptyToolName;
26
27impl ToolName {
28    /// `name`, or [`EmptyToolName`] when it is empty.
29    pub fn new(name: impl Into<String>) -> Result<Self, EmptyToolName> {
30        let name = name.into();
31        if name.is_empty() {
32            Err(EmptyToolName)
33        } else {
34            Ok(Self(name))
35        }
36    }
37
38    /// `name` without the emptiness check. Callers prove `name` is not empty,
39    /// for example with a compile-time assertion on a `Tool::NAME`.
40    pub(crate) fn new_unchecked(name: &str) -> Self {
41        debug_assert!(!name.is_empty());
42        Self(name.to_owned())
43    }
44
45    /// The name.
46    pub fn as_str(&self) -> &str {
47        &self.0
48    }
49}
50
51impl TryFrom<String> for ToolName {
52    type Error = EmptyToolName;
53
54    fn try_from(name: String) -> Result<Self, EmptyToolName> {
55        Self::new(name)
56    }
57}
58
59impl TryFrom<&str> for ToolName {
60    type Error = EmptyToolName;
61
62    fn try_from(name: &str) -> Result<Self, EmptyToolName> {
63        Self::new(name)
64    }
65}
66
67impl From<ToolName> for String {
68    fn from(name: ToolName) -> Self {
69        name.0
70    }
71}
72
73impl std::ops::Deref for ToolName {
74    type Target = str;
75
76    fn deref(&self) -> &str {
77        &self.0
78    }
79}
80
81impl AsRef<str> for ToolName {
82    fn as_ref(&self) -> &str {
83        &self.0
84    }
85}
86
87impl std::borrow::Borrow<str> for ToolName {
88    fn borrow(&self) -> &str {
89        &self.0
90    }
91}
92
93impl std::fmt::Display for ToolName {
94    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
95        f.write_str(&self.0)
96    }
97}
98
99impl PartialEq<str> for ToolName {
100    fn eq(&self, other: &str) -> bool {
101        self.0 == other
102    }
103}
104
105impl PartialEq<&str> for ToolName {
106    fn eq(&self, other: &&str) -> bool {
107        self.0 == *other
108    }
109}
110
111impl PartialEq<String> for ToolName {
112    fn eq(&self, other: &String) -> bool {
113        &self.0 == other
114    }
115}
116
117impl PartialEq<ToolName> for String {
118    fn eq(&self, other: &ToolName) -> bool {
119        *self == other.0
120    }
121}
122
123impl PartialEq<ToolName> for str {
124    fn eq(&self, other: &ToolName) -> bool {
125        self == other.0
126    }
127}
128
129/// A tool call's one identity: the id the provider issued, or one rig issued
130/// because the provider sent none.
131///
132/// A result copies the id of the call it answers ([`ToolCall::result`]), so
133/// the two always match.
134///
135/// ```
136/// use rig_core::message::{CallId, ProviderCallId};
137///
138/// let id = CallId::from(ProviderCallId::new("call_1").ok_or("empty id")?);
139/// assert_eq!(id.to_string(), "call_1");
140/// assert_eq!(id.provider().map(ProviderCallId::as_str), Some("call_1"));
141/// # Ok::<(), Box<dyn std::error::Error>>(())
142/// ```
143///
144/// [`ToolCall::result`]: super::ToolCall::result
145#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
146#[serde(rename_all = "snake_case")]
147pub enum CallId {
148    /// The provider's id.
149    Provider(ProviderCallId),
150    /// An identifier rig issued for a call the provider sent without one.
151    Local(LocalCallId),
152}
153
154impl CallId {
155    /// The provider's id `call_id`, or a rig-issued id when it is empty.
156    pub fn from_wire(call_id: impl Into<String>) -> Self {
157        ProviderCallId::new(call_id).map_or_else(|| Self::Local(LocalCallId::new()), Self::Provider)
158    }
159
160    /// The provider's id, when the provider issued it.
161    pub fn provider(&self) -> Option<&ProviderCallId> {
162        match self {
163            Self::Provider(provider) => Some(provider),
164            Self::Local(_) => None,
165        }
166    }
167
168    /// Whether rig issued this id.
169    pub fn is_local(&self) -> bool {
170        matches!(self, Self::Local(_))
171    }
172
173    /// The id as a wire sends it: the provider's call id, or the rig-issued
174    /// UUID.
175    pub fn wire(&self) -> Cow<'_, str> {
176        match self {
177            Self::Provider(provider) => Cow::Borrowed(provider.as_str()),
178            Self::Local(local) => Cow::Owned(local.to_string()),
179        }
180    }
181}
182
183impl From<ProviderCallId> for CallId {
184    fn from(provider: ProviderCallId) -> Self {
185        Self::Provider(provider)
186    }
187}
188
189impl From<LocalCallId> for CallId {
190    fn from(local: LocalCallId) -> Self {
191        Self::Local(local)
192    }
193}
194
195impl std::fmt::Display for CallId {
196    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
197        f.write_str(&self.wire())
198    }
199}
200
201/// A call id rig issued: a random (v4) UUID.
202#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
203#[serde(transparent)]
204pub struct LocalCallId(uuid::Uuid);
205
206impl LocalCallId {
207    /// A fresh id.
208    pub fn new() -> Self {
209        Self(uuid::Uuid::new_v4())
210    }
211}
212
213impl Default for LocalCallId {
214    fn default() -> Self {
215        Self::new()
216    }
217}
218
219impl std::fmt::Display for LocalCallId {
220    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
221        self.0.fmt(f)
222    }
223}
224
225/// The call-correlation id a provider issued and expects echoed back. Never
226/// empty.
227#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
228#[serde(try_from = "String", into = "String")]
229pub struct ProviderCallId(String);
230
231/// A provider call id was empty.
232#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
233#[error("a provider call id cannot be empty")]
234pub struct EmptyCallId;
235
236impl ProviderCallId {
237    /// Adopt a provider-issued call identifier. `None` for the empty string:
238    /// absence is not an id.
239    pub fn new(call_id: impl Into<String>) -> Option<Self> {
240        let call_id = call_id.into();
241        (!call_id.is_empty()).then_some(Self(call_id))
242    }
243
244    /// The id.
245    pub fn as_str(&self) -> &str {
246        &self.0
247    }
248}
249
250impl TryFrom<String> for ProviderCallId {
251    type Error = EmptyCallId;
252
253    fn try_from(call_id: String) -> Result<Self, EmptyCallId> {
254        Self::new(call_id).ok_or(EmptyCallId)
255    }
256}
257
258impl From<ProviderCallId> for String {
259    fn from(id: ProviderCallId) -> Self {
260        id.0
261    }
262}