1use 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#[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 pub fn len(&self) -> usize {
135 self.map.len()
136 }
137
138 pub fn is_empty(&self) -> bool {
140 self.map.is_empty()
141 }
142}
143
144#[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 #[serde(skip)]
178 scopes: Vec<std::sync::Arc<dyn Any + Send + Sync>>,
179}
180
181#[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 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#[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 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
238fn 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
250fn 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 pub const fn new() -> Self {
267 Self {
268 inbound: BTreeMap::new(),
269 result: BTreeMap::new(),
270 scopes: Vec::new(),
271 }
272 }
273
274 pub fn insert<T: ContextValue>(&mut self, value: T) -> Result<Option<T>, ToolContextError> {
280 insert_slot(&mut self.inbound, value)
281 }
282
283 pub fn get<T: ContextValue>(&self) -> Result<Option<T>, ToolContextError> {
286 get_slot(&self.inbound)
287 }
288
289 pub fn require<T: ContextValue>(&self) -> Result<T, ToolContextError> {
291 require_slot(&self.inbound)
292 }
293
294 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 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 pub fn result<T: ContextValue>(&self) -> Result<Option<T>, ToolContextError> {
313 get_slot(&self.result)
314 }
315
316 pub fn require_result<T: ContextValue>(&self) -> Result<T, ToolContextError> {
318 require_slot(&self.result)
319 }
320
321 pub fn contains<T: ContextValue>(&self) -> bool {
323 self.inbound.contains_key(T::KEY)
324 }
325
326 pub fn is_empty(&self) -> bool {
328 self.inbound.is_empty() && self.result.is_empty()
329 }
330
331 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 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 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 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 pub fn clear_scope(&mut self) {
363 self.scopes.clear();
364 }
365
366 pub fn accept_dispatch_result(&mut self, dispatched: Self) {
369 self.result = dispatched.result;
370 }
371
372 pub fn result_context(&self) -> ToolResultContext {
374 ToolResultContext(self.result.clone())
375 }
376
377 pub fn with_result_context(mut self, result: ToolResultContext) -> Self {
379 self.result = result.0;
380 self
381 }
382
383 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#[derive(Debug, Default)]
404pub struct PublishedContext(std::sync::Mutex<Option<ToolContext>>);
405
406impl PublishedContext {
407 pub fn new() -> std::sync::Arc<Self> {
409 std::sync::Arc::new(Self::default())
410 }
411
412 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 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 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#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
445pub enum ToolContextError {
446 #[error("required tool context value `{0}` was not found")]
448 Missing(&'static str),
449 #[error("tool context value `{key}` could not be encoded: {message}")]
451 Encode {
452 key: &'static str,
454 message: String,
456 },
457 #[error("tool context value `{key}` could not be decoded: {message}")]
459 Decode {
460 key: &'static str,
462 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
473const _: 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;