chore: refactor read-write lock on agent (#2225)

Co-authored-by: Alice Hau <ahau@squareup.com>
This commit is contained in:
Salman Mohammed
2025-04-23 23:46:22 -03:00
committed by GitHub
parent 85e2ee3984
commit 199fa6adbc
24 changed files with 409 additions and 237 deletions
+2 -1
View File
@@ -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",
+62 -36
View File
@@ -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?;
+2 -2
View File
@@ -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());
+26 -18
View File
@@ -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);
+1 -1
View File
@@ -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(),
+4 -7
View File
@@ -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);
+2 -1
View File
@@ -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![