use std::any::{Any, TypeId};
use std::collections::{BTreeMap, HashMap};
use std::hash::{BuildHasherDefault, Hasher};
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use super::ToolExecutionError;
use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
type AnyMap = HashMap<TypeId, Box<dyn AnyClone>, BuildHasherDefault<IdHasher>>;
#[derive(Default)]
struct IdHasher(u64);
impl Hasher for IdHasher {
fn write_u64(&mut self, id: u64) {
self.0 = id;
}
fn write(&mut self, bytes: &[u8]) {
for &byte in bytes {
self.0 = self.0.rotate_left(8) ^ u64::from(byte);
}
}
fn finish(&self) -> u64 {
self.0
}
}
trait AnyClone: Any + WasmCompatSend + WasmCompatSync {
fn clone_box(&self) -> Box<dyn AnyClone>;
fn as_any(&self) -> &dyn Any;
fn as_any_mut(&mut self) -> &mut dyn Any;
fn into_any(self: Box<Self>) -> Box<dyn Any>;
}
impl<T> AnyClone for T
where
T: Clone + WasmCompatSend + WasmCompatSync + 'static,
{
fn clone_box(&self) -> Box<dyn AnyClone> {
Box::new(self.clone())
}
fn as_any(&self) -> &dyn Any {
self
}
fn as_any_mut(&mut self) -> &mut dyn Any {
self
}
fn into_any(self: Box<Self>) -> Box<dyn Any> {
self
}
}
impl Clone for Box<dyn AnyClone> {
fn clone(&self) -> Self {
(**self).clone_box()
}
}
#[derive(Default, Clone)]
pub struct TypeMap {
map: AnyMap,
}
impl TypeMap {
pub fn insert<T>(&mut self, value: T) -> Option<T>
where
T: Clone + WasmCompatSend + WasmCompatSync + 'static,
{
self.map
.insert(TypeId::of::<T>(), Box::new(value))
.and_then(|previous| previous.into_any().downcast::<T>().ok())
.map(|value| *value)
}
pub fn get<T>(&self) -> Option<&T>
where
T: 'static,
{
self.map
.get(&TypeId::of::<T>())
.and_then(|value| (**value).as_any().downcast_ref::<T>())
}
pub fn get_mut<T>(&mut self) -> Option<&mut T>
where
T: 'static,
{
self.map
.get_mut(&TypeId::of::<T>())
.and_then(|value| (**value).as_any_mut().downcast_mut::<T>())
}
pub fn remove<T>(&mut self) -> Option<T>
where
T: 'static,
{
self.map
.remove(&TypeId::of::<T>())
.and_then(|value| value.into_any().downcast::<T>().ok())
.map(|value| *value)
}
pub fn contains<T>(&self) -> bool
where
T: 'static,
{
self.map.contains_key(&TypeId::of::<T>())
}
pub fn len(&self) -> usize {
self.map.len()
}
pub fn is_empty(&self) -> bool {
self.map.is_empty()
}
}
#[derive(Default, Clone, Serialize, Deserialize)]
pub struct ToolContext {
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
inbound: BTreeMap<String, serde_json::Value>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
result: BTreeMap<String, serde_json::Value>,
#[serde(skip)]
scopes: Vec<std::sync::Arc<dyn Any + Send + Sync>>,
}
#[derive(Default, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(transparent)]
pub struct ToolResultContext(BTreeMap<String, serde_json::Value>);
impl std::fmt::Debug for ToolResultContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_tuple("ToolResultContext")
.field(&self.0.keys().collect::<Vec<_>>())
.finish()
}
}
impl ToolResultContext {
pub fn get<T: ContextValue>(&self) -> Result<Option<T>, ToolContextError> {
get_slot(&self.0)
}
}
impl PartialEq for ToolContext {
fn eq(&self, other: &Self) -> bool {
self.inbound == other.inbound && self.result == other.result
}
}
impl Eq for ToolContext {}
#[diagnostic::on_unimplemented(
message = "`{Self}` declares no `ToolContext` key",
label = "not a `ContextValue`",
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"
)]
pub trait ContextValue: Serialize + DeserializeOwned + 'static {
const KEY: &'static str;
}
fn encode<T: ContextValue>(value: &T) -> Result<serde_json::Value, ToolContextError> {
serde_json::to_value(value).map_err(|error| ToolContextError::Encode {
key: T::KEY,
source: error,
})
}
fn decode<T: ContextValue>(value: &serde_json::Value) -> Result<T, ToolContextError> {
serde_json::from_value(value.clone()).map_err(|error| ToolContextError::Decode {
key: T::KEY,
source: error,
})
}
fn insert_slot<T: ContextValue>(
map: &mut BTreeMap<String, serde_json::Value>,
value: T,
) -> Result<Option<T>, ToolContextError> {
let encoded = encode(&value)?;
Ok(map
.insert(T::KEY.to_owned(), encoded)
.and_then(|previous| decode(&previous).ok()))
}
fn get_slot<T: ContextValue>(
map: &BTreeMap<String, serde_json::Value>,
) -> Result<Option<T>, ToolContextError> {
map.get(T::KEY).map(decode).transpose()
}
fn require_slot<T: ContextValue>(
map: &BTreeMap<String, serde_json::Value>,
) -> Result<T, ToolContextError> {
get_slot(map)?.ok_or(ToolContextError::Missing(T::KEY))
}
impl ToolContext {
pub const fn new() -> Self {
Self {
inbound: BTreeMap::new(),
result: BTreeMap::new(),
scopes: Vec::new(),
}
}
pub fn insert<T: ContextValue>(&mut self, value: T) -> Result<Option<T>, ToolContextError> {
insert_slot(&mut self.inbound, value)
}
pub fn get<T: ContextValue>(&self) -> Result<Option<T>, ToolContextError> {
get_slot(&self.inbound)
}
pub fn require<T: ContextValue>(&self) -> Result<T, ToolContextError> {
require_slot(&self.inbound)
}
pub fn remove<T: ContextValue>(&mut self) -> Result<Option<T>, ToolContextError> {
self.inbound
.remove(T::KEY)
.map(|value| decode(&value))
.transpose()
}
pub fn insert_result<T: ContextValue>(
&mut self,
value: T,
) -> Result<Option<T>, ToolContextError> {
insert_slot(&mut self.result, value)
}
pub fn result<T: ContextValue>(&self) -> Result<Option<T>, ToolContextError> {
get_slot(&self.result)
}
pub fn require_result<T: ContextValue>(&self) -> Result<T, ToolContextError> {
require_slot(&self.result)
}
pub fn contains<T: ContextValue>(&self) -> bool {
self.inbound.contains_key(T::KEY)
}
pub fn is_empty(&self) -> bool {
self.inbound.is_empty() && self.result.is_empty()
}
pub fn for_dispatch(&self) -> Self {
Self {
inbound: self.inbound.clone(),
result: BTreeMap::new(),
scopes: self.scopes.clone(),
}
}
pub fn with_scope(mut self, scope: std::sync::Arc<dyn Any + Send + Sync>) -> Self {
self.scopes.push(scope);
self
}
pub fn with_scopes(mut self, scopes: Vec<std::sync::Arc<dyn Any + Send + Sync>>) -> Self {
self.scopes.extend(scopes);
self
}
pub fn scope<T: Any + Send + Sync>(&self) -> Option<std::sync::Arc<T>> {
self.scopes
.iter()
.find_map(|scope| std::sync::Arc::downcast::<T>(scope.clone()).ok())
}
pub fn clear_scope(&mut self) {
self.scopes.clear();
}
pub fn accept_dispatch_result(&mut self, dispatched: Self) {
self.result = dispatched.result;
}
pub fn result_context(&self) -> ToolResultContext {
ToolResultContext(self.result.clone())
}
pub fn with_result_context(mut self, result: ToolResultContext) -> Self {
self.result = result.0;
self
}
pub fn clear_dispatch_result(&mut self) {
self.result.clear();
}
}
impl std::fmt::Debug for ToolContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ToolContext")
.field("inbound_types", &self.inbound.keys().collect::<Vec<_>>())
.field("result_types", &self.result.keys().collect::<Vec<_>>())
.finish()
}
}
#[derive(Debug, Default)]
pub struct PublishedContext(std::sync::Mutex<Option<ToolContext>>);
impl PublishedContext {
pub fn new() -> std::sync::Arc<Self> {
std::sync::Arc::new(Self::default())
}
pub fn result_context(&self) -> Option<ToolResultContext> {
self.0
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.as_ref()
.map(ToolContext::result_context)
}
pub fn publish(&self, mut context: ToolContext) {
context.clear_scope();
*self
.0
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = Some(context);
}
pub fn take(&self) -> Option<ToolContext> {
self.0
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take()
}
}
#[derive(Debug, thiserror::Error)]
pub enum ToolContextError {
#[error("required tool context value `{0}` was not found")]
Missing(&'static str),
#[error("tool context value `{key}` could not be encoded: {source}")]
Encode {
key: &'static str,
#[source]
source: serde_json::Error,
},
#[error("tool context value `{key}` could not be decoded: {source}")]
Decode {
key: &'static str,
#[source]
source: serde_json::Error,
},
}
impl From<ToolContextError> for ToolExecutionError {
fn from(error: ToolContextError) -> Self {
ToolExecutionError::other(error.to_string()).with_source(error)
}
}
const _: fn() = || {
fn assert_wire<T: Send + Sync + 'static + Serialize + DeserializeOwned>() {}
assert_wire::<ToolContext>();
fn assert_shared<T: Send + Sync + 'static>() {}
assert_shared::<PublishedContext>();
};
#[cfg(test)]
mod tests;
#[cfg(test)]
mod migrated_tests;