use std::collections::HashSet;
use std::path::PathBuf;
use std::sync::Arc;
use talos_core::provider::ToolDefinition;
use talos_core::tool::{AgentTool, ToolPresentationPolicy, ToolProtocol, ToolRegistry};
use talos_permission::PermissionEngine;
use talos_plugin::HookRegistry;
use talos_sandbox::SandboxProvider;
use talos_skill::SkillIndex;
use tokio_util::sync::CancellationToken;
use crate::prompt::{ActivatedSkillContext, ContextFile, SystemPromptBuilder, ToolDescription};
use crate::{
Agent, MemoryProviderCallback, RequestBudgetSpec, SandboxFallbackHandler,
SandboxFallbackPolicy, TodoSectionProviderCallback, prompt,
};
impl Agent {
#[deprecated(
note = "Agent::new() has NO permission engine and NO sandbox; use Agent::with_security(). See docs/decisions/007-process-hardening-unsafe.md and ARCH review."
)]
#[must_use]
pub fn new(
provider: Arc<dyn talos_core::provider::LanguageModel>,
tools: ToolRegistry,
) -> Self {
Self {
provider,
tools,
permission_engine: None,
sandbox: None,
sandbox_fallback_policy: SandboxFallbackPolicy::Deny,
sandbox_fallback_handler: None,
workspace_root: PathBuf::from("."),
prompt_builder: SystemPromptBuilder::new().with_workspace_info("Workspace root: ."),
hook_registry: Arc::new(HookRegistry::new()),
workspace_context: None,
tool_definitions: Vec::new(),
presented_tool_names: HashSet::new(),
enforce_tool_presentation_policy: false,
tool_presentation_policy: ToolPresentationPolicy::full(),
cached_stable_prefix: std::sync::Mutex::new(None),
memory_provider: None,
todo_section_provider: None,
provider_key: None,
model_id: None,
replay_reasoning: true,
bash_compression_enabled: false,
tool_output_threshold: 4000,
image_input_supported: false,
request_budget_spec: RequestBudgetSpec::default(),
}
}
#[must_use]
pub fn with_security(
provider: Arc<dyn talos_core::provider::LanguageModel>,
tools: ToolRegistry,
permission_engine: Option<Arc<PermissionEngine>>,
sandbox: Option<Box<dyn SandboxProvider>>,
workspace_root: PathBuf,
) -> Self {
Self::with_security_and_hooks(
provider,
tools,
permission_engine,
sandbox,
workspace_root,
Arc::new(HookRegistry::new()),
)
}
#[must_use]
pub fn with_security_and_sandbox_fallback(
provider: Arc<dyn talos_core::provider::LanguageModel>,
tools: ToolRegistry,
permission_engine: Option<Arc<PermissionEngine>>,
sandbox: Option<Box<dyn SandboxProvider>>,
workspace_root: PathBuf,
sandbox_fallback_policy: SandboxFallbackPolicy,
sandbox_fallback_handler: Option<Arc<dyn SandboxFallbackHandler>>,
) -> Self {
Self::with_security_and_hooks_and_sandbox_fallback(
provider,
tools,
permission_engine,
sandbox,
workspace_root,
Arc::new(HookRegistry::new()),
sandbox_fallback_policy,
sandbox_fallback_handler,
)
}
#[must_use]
pub fn with_security_and_hooks(
provider: Arc<dyn talos_core::provider::LanguageModel>,
tools: ToolRegistry,
permission_engine: Option<Arc<PermissionEngine>>,
sandbox: Option<Box<dyn SandboxProvider>>,
workspace_root: PathBuf,
hook_registry: Arc<HookRegistry>,
) -> Self {
Self::with_security_and_hooks_and_sandbox_fallback(
provider,
tools,
permission_engine,
sandbox,
workspace_root,
hook_registry,
SandboxFallbackPolicy::Deny,
None,
)
}
#[must_use]
#[allow(clippy::too_many_arguments)]
pub fn with_security_and_hooks_and_sandbox_fallback(
provider: Arc<dyn talos_core::provider::LanguageModel>,
tools: ToolRegistry,
permission_engine: Option<Arc<PermissionEngine>>,
sandbox: Option<Box<dyn SandboxProvider>>,
workspace_root: PathBuf,
hook_registry: Arc<HookRegistry>,
sandbox_fallback_policy: SandboxFallbackPolicy,
sandbox_fallback_handler: Option<Arc<dyn SandboxFallbackHandler>>,
) -> Self {
let tool_presentation_policy = ToolPresentationPolicy::runtime_default();
let (descriptions, tool_definitions, presented_tool_names) =
describe_presented_tools(&tools, &tool_presentation_policy);
let descriptions: Vec<_> = descriptions
.into_iter()
.filter(|d| d.name != "read_image")
.collect();
let tool_definitions: Vec<_> = tool_definitions
.into_iter()
.filter(|td| td.name != "read_image")
.collect();
let presented_tool_names: HashSet<_> = presented_tool_names
.into_iter()
.filter(|n| n != "read_image")
.collect();
let prompt_builder = SystemPromptBuilder::new()
.with_workspace_info(format!("Workspace root: {}", workspace_root.display()))
.with_tools(descriptions.clone());
Self {
provider,
tools,
permission_engine,
sandbox: sandbox.map(Arc::from),
sandbox_fallback_policy,
sandbox_fallback_handler,
workspace_root,
prompt_builder,
hook_registry,
workspace_context: None,
tool_definitions,
presented_tool_names,
enforce_tool_presentation_policy: true,
tool_presentation_policy,
cached_stable_prefix: std::sync::Mutex::new(None),
memory_provider: None,
todo_section_provider: None,
provider_key: None,
model_id: None,
replay_reasoning: true,
bash_compression_enabled: false,
tool_output_threshold: 4000,
image_input_supported: false,
request_budget_spec: RequestBudgetSpec::default(),
}
}
#[must_use]
pub fn with_reasoning_identity(
mut self,
provider_key: Option<String>,
model_id: Option<String>,
replay: bool,
) -> Self {
self.provider_key = provider_key;
self.model_id = model_id;
self.replay_reasoning = replay;
self
}
pub fn set_request_budget_spec(&mut self, spec: RequestBudgetSpec) {
self.request_budget_spec = spec;
}
#[must_use]
pub fn request_budget_spec(&self) -> RequestBudgetSpec {
self.request_budget_spec
}
pub fn set_memory_provider(&mut self, provider: Arc<MemoryProviderCallback>) {
self.memory_provider = Some(provider);
}
pub fn set_todo_section_provider(&mut self, provider: Arc<TodoSectionProviderCallback>) {
self.todo_section_provider = Some(provider);
}
#[must_use]
pub fn with_bash_compression(mut self, enabled: bool) -> Self {
self.bash_compression_enabled = enabled;
self
}
#[must_use]
pub fn with_image_input_supported(mut self, supported: bool) -> Self {
self.image_input_supported = supported;
self
}
pub fn set_image_input_supported(&mut self, supported: bool) {
self.image_input_supported = supported;
let (descs, defs, names) =
describe_presented_tools(&self.tools, &self.tool_presentation_policy);
let descs: Vec<_> = descs
.into_iter()
.filter(|d| supported || d.name != "read_image")
.collect();
self.tool_definitions = defs
.into_iter()
.filter(|td| supported || td.name != "read_image")
.collect();
self.presented_tool_names = names
.into_iter()
.filter(|n| supported || n != "read_image")
.collect();
self.enforce_tool_presentation_policy = true;
self.update_prompt_builder(true, |builder| builder.with_tools(descs));
}
pub fn set_tools(&mut self, tools: Vec<ToolDescription>) {
self.tool_definitions = tools
.iter()
.map(|tool| ToolDefinition {
name: tool.name.clone(),
description: tool.description.clone(),
parameters: tool.parameters.clone(),
})
.collect();
self.presented_tool_names = tools.iter().map(|tool| tool.name.clone()).collect();
self.enforce_tool_presentation_policy = true;
self.update_prompt_builder(true, |builder| builder.with_tools(tools));
}
pub fn set_tool_presentation_policy(&mut self, policy: ToolPresentationPolicy) {
self.tool_presentation_policy = policy;
let (descriptions, tool_definitions, presented_tool_names) =
describe_presented_tools(&self.tools, &self.tool_presentation_policy);
self.tool_definitions = tool_definitions;
self.presented_tool_names = presented_tool_names;
self.enforce_tool_presentation_policy = true;
self.update_prompt_builder(true, |builder| builder.with_tools(descriptions));
}
pub fn set_tool_protocol(&mut self, protocol: ToolProtocol) {
self.update_prompt_builder(true, |builder| match protocol {
ToolProtocol::TalosStrict => builder.with_strict_tool_format(),
ToolProtocol::Compat => builder.with_tool_format(prompt::TOOL_CALLING_FORMAT),
ToolProtocol::Native => builder.with_tool_format(""),
});
}
pub fn set_skill_index(&mut self, skills: Vec<SkillIndex>) {
self.update_prompt_builder(true, |builder| builder.with_skill_index(skills));
}
pub fn set_activated_skill_context(&mut self, context: Option<ActivatedSkillContext>) {
self.update_prompt_builder(true, |builder| builder.with_activated_skill(context));
}
pub fn set_context_files(&mut self, files: Vec<ContextFile>) {
self.update_prompt_builder(false, |builder| builder.with_context_files(files));
}
pub fn set_user_preferences(&mut self, prefs: String) {
self.update_prompt_builder(false, |builder| builder.with_user_preferences(prefs));
}
pub fn set_custom_prompt(&mut self, prompt: String) {
self.update_prompt_builder(true, |builder| builder.with_custom_prompt(prompt));
}
pub fn set_append_prompt(&mut self, prompt: String) {
self.update_prompt_builder(false, |builder| builder.with_append_prompt(prompt));
}
pub fn clear_append_prompt(&mut self) {
self.prompt_builder.clear_append_prompt();
}
pub fn set_append_prompt_opt(&mut self, prompt: Option<String>) {
self.prompt_builder.set_append_prompt_opt(prompt);
}
#[must_use]
pub fn build_system_prompt(&self) -> String {
self.prompt_builder.build()
}
#[must_use]
pub fn cancellation_token(&self) -> CancellationToken {
CancellationToken::new()
}
fn update_prompt_builder(
&mut self,
invalidate_stable_prefix: bool,
update: impl FnOnce(SystemPromptBuilder) -> SystemPromptBuilder,
) {
self.prompt_builder = update(std::mem::take(&mut self.prompt_builder));
if invalidate_stable_prefix {
self.invalidate_stable_prefix_cache();
}
}
pub(crate) fn invalidate_stable_prefix_cache(&self) {
*self
.cached_stable_prefix
.lock()
.expect("cache lock poisoned") = None;
}
}
pub(crate) fn describe_presented_tools(
tools: &ToolRegistry,
policy: &ToolPresentationPolicy,
) -> (Vec<ToolDescription>, Vec<ToolDefinition>, HashSet<String>) {
let mut selected: Vec<&dyn AgentTool> = tools
.list()
.into_iter()
.filter(|tool| policy.allows_tool(*tool))
.collect();
selected.sort_by(|a, b| a.name().cmp(b.name()));
let descriptions: Vec<ToolDescription> = selected
.iter()
.map(|tool| {
let backends = policy.backend_set_for(tool.name());
ToolDescription {
name: tool.name().to_string(),
description: tool.description_for_backends(&backends),
parameters: tool.parameters_for_backends(&backends),
family: tool.family(),
}
})
.collect();
let tool_definitions = descriptions
.iter()
.map(|tool| ToolDefinition {
name: tool.name.clone(),
description: tool.description.clone(),
parameters: tool.parameters.clone(),
})
.collect();
let presented_tool_names = descriptions.iter().map(|tool| tool.name.clone()).collect();
(descriptions, tool_definitions, presented_tool_names)
}