fix: update index when tool selection strategy changes (#2991)

This commit is contained in:
Wendy Tang
2025-06-23 11:45:54 -07:00
committed by GitHub
parent ce7eabf20d
commit cebdbdb3d2
5 changed files with 213 additions and 24 deletions
+39 -5
View File
@@ -842,12 +842,23 @@ impl Agent {
/// Update the provider used by this agent
pub async fn update_provider(&self, provider: Arc<dyn Provider>) -> Result<()> {
*self.provider.lock().await = Some(provider.clone());
self.update_router_tool_selector(provider).await?;
self.update_router_tool_selector(Some(provider), None)
.await?;
Ok(())
}
async fn update_router_tool_selector(&self, provider: Arc<dyn Provider>) -> Result<()> {
pub async fn update_router_tool_selector(
&self,
provider: Option<Arc<dyn Provider>>,
reindex_all: Option<bool>,
) -> Result<()> {
let config = Config::global();
let extension_manager = self.extension_manager.lock().await;
let provider = match provider {
Some(p) => p,
None => self.provider().await?,
};
let router_tool_selection_strategy = config
.get_param("GOOSE_ROUTER_TOOL_SELECTION_STRATEGY")
.unwrap_or_else(|_| "default".to_string());
@@ -861,21 +872,44 @@ impl Agent {
let selector = match strategy {
Some(RouterToolSelectionStrategy::Vector) => {
let table_name = generate_table_id();
let selector = create_tool_selector(strategy, provider, Some(table_name))
let selector = create_tool_selector(strategy, provider.clone(), Some(table_name))
.await
.map_err(|e| anyhow!("Failed to create tool selector: {}", e))?;
Arc::new(selector)
}
Some(RouterToolSelectionStrategy::Llm) => {
let selector = create_tool_selector(strategy, provider, None)
let selector = create_tool_selector(strategy, provider.clone(), None)
.await
.map_err(|e| anyhow!("Failed to create tool selector: {}", e))?;
Arc::new(selector)
}
None => return Ok(()),
};
let extension_manager = self.extension_manager.lock().await;
// First index platform tools
ToolRouterIndexManager::index_platform_tools(&selector, &extension_manager).await?;
if reindex_all.unwrap_or(false) {
let enabled_extensions = extension_manager.list_extensions().await?;
for extension_name in enabled_extensions {
if let Err(e) = ToolRouterIndexManager::update_extension_tools(
&selector,
&extension_manager,
&extension_name,
"add",
)
.await
{
tracing::error!(
"Failed to index tools for extension {}: {}",
extension_name,
e
);
}
}
}
// Update the selector
*self.router_tool_selector.lock().await = Some(selector.clone());
Ok(())
}