chore: refactor read-write lock on agent (#2225)
Co-authored-by: Alice Hau <ahau@squareup.com>
This commit is contained in:
@@ -15,7 +15,8 @@ async fn main() {
|
||||
let provider = Arc::new(DatabricksProvider::default());
|
||||
|
||||
// Setup an agent with the developer extension
|
||||
let mut agent = Agent::new(provider);
|
||||
let agent = Agent::new();
|
||||
let _ = agent.update_provider(provider).await;
|
||||
|
||||
let config = ExtensionConfig::stdio(
|
||||
"developer",
|
||||
|
||||
@@ -35,12 +35,11 @@ use super::tool_execution::{ToolFuture, CHAT_MODE_TOOL_SKIPPED_RESPONSE, DECLINE
|
||||
|
||||
/// The main goose Agent
|
||||
pub struct Agent {
|
||||
pub(super) provider: Arc<dyn Provider>,
|
||||
pub(super) provider: Mutex<Option<Arc<dyn Provider>>>,
|
||||
pub(super) extension_manager: Mutex<ExtensionManager>,
|
||||
pub(super) frontend_tools: HashMap<String, FrontendTool>,
|
||||
pub(super) frontend_instructions: Option<String>,
|
||||
pub(super) prompt_manager: PromptManager,
|
||||
// Channels for tool results and confirmations
|
||||
pub(super) frontend_tools: Mutex<HashMap<String, FrontendTool>>,
|
||||
pub(super) frontend_instructions: Mutex<Option<String>>,
|
||||
pub(super) prompt_manager: Mutex<PromptManager>,
|
||||
pub(super) confirmation_tx: mpsc::Sender<(String, PermissionConfirmation)>,
|
||||
pub(super) confirmation_rx: Mutex<mpsc::Receiver<(String, PermissionConfirmation)>>,
|
||||
pub(super) tool_result_tx: mpsc::Sender<(String, ToolResult<Vec<Content>>)>,
|
||||
@@ -48,41 +47,52 @@ pub struct Agent {
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
pub fn new(provider: Arc<dyn Provider>) -> Self {
|
||||
pub fn new() -> Self {
|
||||
// Create channels with buffer size 32 (adjust if needed)
|
||||
let (confirm_tx, confirm_rx) = mpsc::channel(32);
|
||||
let (tool_tx, tool_rx) = mpsc::channel(32);
|
||||
|
||||
Self {
|
||||
provider,
|
||||
provider: Mutex::new(None),
|
||||
extension_manager: Mutex::new(ExtensionManager::new()),
|
||||
frontend_tools: HashMap::new(),
|
||||
frontend_instructions: None,
|
||||
prompt_manager: PromptManager::new(),
|
||||
frontend_tools: Mutex::new(HashMap::new()),
|
||||
frontend_instructions: Mutex::new(None),
|
||||
prompt_manager: Mutex::new(PromptManager::new()),
|
||||
confirmation_tx: confirm_tx,
|
||||
confirmation_rx: Mutex::new(confirm_rx),
|
||||
tool_result_tx: tool_tx,
|
||||
tool_result_rx: Arc::new(Mutex::new(tool_rx)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for Agent {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
/// Get a reference count clone to the provider
|
||||
pub fn provider(&self) -> Arc<dyn Provider> {
|
||||
Arc::clone(&self.provider)
|
||||
pub async fn provider(&self) -> Result<Arc<dyn Provider>, anyhow::Error> {
|
||||
match &*self.provider.lock().await {
|
||||
Some(provider) => Ok(Arc::clone(provider)),
|
||||
None => Err(anyhow!("Provider not set")),
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if a tool is a frontend tool
|
||||
pub fn is_frontend_tool(&self, name: &str) -> bool {
|
||||
self.frontend_tools.contains_key(name)
|
||||
pub async fn is_frontend_tool(&self, name: &str) -> bool {
|
||||
self.frontend_tools.lock().await.contains_key(name)
|
||||
}
|
||||
|
||||
/// Get a reference to a frontend tool
|
||||
pub fn get_frontend_tool(&self, name: &str) -> Option<&FrontendTool> {
|
||||
self.frontend_tools.get(name)
|
||||
pub async fn get_frontend_tool(&self, name: &str) -> Option<FrontendTool> {
|
||||
self.frontend_tools.lock().await.get(name).cloned()
|
||||
}
|
||||
|
||||
/// Get all tools from all clients with proper prefixing
|
||||
pub async fn get_prefixed_tools(&mut self) -> ExtensionResult<Vec<Tool>> {
|
||||
pub async fn get_prefixed_tools(&self) -> ExtensionResult<Vec<Tool>> {
|
||||
let mut tools = self
|
||||
.extension_manager
|
||||
.lock()
|
||||
@@ -91,7 +101,8 @@ impl Agent {
|
||||
.await?;
|
||||
|
||||
// Add frontend tools directly - they don't need prefixing since they're already uniquely named
|
||||
for frontend_tool in self.frontend_tools.values() {
|
||||
let frontend_tools = self.frontend_tools.lock().await;
|
||||
for frontend_tool in frontend_tools.values() {
|
||||
tools.push(frontend_tool.tool.clone());
|
||||
}
|
||||
|
||||
@@ -135,7 +146,7 @@ impl Agent {
|
||||
.await
|
||||
} else if tool_call.name == PLATFORM_SEARCH_AVAILABLE_EXTENSIONS_TOOL_NAME {
|
||||
extension_manager.search_available_extensions().await
|
||||
} else if self.is_frontend_tool(&tool_call.name) {
|
||||
} else if self.is_frontend_tool(&tool_call.name).await {
|
||||
// For frontend tools, return an error indicating we need frontend execution
|
||||
Err(ToolError::ExecutionError(
|
||||
"Frontend tool execution required".to_string(),
|
||||
@@ -212,7 +223,7 @@ impl Agent {
|
||||
(request_id, result)
|
||||
}
|
||||
|
||||
pub async fn add_extension(&mut self, extension: ExtensionConfig) -> ExtensionResult<()> {
|
||||
pub async fn add_extension(&self, extension: ExtensionConfig) -> ExtensionResult<()> {
|
||||
match &extension {
|
||||
ExtensionConfig::Frontend {
|
||||
name: _,
|
||||
@@ -221,19 +232,21 @@ impl Agent {
|
||||
bundled: _,
|
||||
} => {
|
||||
// For frontend tools, just store them in the frontend_tools map
|
||||
let mut frontend_tools = self.frontend_tools.lock().await;
|
||||
for tool in tools {
|
||||
let frontend_tool = FrontendTool {
|
||||
name: tool.name.clone(),
|
||||
tool: tool.clone(),
|
||||
};
|
||||
self.frontend_tools.insert(tool.name.clone(), frontend_tool);
|
||||
frontend_tools.insert(tool.name.clone(), frontend_tool);
|
||||
}
|
||||
// Store instructions if provided, using "frontend" as the key
|
||||
let mut frontend_instructions = self.frontend_instructions.lock().await;
|
||||
if let Some(instructions) = instructions {
|
||||
self.frontend_instructions = Some(instructions.clone());
|
||||
*frontend_instructions = Some(instructions.clone());
|
||||
} else {
|
||||
// Default frontend instructions if none provided
|
||||
self.frontend_instructions = Some(
|
||||
*frontend_instructions = Some(
|
||||
"The following tools are provided directly by the frontend and will be executed by the frontend when called.".to_string(),
|
||||
);
|
||||
}
|
||||
@@ -269,7 +282,7 @@ impl Agent {
|
||||
prefixed_tools
|
||||
}
|
||||
|
||||
pub async fn remove_extension(&mut self, name: &str) {
|
||||
pub async fn remove_extension(&self, name: &str) {
|
||||
let mut extension_manager = self.extension_manager.lock().await;
|
||||
extension_manager
|
||||
.remove_extension(name)
|
||||
@@ -329,7 +342,7 @@ impl Agent {
|
||||
let _ = reply_span.enter();
|
||||
loop {
|
||||
match Self::generate_response_from_provider(
|
||||
self.provider(),
|
||||
self.provider().await?,
|
||||
&system_prompt,
|
||||
&messages,
|
||||
&tools,
|
||||
@@ -345,7 +358,7 @@ impl Agent {
|
||||
let (frontend_requests,
|
||||
remaining_requests,
|
||||
filtered_response) =
|
||||
self.categorize_tool_requests(&response);
|
||||
self.categorize_tool_requests(&response).await;
|
||||
|
||||
|
||||
// Yield the assistant's response with frontend tool requests filtered out
|
||||
@@ -396,8 +409,7 @@ impl Agent {
|
||||
tools_with_readonly_annotation.clone(),
|
||||
tools_without_annotation.clone(),
|
||||
&mut permission_manager,
|
||||
self.provider(),
|
||||
).await;
|
||||
self.provider().await?).await;
|
||||
|
||||
// Handle pre-approved and read-only tools in parallel
|
||||
let mut tool_futures: Vec<ToolFuture> = Vec::new();
|
||||
@@ -492,13 +504,21 @@ impl Agent {
|
||||
}
|
||||
|
||||
/// Extend the system prompt with one line of additional instruction
|
||||
pub async fn extend_system_prompt(&mut self, instruction: String) {
|
||||
self.prompt_manager.add_system_prompt_extra(instruction);
|
||||
pub async fn extend_system_prompt(&self, instruction: String) {
|
||||
let mut prompt_manager = self.prompt_manager.lock().await;
|
||||
prompt_manager.add_system_prompt_extra(instruction);
|
||||
}
|
||||
|
||||
/// Update the provider used by this agent
|
||||
pub async fn update_provider(&self, provider: Arc<dyn Provider>) -> Result<()> {
|
||||
*self.provider.lock().await = Some(provider);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Override the system prompt with a custom template
|
||||
pub async fn override_system_prompt(&mut self, template: String) {
|
||||
self.prompt_manager.set_system_prompt_override(template);
|
||||
pub async fn override_system_prompt(&self, template: String) {
|
||||
let mut prompt_manager = self.prompt_manager.lock().await;
|
||||
prompt_manager.set_system_prompt_override(template);
|
||||
}
|
||||
|
||||
pub async fn list_extension_prompts(&self) -> HashMap<String, Vec<Prompt>> {
|
||||
@@ -563,23 +583,29 @@ impl Agent {
|
||||
let extensions_info = extension_manager.get_extensions_info().await;
|
||||
|
||||
// Get model name from provider
|
||||
let model_config = self.provider.get_model_config();
|
||||
let provider = self.provider().await?;
|
||||
let model_config = provider.get_model_config();
|
||||
let model_name = &model_config.model_name;
|
||||
|
||||
let system_prompt = self.prompt_manager.build_system_prompt(
|
||||
let prompt_manager = self.prompt_manager.lock().await;
|
||||
let system_prompt = prompt_manager.build_system_prompt(
|
||||
extensions_info,
|
||||
self.frontend_instructions.clone(),
|
||||
self.frontend_instructions.lock().await.clone(),
|
||||
extension_manager.suggest_disable_extensions_prompt().await,
|
||||
Some(model_name),
|
||||
);
|
||||
|
||||
let recipe_prompt = self.prompt_manager.get_recipe_prompt().await;
|
||||
let recipe_prompt = prompt_manager.get_recipe_prompt().await;
|
||||
let tools = extension_manager.get_prefixed_tools(None).await?;
|
||||
|
||||
messages.push(Message::user().with_text(recipe_prompt));
|
||||
|
||||
let (result, _usage) = self
|
||||
.provider
|
||||
.lock()
|
||||
.await
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.complete(&system_prompt, &messages, &tools)
|
||||
.await?;
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ impl Agent {
|
||||
&self,
|
||||
messages: &[Message], // last message is a user msg that led to assistant message with_context_length_exceeded
|
||||
) -> Result<(Vec<Message>, Vec<usize>), anyhow::Error> {
|
||||
let provider = self.provider.clone();
|
||||
let provider = self.provider().await?;
|
||||
let token_counter = TokenCounter::new(provider.get_model_config().tokenizer_name());
|
||||
let target_context_limit = estimate_target_context_limit(provider);
|
||||
let token_counts = get_messages_token_counts(&token_counter, messages);
|
||||
@@ -41,7 +41,7 @@ impl Agent {
|
||||
&self,
|
||||
messages: &[Message], // last message is a user msg that led to assistant message with_context_length_exceeded
|
||||
) -> Result<(Vec<Message>, Vec<usize>), anyhow::Error> {
|
||||
let provider = self.provider.clone();
|
||||
let provider = self.provider().await?;
|
||||
let token_counter = TokenCounter::new(provider.get_model_config().tokenizer_name());
|
||||
let target_context_limit = estimate_target_context_limit(provider.clone());
|
||||
|
||||
|
||||
@@ -22,7 +22,8 @@ impl Agent {
|
||||
let mut tools = self.list_tools(None).await;
|
||||
|
||||
// Add frontend tools
|
||||
for frontend_tool in self.frontend_tools.values() {
|
||||
let frontend_tools = self.frontend_tools.lock().await;
|
||||
for frontend_tool in frontend_tools.values() {
|
||||
tools.push(frontend_tool.tool.clone());
|
||||
}
|
||||
|
||||
@@ -31,19 +32,21 @@ impl Agent {
|
||||
let extensions_info = extension_manager.get_extensions_info().await;
|
||||
|
||||
// Get model name from provider
|
||||
let model_config = self.provider.get_model_config();
|
||||
let provider = self.provider().await?;
|
||||
let model_config = provider.get_model_config();
|
||||
let model_name = &model_config.model_name;
|
||||
|
||||
let mut system_prompt = self.prompt_manager.build_system_prompt(
|
||||
let prompt_manager = self.prompt_manager.lock().await;
|
||||
let mut system_prompt = prompt_manager.build_system_prompt(
|
||||
extensions_info,
|
||||
self.frontend_instructions.clone(),
|
||||
self.frontend_instructions.lock().await.clone(),
|
||||
extension_manager.suggest_disable_extensions_prompt().await,
|
||||
Some(model_name),
|
||||
);
|
||||
|
||||
// Handle toolshim if enabled
|
||||
let mut toolshim_tools = vec![];
|
||||
if self.provider.get_model_config().toolshim {
|
||||
if model_config.toolshim {
|
||||
// If tool interpretation is enabled, modify the system prompt
|
||||
system_prompt = modify_system_prompt_for_tool_json(&system_prompt, &tools);
|
||||
// Make a copy of tools before emptying
|
||||
@@ -115,7 +118,7 @@ impl Agent {
|
||||
/// - frontend_requests: Tool requests that should be handled by the frontend
|
||||
/// - other_requests: All other tool requests (including requests to enable extensions)
|
||||
/// - filtered_message: The original message with frontend tool requests removed
|
||||
pub(crate) fn categorize_tool_requests(
|
||||
pub(crate) async fn categorize_tool_requests(
|
||||
&self,
|
||||
response: &Message,
|
||||
) -> (Vec<ToolRequest>, Vec<ToolRequest>, Message) {
|
||||
@@ -133,20 +136,25 @@ impl Agent {
|
||||
.collect();
|
||||
|
||||
// Create a filtered message with frontend tool requests removed
|
||||
let filtered_content = response
|
||||
.content
|
||||
.iter()
|
||||
.filter(|c| {
|
||||
if let MessageContent::ToolRequest(req) = c {
|
||||
// Only filter out frontend tool requests
|
||||
let mut filtered_content = Vec::new();
|
||||
|
||||
// Process each content item one by one
|
||||
for content in &response.content {
|
||||
let should_include = match content {
|
||||
MessageContent::ToolRequest(req) => {
|
||||
if let Ok(tool_call) = &req.tool_call {
|
||||
return !self.is_frontend_tool(&tool_call.name);
|
||||
!self.is_frontend_tool(&tool_call.name).await
|
||||
} else {
|
||||
true
|
||||
}
|
||||
}
|
||||
true
|
||||
})
|
||||
.cloned()
|
||||
.collect();
|
||||
_ => true,
|
||||
};
|
||||
|
||||
if should_include {
|
||||
filtered_content.push(content.clone());
|
||||
}
|
||||
}
|
||||
|
||||
let filtered_message = Message {
|
||||
role: response.role.clone(),
|
||||
@@ -160,7 +168,7 @@ impl Agent {
|
||||
|
||||
for request in tool_requests {
|
||||
if let Ok(tool_call) = &request.tool_call {
|
||||
if self.is_frontend_tool(&tool_call.name) {
|
||||
if self.is_frontend_tool(&tool_call.name).await {
|
||||
frontend_requests.push(request);
|
||||
} else {
|
||||
other_requests.push(request);
|
||||
|
||||
@@ -87,7 +87,7 @@ impl Agent {
|
||||
try_stream! {
|
||||
for request in tool_requests {
|
||||
if let Ok(tool_call) = request.tool_call.clone() {
|
||||
if self.is_frontend_tool(&tool_call.name) {
|
||||
if self.is_frontend_tool(&tool_call.name).await {
|
||||
// Send frontend tool request and wait for response
|
||||
yield Message::assistant().with_frontend_tool_request(
|
||||
request.id.clone(),
|
||||
|
||||
@@ -739,10 +739,10 @@ mod tests {
|
||||
thread::sleep(Duration::from_millis(i * 10));
|
||||
|
||||
let extension_key = format!("extension_{}", i);
|
||||
let mut values = config.load_values()?;
|
||||
|
||||
values.insert(
|
||||
extension_key.clone(),
|
||||
// Use set_param which handles concurrent access properly
|
||||
config.set_param(
|
||||
&extension_key,
|
||||
serde_json::json!({
|
||||
"name": format!("test_extension_{}", i),
|
||||
"version": format!("1.0.{}", i),
|
||||
@@ -752,10 +752,7 @@ mod tests {
|
||||
"option2": i
|
||||
}
|
||||
}),
|
||||
);
|
||||
|
||||
// Write all values atomically
|
||||
config.save_values(values)?;
|
||||
)?;
|
||||
Ok(())
|
||||
});
|
||||
handles.push(handle);
|
||||
|
||||
@@ -110,7 +110,8 @@ async fn run_truncate_test(
|
||||
.with_temperature(Some(0.0));
|
||||
let provider = provider_type.create_provider(model_config)?;
|
||||
|
||||
let agent = Agent::new(provider);
|
||||
let agent = Agent::new();
|
||||
agent.update_provider(provider).await?;
|
||||
let repeat_count = context_window + 10_000;
|
||||
let large_message_content = "hello ".repeat(repeat_count);
|
||||
let messages = vec![
|
||||
|
||||
Reference in New Issue
Block a user