Files
codex/codex-rs/core/src/tools/registry_tests.rs
T
jif-oai 516f134641 Make tool executor specs mandatory (#23870)
## Why

`ToolExecutor` is the runtime contract that keeps a callable tool and
its model-visible spec together. Leaving `spec()` optional lets a
registered runtime silently omit that half of the contract, and it also
overloads a missing spec as an exposure decision for tools that should
stay dispatchable without being shown to the model.

## What

- Make `ToolExecutor::spec()` required and update core, extension, and
test tool executors to return a concrete `ToolSpec`.
- Add `ToolExposure::Hidden` for dispatch-only tools. The legacy
`shell_command` runtime in unified-exec sessions now uses that explicit
exposure instead of hiding itself by omitting a spec.
- Build MCP tool specs when `McpHandler` is constructed so invalid MCP
specs are skipped before the handler is registered.
- Keep tool planning aligned with the new contract for direct, deferred,
hidden, code-mode, dynamic, and namespaced tool paths.

## Testing

- Added tool-plan coverage that invalid MCP tool specs are not
registered.
- Updated shell-family coverage for the hidden legacy `shell_command`
runtime and the affected tool executor test fixtures.
2026-05-21 15:25:56 +02:00

277 lines
8.5 KiB
Rust

use super::*;
use pretty_assertions::assert_eq;
struct TestHandler {
tool_name: codex_tools::ToolName,
}
#[async_trait::async_trait]
impl ToolExecutor<ToolInvocation> for TestHandler {
fn tool_name(&self) -> codex_tools::ToolName {
self.tool_name.clone()
}
fn spec(&self) -> codex_tools::ToolSpec {
test_spec(&self.tool_name)
}
async fn handle(
&self,
_invocation: ToolInvocation,
) -> Result<Box<dyn crate::tools::context::ToolOutput>, FunctionCallError> {
Ok(Box::new(
crate::tools::context::FunctionToolOutput::from_text("ok".to_string(), Some(true)),
))
}
}
impl CoreToolRuntime for TestHandler {}
#[derive(Clone)]
enum LifecycleTestResult {
Ok { success: bool },
Err,
}
struct LifecycleTestHandler {
tool_name: codex_tools::ToolName,
result: LifecycleTestResult,
}
#[async_trait::async_trait]
impl ToolExecutor<ToolInvocation> for LifecycleTestHandler {
fn tool_name(&self) -> codex_tools::ToolName {
self.tool_name.clone()
}
fn spec(&self) -> codex_tools::ToolSpec {
test_spec(&self.tool_name)
}
async fn handle(
&self,
_invocation: ToolInvocation,
) -> Result<Box<dyn crate::tools::context::ToolOutput>, FunctionCallError> {
match self.result.clone() {
LifecycleTestResult::Ok { success } => Ok(Box::new(
crate::tools::context::FunctionToolOutput::from_text(
"ok".to_string(),
Some(success),
),
)),
LifecycleTestResult::Err => Err(FunctionCallError::RespondToModel(
"handler failed".to_string(),
)),
}
}
}
impl CoreToolRuntime for LifecycleTestHandler {}
fn test_spec(tool_name: &codex_tools::ToolName) -> codex_tools::ToolSpec {
codex_tools::ToolSpec::Function(codex_tools::ResponsesApiTool {
name: tool_name.name.clone(),
description: "Test tool.".to_string(),
strict: false,
defer_loading: None,
parameters: codex_tools::JsonSchema::default(),
output_schema: None,
})
}
#[derive(Debug, PartialEq, Eq)]
enum RecordedToolLifecycle {
Start {
call_id: String,
tool_name: codex_tools::ToolName,
},
Finish {
call_id: String,
tool_name: codex_tools::ToolName,
outcome: codex_extension_api::ToolCallOutcome,
},
}
struct ToolLifecycleRecorder {
records: Arc<std::sync::Mutex<Vec<RecordedToolLifecycle>>>,
}
impl codex_extension_api::ToolLifecycleContributor for ToolLifecycleRecorder {
fn on_tool_start<'a>(
&'a self,
input: codex_extension_api::ToolStartInput<'a>,
) -> codex_extension_api::ToolLifecycleFuture<'a> {
let records = Arc::clone(&self.records);
let record = RecordedToolLifecycle::Start {
call_id: input.call_id.to_string(),
tool_name: input.tool_name.clone(),
};
Box::pin(async move {
records
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.push(record);
})
}
fn on_tool_finish<'a>(
&'a self,
input: codex_extension_api::ToolFinishInput<'a>,
) -> codex_extension_api::ToolLifecycleFuture<'a> {
let records = Arc::clone(&self.records);
let record = RecordedToolLifecycle::Finish {
call_id: input.call_id.to_string(),
tool_name: input.tool_name.clone(),
outcome: input.outcome,
};
Box::pin(async move {
records
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.push(record);
})
}
}
#[test]
fn handler_looks_up_namespaced_aliases_explicitly() {
let namespace = "mcp__codex_apps__gmail";
let tool_name = "gmail_get_recent_emails";
let plain_name = codex_tools::ToolName::plain(tool_name);
let namespaced_name = codex_tools::ToolName::namespaced(namespace, tool_name);
let plain_handler = Arc::new(TestHandler {
tool_name: plain_name.clone(),
}) as Arc<dyn CoreToolRuntime>;
let namespaced_handler = Arc::new(TestHandler {
tool_name: namespaced_name.clone(),
}) as Arc<dyn CoreToolRuntime>;
let registry = ToolRegistry::new(HashMap::from([
(plain_name.clone(), Arc::clone(&plain_handler)),
(namespaced_name.clone(), Arc::clone(&namespaced_handler)),
]));
let plain = registry.tool(&plain_name);
let namespaced = registry.tool(&namespaced_name);
let missing_namespaced = registry.tool(&codex_tools::ToolName::namespaced(
"mcp__codex_apps__calendar",
tool_name,
));
assert_eq!(plain.is_some(), true);
assert_eq!(namespaced.is_some(), true);
assert_eq!(missing_namespaced.is_none(), true);
assert!(
plain
.as_ref()
.is_some_and(|handler| Arc::ptr_eq(handler, &plain_handler))
);
assert!(
namespaced
.as_ref()
.is_some_and(|handler| Arc::ptr_eq(handler, &namespaced_handler))
);
}
#[tokio::test]
async fn dispatch_notifies_tool_lifecycle_contributors() -> anyhow::Result<()> {
let (mut session, turn) = crate::session::tests::make_session_and_context().await;
let records = Arc::new(std::sync::Mutex::new(Vec::new()));
let mut builder = codex_extension_api::ExtensionRegistryBuilder::<crate::config::Config>::new();
builder.tool_lifecycle_contributor(Arc::new(ToolLifecycleRecorder {
records: Arc::clone(&records),
}));
session.services.extensions = Arc::new(builder.build());
let ok_tool = codex_tools::ToolName::plain("ok_tool");
let failing_tool = codex_tools::ToolName::plain("failing_tool");
let ok_handler = Arc::new(LifecycleTestHandler {
tool_name: ok_tool.clone(),
result: LifecycleTestResult::Ok { success: false },
}) as Arc<dyn CoreToolRuntime>;
let failing_handler = Arc::new(LifecycleTestHandler {
tool_name: failing_tool.clone(),
result: LifecycleTestResult::Err,
}) as Arc<dyn CoreToolRuntime>;
let registry = ToolRegistry::new(HashMap::from([
(ok_tool.clone(), ok_handler),
(failing_tool.clone(), failing_handler),
]));
let session = Arc::new(session);
let turn = Arc::new(turn);
registry
.dispatch_any(test_invocation(
Arc::clone(&session),
Arc::clone(&turn),
"ok-call",
ok_tool.clone(),
))
.await?;
let err = match registry
.dispatch_any(test_invocation(
Arc::clone(&session),
Arc::clone(&turn),
"failing-call",
failing_tool.clone(),
))
.await
{
Ok(_) => panic!("failing handler should return an error"),
Err(err) => err,
};
assert_eq!(err.to_string(), "handler failed");
let expected = vec![
RecordedToolLifecycle::Start {
call_id: "ok-call".to_string(),
tool_name: ok_tool.clone(),
},
RecordedToolLifecycle::Finish {
call_id: "ok-call".to_string(),
tool_name: ok_tool,
outcome: codex_extension_api::ToolCallOutcome::Completed { success: false },
},
RecordedToolLifecycle::Start {
call_id: "failing-call".to_string(),
tool_name: failing_tool.clone(),
},
RecordedToolLifecycle::Finish {
call_id: "failing-call".to_string(),
tool_name: failing_tool,
outcome: codex_extension_api::ToolCallOutcome::Failed {
handler_executed: true,
},
},
];
let actual = records
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.drain(..)
.collect::<Vec<_>>();
assert_eq!(expected, actual);
Ok(())
}
fn test_invocation(
session: Arc<crate::session::session::Session>,
turn: Arc<crate::session::turn_context::TurnContext>,
call_id: &str,
tool_name: codex_tools::ToolName,
) -> ToolInvocation {
ToolInvocation {
session,
turn,
cancellation_token: tokio_util::sync::CancellationToken::new(),
tracker: Arc::new(tokio::sync::Mutex::new(
crate::turn_diff_tracker::TurnDiffTracker::new(),
)),
call_id: call_id.to_string(),
tool_name,
source: crate::tools::context::ToolCallSource::Direct,
payload: ToolPayload::Function {
arguments: "{}".to_string(),
},
}
}