use anda_core::{BoxError, FunctionDefinition, StateFeatures, ToolOutput, gen_schema_for};
use schemars::JsonSchema;
use std::future::Future;
use crate::{
context::BaseCtx,
hook::{DynToolHook, ToolHook},
};
pub mod fetch;
pub mod fs;
pub mod mcp;
pub mod note;
pub mod shell;
pub mod skill;
pub mod todo;
pub mod workspace;
pub fn tool_definition<A: JsonSchema>(name: String, description: String) -> FunctionDefinition {
FunctionDefinition {
name,
description,
parameters: gen_schema_for::<A>(),
strict: Some(true),
}
}
pub async fn hooked_call<I, O, F, Fut>(
ctx: &BaseCtx,
args: I,
run: F,
) -> Result<ToolOutput<O>, BoxError>
where
I: Send + Sync + 'static,
O: Send + Sync + 'static,
F: FnOnce(I) -> Fut,
Fut: Future<Output = Result<ToolOutput<O>, BoxError>> + Send,
{
if ctx.cancellation_token().is_cancelled() {
return Err("call was cancelled".into());
}
let hook = ctx.get_state::<DynToolHook<I, O>>();
let args = match &hook {
Some(hook) => hook.before_tool_call(ctx, args).await?,
None => args,
};
let output = run(args).await?;
match &hook {
Some(hook) => hook.after_tool_call(ctx, output).await,
None => Ok(output),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::engine::EngineBuilder;
use crate::extension::{
fetch::FetchWebResourcesTool,
fs::{EditFileTool, ReadFileTool, SearchFileTool, WriteFileTool},
note::NoteTool,
shell::{ExecArgs, ExecOutput, Executor, ShellTool},
skill::{SkillManager, SkillsListTool, SkillsReadTool},
todo::TodoTool,
};
use anda_core::{Json, Tool};
use async_trait::async_trait;
use serde_json::json;
use std::{collections::HashMap, path::PathBuf, sync::Arc};
struct StubRuntime {
workspace: PathBuf,
}
#[async_trait]
impl Executor for StubRuntime {
fn name(&self) -> &str {
"stub"
}
fn workspace(&self) -> &PathBuf {
&self.workspace
}
fn shell(&self) -> &str {
"sh"
}
async fn execute(
&self,
_ctx: BaseCtx,
_input: ExecArgs,
_envs: HashMap<String, String>,
) -> Result<ExecOutput, BoxError> {
unreachable!("definition-only test")
}
}
fn extension_tool_definitions() -> Vec<FunctionDefinition> {
let dir = std::env::temp_dir();
vec![
FetchWebResourcesTool::new().definition(),
ReadFileTool::new(dir.clone()).definition(),
EditFileTool::new(dir.clone()).definition(),
WriteFileTool::new(dir.clone()).definition(),
SearchFileTool::new(dir.clone()).definition(),
fs::ApplyPatchTool::new(dir.clone()).definition(),
shell::ShellCommandTool::new(
Arc::new(StubRuntime {
workspace: dir.clone(),
}),
vec![],
)
.definition(),
shell::ShellSessionTool::new(Arc::new(StubRuntime {
workspace: dir.clone(),
}))
.definition(),
TodoTool::new().definition(),
NoteTool::new().definition(),
ShellTool::new(
Arc::new(StubRuntime {
workspace: dir.clone(),
}),
HashMap::new(),
None,
)
.definition(),
SkillManager::new(dir.clone()).definition(),
SkillsListTool::new(Arc::new(SkillManager::new(dir.clone()))).definition(),
SkillsReadTool::new(Arc::new(SkillManager::new(dir))).definition(),
]
}
fn assert_union_free(tool: &str, value: &Json) {
match value {
Json::Object(map) => {
for key in ["anyOf", "oneOf", "allOf"] {
assert!(
!map.contains_key(key),
"tool {tool}: derived schema contains {key}"
);
}
for child in map.values() {
assert_union_free(tool, child);
}
}
Json::Array(items) => {
for child in items {
assert_union_free(tool, child);
}
}
_ => {}
}
}
#[test]
fn extension_tool_definitions_are_strict_and_union_free() {
let definitions = extension_tool_definitions();
assert_eq!(definitions.len(), 14);
for definition in definitions {
let tool = &definition.name;
assert_eq!(definition.strict, Some(true), "tool {tool}");
assert_eq!(
definition.parameters["type"],
json!("object"),
"tool {tool}"
);
assert_eq!(
definition.parameters["additionalProperties"],
json!(false),
"tool {tool}"
);
let mut properties: Vec<&str> = definition.parameters["properties"]
.as_object()
.unwrap_or_else(|| panic!("tool {tool}: missing properties object"))
.keys()
.map(String::as_str)
.collect();
let mut required: Vec<&str> = definition.parameters["required"]
.as_array()
.unwrap_or_else(|| panic!("tool {tool}: missing required array"))
.iter()
.map(|value| value.as_str().unwrap())
.collect();
properties.sort_unstable();
required.sort_unstable();
assert_eq!(
required, properties,
"tool {tool}: strict schemas require every property"
);
assert_union_free(tool, &definition.parameters);
}
}
#[test]
fn nullable_enum_args_derive_flat_type_unions() {
let todo = TodoTool::new().definition().parameters;
assert_eq!(todo["properties"]["op"]["type"], json!(["string", "null"]));
assert_eq!(
todo["properties"]["op"]["enum"],
json!(["read", "set", "update", null])
);
assert_eq!(todo["properties"]["op"]["default"], json!("read"));
assert_eq!(
todo["properties"]["items"]["type"],
json!(["array", "null"])
);
assert_eq!(
todo["properties"]["items"]["items"]["properties"]["status"]["enum"],
json!(["pending", "in_progress", "completed", "cancelled", null])
);
let note = NoteTool::new().definition().parameters;
assert_eq!(
note["properties"]["op"]["enum"],
json!(["read", "list", "search", "set", "upsert", "delete", null])
);
assert_eq!(
note["properties"]["items"]["items"]["properties"]["content"]["type"],
json!(["string", "null"])
);
}
#[test]
fn instance_dependent_schema_patches_apply() {
let shell = ShellTool::new(
Arc::new(StubRuntime {
workspace: std::env::temp_dir(),
}),
HashMap::new(),
None,
)
.definition()
.parameters;
assert!(
shell["properties"]["env_keys"]["description"]
.as_str()
.unwrap()
.contains("environment variable")
);
assert!(
shell["properties"]["background"]["description"]
.as_str()
.unwrap()
.contains("background")
);
let skill = SkillManager::new(std::env::temp_dir())
.definition()
.parameters;
assert!(
skill["description"]
.as_str()
.unwrap()
.contains("SKILL.md file content")
);
}
struct RewritingHook;
#[async_trait]
impl ToolHook<String, String> for RewritingHook {
async fn before_tool_call(&self, _ctx: &BaseCtx, args: String) -> Result<String, BoxError> {
Ok(format!("{args}+before"))
}
async fn after_tool_call(
&self,
_ctx: &BaseCtx,
output: ToolOutput<String>,
) -> Result<ToolOutput<String>, BoxError> {
Ok(ToolOutput::new(format!("{}+after", output.output)))
}
}
#[tokio::test(flavor = "current_thread")]
async fn hooked_call_applies_the_hook_and_gates_on_cancellation() {
let ctx = EngineBuilder::new().mock_ctx().base;
let output = hooked_call(&ctx, "args".to_string(), |args| async move {
Ok(ToolOutput::new(format!("{args}+run")))
})
.await
.unwrap();
assert_eq!(output.output, "args+run");
ctx.set_state(DynToolHook::new(
Arc::new(RewritingHook) as Arc<dyn ToolHook<String, String>>
));
let output = hooked_call(&ctx, "args".to_string(), |args| async move {
Ok(ToolOutput::new(format!("{args}+run")))
})
.await
.unwrap();
assert_eq!(output.output, "args+before+run+after");
ctx.cancellation_token().cancel();
let err = hooked_call(&ctx, "args".to_string(), |_args: String| async move {
unreachable!("the body must not run on a cancelled context")
as Result<ToolOutput<String>, BoxError>
})
.await
.unwrap_err();
assert!(err.to_string().contains("cancelled"));
}
}