use super::*;
#[async_trait]
pub trait DynamicToolProvider: Send + Sync {
fn start(&self);
fn direct_tools(&self) -> Vec<Arc<dyn Tool>>;
fn available_definitions(&self) -> Vec<ToolDefinition>;
fn contains(&self, name: &str) -> bool {
self.available_definitions()
.iter()
.any(|definition| definition.name() == name)
}
fn supports_parallel_tool_calls(&self, _name: &str) -> bool {
false
}
async fn execute(
&self,
name: &str,
input: Value,
context: ToolContext<'_>,
) -> Option<ToolOutput>;
}
#[derive(Clone)]
pub struct Tools {
workspace: bool,
web_search: bool,
image_generation: bool,
pub(super) working_directory: Option<Arc<str>>,
pub(super) default_shell: Option<Arc<str>>,
process_environment: Arc<Vec<(OsString, OsString)>>,
remote_http_client: Option<reqwest::Client>,
pub(super) registered: Vec<Arc<dyn Tool>>,
pub(super) providers: Vec<Arc<dyn DynamicToolProvider>>,
}
impl Default for Tools {
fn default() -> Self {
Self {
workspace: true,
web_search: true,
image_generation: true,
working_directory: None,
default_shell: None,
process_environment: Arc::new(Vec::new()),
remote_http_client: None,
registered: Vec::new(),
providers: Vec::new(),
}
}
}
impl fmt::Debug for Tools {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
let remote_http_client_configured = self.remote_http_client.is_some();
formatter
.debug_struct("Tools")
.field("workspace", &self.workspace)
.field("web_search", &self.web_search)
.field("image_generation", &self.image_generation)
.field("working_directory", &self.working_directory)
.field("default_shell", &self.default_shell)
.field("process_environment_count", &self.process_environment.len())
.field(
"remote_http_client_configured",
&remote_http_client_configured,
)
.field(
"registered",
&self
.registered
.iter()
.map(|tool| tool.definition().name().to_owned())
.collect::<Vec<_>>(),
)
.field("provider_count", &self.providers.len())
.finish()
}
}
impl Tools {
#[must_use]
pub fn builder() -> ToolsBuilder {
ToolsBuilder::default()
}
#[must_use]
pub const fn into_builder(self) -> ToolsBuilder {
ToolsBuilder { tools: self }
}
#[must_use]
pub const fn workspace_enabled(&self) -> bool {
self.workspace
}
#[must_use]
pub const fn web_search_enabled(&self) -> bool {
self.web_search
}
#[must_use]
pub const fn image_generation_enabled(&self) -> bool {
self.image_generation
}
#[must_use]
pub fn for_session(mut self, session_id: &str) -> Self {
self.insert_process_environment(CODEX_THREAD_ID_ENV_VAR.into(), session_id.into());
self
}
pub(super) fn process_environment(&self) -> Arc<Vec<(OsString, OsString)>> {
Arc::clone(&self.process_environment)
}
fn insert_process_environment(&mut self, name: OsString, value: OsString) {
let environment = Arc::make_mut(&mut self.process_environment);
environment.retain(|(candidate, _)| candidate != &name);
environment.push((name, value));
}
pub(super) fn remote_http_client(&self) -> Option<reqwest::Client> {
self.remote_http_client.clone()
}
pub fn start_providers(&self) {
for provider in &self.providers {
provider.start();
}
}
}
#[derive(Default)]
pub struct ToolsBuilder {
tools: Tools,
}
#[derive(Debug, thiserror::Error)]
pub enum ToolsBuildError {
#[error("tool name must not be empty")]
EmptyName,
#[error("working directory override must not be empty")]
EmptyWorkingDirectory,
#[error("default shell override must not be empty")]
EmptyDefaultShell,
#[error("tool name `{0}` is registered more than once")]
DuplicateName(Box<str>),
#[error("tool name `{0}` conflicts with an enabled built-in tool")]
BuiltInName(Box<str>),
}
impl ToolsBuilder {
#[must_use]
pub const fn without_defaults(mut self) -> Self {
self.tools.workspace = false;
self.tools.web_search = false;
self.tools.image_generation = false;
self
}
#[must_use]
pub const fn workspace(mut self, enabled: bool) -> Self {
self.tools.workspace = enabled;
self
}
#[must_use]
pub const fn web_search(mut self, enabled: bool) -> Self {
self.tools.web_search = enabled;
self
}
#[must_use]
pub const fn image_generation(mut self, enabled: bool) -> Self {
self.tools.image_generation = enabled;
self
}
#[must_use]
pub fn working_directory(mut self, directory: impl Into<Arc<str>>) -> Self {
self.tools.working_directory = Some(directory.into());
self
}
#[must_use]
pub fn default_shell(mut self, shell: impl Into<Arc<str>>) -> Self {
self.tools.default_shell = Some(shell.into());
self
}
#[must_use]
pub fn process_environment<I, K, V>(mut self, variables: I) -> Self
where
I: IntoIterator<Item = (K, V)>,
K: Into<OsString>,
V: Into<OsString>,
{
for (name, value) in variables {
self.tools
.insert_process_environment(name.into(), value.into());
}
self
}
#[must_use]
pub fn remote_http_client(mut self, client: reqwest::Client) -> Self {
self.tools.remote_http_client = Some(client);
self
}
#[must_use]
pub fn tool<T: Tool + 'static>(mut self, tool: T) -> Self {
self.tools.registered.push(Arc::new(tool));
self
}
#[must_use]
pub fn provider<P: DynamicToolProvider + 'static>(mut self, provider: P) -> Self {
let provider: Arc<dyn DynamicToolProvider> = Arc::new(provider);
self.tools.registered.extend(provider.direct_tools());
self.tools.providers.push(provider);
self
}
pub fn build(self) -> Result<Tools, ToolsBuildError> {
if self
.tools
.working_directory
.as_deref()
.is_some_and(|directory| directory.trim().is_empty())
{
return Err(ToolsBuildError::EmptyWorkingDirectory);
}
if self
.tools
.default_shell
.as_deref()
.is_some_and(|shell| shell.trim().is_empty())
{
return Err(ToolsBuildError::EmptyDefaultShell);
}
let mut names = HashSet::with_capacity(self.tools.registered.len());
for tool in &self.tools.registered {
let definition = tool.definition();
let name = definition.name();
if name.is_empty() {
return Err(ToolsBuildError::EmptyName);
}
if built_in_name(&self.tools, name) {
return Err(ToolsBuildError::BuiltInName(name.into()));
}
if !names.insert(name.to_owned()) {
return Err(ToolsBuildError::DuplicateName(name.into()));
}
}
Ok(self.tools)
}
}
fn built_in_name(tools: &Tools, name: &str) -> bool {
(tools.workspace
&& matches!(
name,
"exec_command" | "write_stdin" | "update_plan" | "apply_patch" | "view_image"
))
|| (tools.web_search && name == "web__run")
|| (tools.image_generation && name == "image_gen__imagegen")
}