Skip to main content

rig_core/tool/
context.rs

1//! Typed inbound context and explicitly published host-only tool-result metadata.
2//!
3//! ```
4//! use rig_core::tool::{ContextValue, ToolContext};
5//! #[derive(serde::Serialize, serde::Deserialize)]
6//! struct Session(String);
7//! impl ContextValue for Session { const KEY: &'static str = "session"; }
8//! let mut context = ToolContext::new();
9//! context.insert(Session("example".into()))?;
10//! assert_eq!(context.require::<Session>()?.0, "example");
11//! # Ok::<(), rig_core::tool::ToolContextError>(())
12//! ```
13
14use std::any::{Any, TypeId};
15use std::collections::{BTreeMap, HashMap};
16use std::hash::{BuildHasherDefault, Hasher};
17
18use serde::{Deserialize, Serialize, de::DeserializeOwned};
19
20use super::ToolExecutionError;
21use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
22
23type AnyMap = HashMap<TypeId, Box<dyn AnyClone>, BuildHasherDefault<IdHasher>>;
24
25#[derive(Default)]
26struct IdHasher(u64);
27
28impl Hasher for IdHasher {
29    fn write_u64(&mut self, id: u64) {
30        self.0 = id;
31    }
32
33    fn write(&mut self, bytes: &[u8]) {
34        for &byte in bytes {
35            self.0 = self.0.rotate_left(8) ^ u64::from(byte);
36        }
37    }
38
39    fn finish(&self) -> u64 {
40        self.0
41    }
42}
43
44trait AnyClone: Any + WasmCompatSend + WasmCompatSync {
45    fn clone_box(&self) -> Box<dyn AnyClone>;
46    fn as_any(&self) -> &dyn Any;
47    fn as_any_mut(&mut self) -> &mut dyn Any;
48    fn into_any(self: Box<Self>) -> Box<dyn Any>;
49}
50
51impl<T> AnyClone for T
52where
53    T: Clone + WasmCompatSend + WasmCompatSync + 'static,
54{
55    fn clone_box(&self) -> Box<dyn AnyClone> {
56        Box::new(self.clone())
57    }
58
59    fn as_any(&self) -> &dyn Any {
60        self
61    }
62
63    fn as_any_mut(&mut self) -> &mut dyn Any {
64        self
65    }
66
67    fn into_any(self: Box<Self>) -> Box<dyn Any> {
68        self
69    }
70}
71
72impl Clone for Box<dyn AnyClone> {
73    fn clone(&self) -> Self {
74        (**self).clone_box()
75    }
76}
77
78/// A clone-on-dispatch map of live values keyed by type. Not the storage
79/// behind [`ToolContext`] (which is serde data); runtimes use it for
80/// in-process per-run state that never crosses a wire (rig-agent's hook state
81/// does).
82#[derive(Default, Clone)]
83pub struct TypeMap {
84    map: AnyMap,
85}
86
87impl TypeMap {
88    pub fn insert<T>(&mut self, value: T) -> Option<T>
89    where
90        T: Clone + WasmCompatSend + WasmCompatSync + 'static,
91    {
92        self.map
93            .insert(TypeId::of::<T>(), Box::new(value))
94            .and_then(|previous| previous.into_any().downcast::<T>().ok())
95            .map(|value| *value)
96    }
97
98    pub fn get<T>(&self) -> Option<&T>
99    where
100        T: 'static,
101    {
102        self.map
103            .get(&TypeId::of::<T>())
104            .and_then(|value| (**value).as_any().downcast_ref::<T>())
105    }
106
107    pub fn get_mut<T>(&mut self) -> Option<&mut T>
108    where
109        T: 'static,
110    {
111        self.map
112            .get_mut(&TypeId::of::<T>())
113            .and_then(|value| (**value).as_any_mut().downcast_mut::<T>())
114    }
115
116    pub fn remove<T>(&mut self) -> Option<T>
117    where
118        T: 'static,
119    {
120        self.map
121            .remove(&TypeId::of::<T>())
122            .and_then(|value| value.into_any().downcast::<T>().ok())
123            .map(|value| *value)
124    }
125
126    pub fn contains<T>(&self) -> bool
127    where
128        T: 'static,
129    {
130        self.map.contains_key(&TypeId::of::<T>())
131    }
132
133    /// Number of values held.
134    pub fn len(&self) -> usize {
135        self.map.len()
136    }
137
138    /// Whether the map holds no values.
139    pub fn is_empty(&self) -> bool {
140        self.map.is_empty()
141    }
142}
143
144/// Context passed to every tool execution.
145///
146/// Callers insert typed inbound values with [`insert`](Self::insert). Tools read
147/// those values with [`get`](Self::get) or [`require`](Self::require), and attach
148/// host-only result metadata with [`insert_result`](Self::insert_result). Result
149/// hooks inspect that metadata through [`result`](Self::result). Neither inbound
150/// values nor result metadata are sent to the model.
151///
152/// Every value is **data**: inserting serializes it, reading deserializes it,
153/// and the whole context is itself `Serialize + Deserialize`. Explicit scene
154/// serialization can therefore contain sensitive inbound values and must be
155/// protected by the host. Effect logs record only [`ToolResultContext`]:
156/// values explicitly published through `insert_result`, never inbound slots
157/// or runtime scopes. Replay uses the current dispatch's inbound context;
158/// code must not treat mutations to inbound slots as durable output. Values with
159/// interior sharing (`Arc<Mutex<_>>`, atomics, channels) do not belong here;
160/// they belong to the tool instance.
161///
162/// Registry, server, and agent dispatch clone inbound values once per call.
163/// Map-level mutations (inserting, replacing, or removing typed slots) affect
164/// only that execution. The dispatch surface returns result metadata without
165/// replacing the caller's inbound slots.
166///
167/// Slots use the value type's declared [`ContextValue::KEY`]. Accessors need
168/// no key argument; types must use distinct keys to occupy distinct slots.
169#[derive(Default, Clone, Serialize, Deserialize)]
170pub struct ToolContext {
171    #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
172    inbound: BTreeMap<String, serde_json::Value>,
173    #[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
174    result: BTreeMap<String, serde_json::Value>,
175    /// Live runtime scopes, excluded from serialization and equality.
176    /// Publication clears them before the result is resolved.
177    #[serde(skip)]
178    scopes: Vec<std::sync::Arc<dyn Any + Send + Sync>>,
179}
180
181/// Explicitly published, durable tool-result metadata. Unlike [`ToolContext`],
182/// this contains no inbound credentials or live runtime scopes. Values are
183/// keyed by [`ContextValue::KEY`]; applications own their schema versions.
184#[derive(Default, Clone, PartialEq, Eq, Serialize, Deserialize)]
185#[serde(transparent)]
186pub struct ToolResultContext(BTreeMap<String, serde_json::Value>);
187
188impl std::fmt::Debug for ToolResultContext {
189    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
190        f.debug_tuple("ToolResultContext")
191            .field(&self.0.keys().collect::<Vec<_>>())
192            .finish()
193    }
194}
195
196impl ToolResultContext {
197    /// Read a published value without exposing ambient execution inputs.
198    pub fn get<T: ContextValue>(&self) -> Result<Option<T>, ToolContextError> {
199        get_slot(&self.0)
200    }
201}
202
203impl PartialEq for ToolContext {
204    fn eq(&self, other: &Self) -> bool {
205        self.inbound == other.inbound && self.result == other.result
206    }
207}
208
209impl Eq for ToolContext {}
210
211/// Serializable context data with a stable slot key. Distinct value types must
212/// use distinct keys, and persisted keys must remain stable across schema changes.
213/// The derive macro defaults to the type name; `#[context(key = "…")]` overrides it.
214#[diagnostic::on_unimplemented(
215    message = "`{Self}` declares no `ToolContext` key",
216    label = "not a `ContextValue`",
217    note = "derive it (`#[derive(rig::ContextValue)]`, optionally `#[context(key = \"…\")]`) or write `impl ContextValue for {Self} {{ const KEY: &'static str = \"…\"; }}`; a bare `String`, integer or `serde_json::Value` cannot be stored — wrap it in a newtype"
218)]
219pub trait ContextValue: Serialize + DeserializeOwned + 'static {
220    /// The slot this value lives under.
221    const KEY: &'static str;
222}
223
224fn encode<T: ContextValue>(value: &T) -> Result<serde_json::Value, ToolContextError> {
225    serde_json::to_value(value).map_err(|error| ToolContextError::Encode {
226        key: T::KEY,
227        message: error.to_string(),
228    })
229}
230
231fn decode<T: ContextValue>(value: &serde_json::Value) -> Result<T, ToolContextError> {
232    serde_json::from_value(value.clone()).map_err(|error| ToolContextError::Decode {
233        key: T::KEY,
234        message: error.to_string(),
235    })
236}
237
238/// Replaces a slot after successful encoding, returning the displaced value
239/// only if it decodes as `T`.
240fn insert_slot<T: ContextValue>(
241    map: &mut BTreeMap<String, serde_json::Value>,
242    value: T,
243) -> Result<Option<T>, ToolContextError> {
244    let encoded = encode(&value)?;
245    Ok(map
246        .insert(T::KEY.to_owned(), encoded)
247        .and_then(|previous| decode(&previous).ok()))
248}
249
250/// `Ok(None)` when the slot is empty, `Err(Decode)` when it holds something
251/// that is not a `T`: absence and a shape mismatch are different facts.
252fn get_slot<T: ContextValue>(
253    map: &BTreeMap<String, serde_json::Value>,
254) -> Result<Option<T>, ToolContextError> {
255    map.get(T::KEY).map(decode).transpose()
256}
257
258fn require_slot<T: ContextValue>(
259    map: &BTreeMap<String, serde_json::Value>,
260) -> Result<T, ToolContextError> {
261    get_slot(map)?.ok_or(ToolContextError::Missing(T::KEY))
262}
263
264impl ToolContext {
265    /// Create an empty context.
266    pub const fn new() -> Self {
267        Self {
268            inbound: BTreeMap::new(),
269            result: BTreeMap::new(),
270            scopes: Vec::new(),
271        }
272    }
273
274    /// Insert an inbound typed value, returning the displaced value if it decodes
275    /// as `T`. A differently shaped displaced value is replaced successfully and
276    /// returns `None`; decoding the previous value cannot undo a successful write.
277    ///
278    /// Returns an encoding error if serialization fails, leaving the slot unchanged.
279    pub fn insert<T: ContextValue>(&mut self, value: T) -> Result<Option<T>, ToolContextError> {
280        insert_slot(&mut self.inbound, value)
281    }
282
283    /// Read an inbound typed value: `Ok(None)` when absent, a decode error when
284    /// the stored value does not match `T`.
285    pub fn get<T: ContextValue>(&self) -> Result<Option<T>, ToolContextError> {
286        get_slot(&self.inbound)
287    }
288
289    /// Require an inbound typed value.
290    pub fn require<T: ContextValue>(&self) -> Result<T, ToolContextError> {
291        require_slot(&self.inbound)
292    }
293
294    /// Remove an inbound typed value. The slot is removed even if decoding its
295    /// former value fails.
296    pub fn remove<T: ContextValue>(&mut self) -> Result<Option<T>, ToolContextError> {
297        self.inbound
298            .remove(T::KEY)
299            .map(|value| decode(&value))
300            .transpose()
301    }
302
303    /// Attach host-only metadata to this execution's result.
304    pub fn insert_result<T: ContextValue>(
305        &mut self,
306        value: T,
307    ) -> Result<Option<T>, ToolContextError> {
308        insert_slot(&mut self.result, value)
309    }
310
311    /// Read host-only result metadata.
312    pub fn result<T: ContextValue>(&self) -> Result<Option<T>, ToolContextError> {
313        get_slot(&self.result)
314    }
315
316    /// Require host-only result metadata.
317    pub fn require_result<T: ContextValue>(&self) -> Result<T, ToolContextError> {
318        require_slot(&self.result)
319    }
320
321    /// Whether this context contains the inbound type `T`.
322    pub fn contains<T: ContextValue>(&self) -> bool {
323        self.inbound.contains_key(T::KEY)
324    }
325
326    /// Whether both maps are empty.
327    pub fn is_empty(&self) -> bool {
328        self.inbound.is_empty() && self.result.is_empty()
329    }
330
331    /// Build the snapshot one dispatch runs against: the same inbound values
332    /// and no result metadata. Runtimes call this per tool call so map-level
333    /// mutations inside the call cannot leak into the caller's context.
334    pub fn for_dispatch(&self) -> Self {
335        Self {
336            inbound: self.inbound.clone(),
337            result: BTreeMap::new(),
338            scopes: self.scopes.clone(),
339        }
340    }
341
342    /// Attach one of the driver's scopes for the call (see the field).
343    pub fn with_scope(mut self, scope: std::sync::Arc<dyn Any + Send + Sync>) -> Self {
344        self.scopes.push(scope);
345        self
346    }
347
348    /// Attach the driver's scopes for the call, as the sink carries them.
349    pub fn with_scopes(mut self, scopes: Vec<std::sync::Arc<dyn Any + Send + Sync>>) -> Self {
350        self.scopes.extend(scopes);
351        self
352    }
353
354    /// Returns the first attached scope of type `T`, or `None` if absent.
355    pub fn scope<T: Any + Send + Sync>(&self) -> Option<std::sync::Arc<T>> {
356        self.scopes
357            .iter()
358            .find_map(|scope| std::sync::Arc::downcast::<T>(scope.clone()).ok())
359    }
360
361    /// Drop the scopes: the call is over and the context is data again.
362    pub fn clear_scope(&mut self) {
363        self.scopes.clear();
364    }
365
366    /// Publish the result metadata a dispatch produced (see
367    /// [`Self::for_dispatch`]) while keeping the caller's inbound values.
368    pub fn accept_dispatch_result(&mut self, dispatched: Self) {
369        self.result = dispatched.result;
370    }
371
372    /// Snapshot only the values explicitly published with `insert_result`.
373    pub fn result_context(&self) -> ToolResultContext {
374        ToolResultContext(self.result.clone())
375    }
376
377    /// Restore durable output while preserving this dispatch's inbound values.
378    pub fn with_result_context(mut self, result: ToolResultContext) -> Self {
379        self.result = result.0;
380        self
381    }
382
383    /// Drop the previous dispatch's result metadata before starting another.
384    pub fn clear_dispatch_result(&mut self) {
385        self.result.clear();
386    }
387}
388
389impl std::fmt::Debug for ToolContext {
390    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
391        f.debug_struct("ToolContext")
392            .field("inbound_types", &self.inbound.keys().collect::<Vec<_>>())
393            .field("result_types", &self.result.keys().collect::<Vec<_>>())
394            .finish()
395    }
396}
397
398/// What a tool published into its dispatch context, handed back beside the
399/// sink rather than on the wire: the driver attaches an empty one to the
400/// sink's scope (`Dispatch::with_scope`), the adapter fills it once the
401/// tool ran, the driver reads it after the outcome. One per dispatch; a
402/// second publish replaces the first.
403#[derive(Debug, Default)]
404pub struct PublishedContext(std::sync::Mutex<Option<ToolContext>>);
405
406impl PublishedContext {
407    /// An empty one, shared between the driver and the sink.
408    pub fn new() -> std::sync::Arc<Self> {
409        std::sync::Arc::new(Self::default())
410    }
411
412    /// The recorder's non-consuming snapshot. Publish before resolving the
413    /// sink; values published after resolution are not part of that outcome.
414    pub fn result_context(&self) -> Option<ToolResultContext> {
415        self.0
416            .lock()
417            .unwrap_or_else(std::sync::PoisonError::into_inner)
418            .as_ref()
419            .map(ToolContext::result_context)
420    }
421
422    /// The adapter's write: the context after the tool ran, scopes cleared.
423    /// All publication must finish before resolving the sink. Publishing
424    /// afterward violates the delivery contract: consumers and recorders may
425    /// observe different versions. The slot does not enforce this ordering.
426    pub fn publish(&self, mut context: ToolContext) {
427        context.clear_scope();
428        *self
429            .0
430            .lock()
431            .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(context);
432    }
433
434    /// The driver's read: what was published, if anything, leaving nothing.
435    pub fn take(&self) -> Option<ToolContext> {
436        self.0
437            .lock()
438            .unwrap_or_else(std::sync::PoisonError::into_inner)
439            .take()
440    }
441}
442
443/// A [`ToolContext`] slot could not be read or written.
444#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
445pub enum ToolContextError {
446    /// A required typed value was absent.
447    #[error("required tool context value `{0}` was not found")]
448    Missing(&'static str),
449    /// A value could not be represented as JSON.
450    #[error("tool context value `{key}` could not be encoded: {message}")]
451    Encode {
452        /// The slot's type name.
453        key: &'static str,
454        /// The serializer's message.
455        message: String,
456    },
457    /// A stored value could not be decoded as the requested type.
458    #[error("tool context value `{key}` could not be decoded: {message}")]
459    Decode {
460        /// The slot's type name.
461        key: &'static str,
462        /// The deserializer's message.
463        message: String,
464    },
465}
466
467impl From<ToolContextError> for ToolExecutionError {
468    fn from(error: ToolContextError) -> Self {
469        ToolExecutionError::other(error.to_string()).with_source(error)
470    }
471}
472
473// The context crosses the effect wire: it must serialize and cross threads on
474// every target, browser wasm included.
475const _: fn() = || {
476    fn assert_wire<T: Send + Sync + 'static + Serialize + DeserializeOwned>() {}
477    assert_wire::<ToolContext>();
478    fn assert_shared<T: Send + Sync + 'static>() {}
479    assert_shared::<PublishedContext>();
480};
481
482#[cfg(test)]
483mod tests;
484
485#[cfg(test)]
486mod migrated_tests;