Fix tool vector tests (#3709)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2025-07-29 19:31:29 +02:00
committed by GitHub
parent 23ec0177bb
commit eb0048018f
2 changed files with 122 additions and 105 deletions
+57 -39
View File
@@ -382,33 +382,57 @@ pub fn generate_table_id() -> String {
mod tests { mod tests {
use super::*; use super::*;
impl ToolVectorDB {
async fn new_test_db(
base_name: &str,
) -> Result<(Self, impl std::future::Future<Output = ()>)> {
let unique_name = format!("{}_{}", base_name, uuid::Uuid::new_v4().simple());
let db = Self::new(Some(unique_name)).await?;
let table_name = db.table_name.clone();
let connection = db.connection.clone();
let cleanup = async move {
let _ = async move {
let _ = connection.read().await.drop_table(&table_name).await;
};
};
Ok((db, cleanup))
}
}
#[tokio::test] #[tokio::test]
#[serial_test::serial] #[serial_test::serial]
async fn test_tool_vectordb_creation() { async fn test_tool_vectordb_creation() -> Result<()> {
let db = ToolVectorDB::new(Some("test_tools_vectordb_creation".to_string())) let (db, cleanup) = ToolVectorDB::new_test_db("test_tools_vectordb_creation").await?;
.await
.unwrap(); let result = async {
db.clear_tools().await.unwrap(); db.clear_tools().await?;
assert_eq!(db.table_name, "test_tools_vectordb_creation"); assert!(db.table_name.contains("test_tools_vectordb_creation"));
Ok(())
}
.await;
cleanup.await;
result
} }
#[tokio::test] #[tokio::test]
#[serial_test::serial] #[serial_test::serial]
async fn test_tool_vectordb_operations() -> Result<()> { async fn test_tool_vectordb_operations() -> Result<()> {
// Create a new database instance with a unique table name let (db, cleanup) = ToolVectorDB::new_test_db("test_tool_vectordb_operations").await?;
let db = ToolVectorDB::new(Some("test_tool_vectordb_operations".to_string())).await?;
// Clear any existing tools let result = async {
db.clear_tools().await?; db.clear_tools().await?;
// Create test tool records
let test_tools = vec![ let test_tools = vec![
ToolRecord { ToolRecord {
tool_name: "test_tool_1".to_string(), tool_name: "test_tool_1".to_string(),
description: "A test tool for reading files".to_string(), description: "A test tool for reading files".to_string(),
schema: r#"{"type": "object", "properties": {"path": {"type": "string"}}}"# schema: r#"{"type": "object", "properties": {"path": {"type": "string"}}}"#
.to_string(), .to_string(),
vector: vec![0.1; 1536], // Mock embedding vector vector: vec![0.1; 1536],
extension_name: "test_extension".to_string(), extension_name: "test_extension".to_string(),
}, },
ToolRecord { ToolRecord {
@@ -416,19 +440,16 @@ mod tests {
description: "A test tool for writing files".to_string(), description: "A test tool for writing files".to_string(),
schema: r#"{"type": "object", "properties": {"path": {"type": "string"}}}"# schema: r#"{"type": "object", "properties": {"path": {"type": "string"}}}"#
.to_string(), .to_string(),
vector: vec![0.2; 1536], // Different mock embedding vector vector: vec![0.2; 1536],
extension_name: "test_extension".to_string(), extension_name: "test_extension".to_string(),
}, },
]; ];
// Index the test tools
db.index_tools(test_tools).await?; db.index_tools(test_tools).await?;
// Search for tools using a query vector similar to test_tool_1
let query_vector = vec![0.1; 1536]; let query_vector = vec![0.1; 1536];
let results = db.search_tools(query_vector.clone(), 2, None).await?; let results = db.search_tools(query_vector.clone(), 2, None).await?;
// Verify results
assert_eq!(results.len(), 2, "Should find both tools"); assert_eq!(results.len(), 2, "Should find both tools");
assert_eq!( assert_eq!(
results[0].tool_name, "test_tool_1", results[0].tool_name, "test_tool_1",
@@ -439,7 +460,6 @@ mod tests {
"Second result should be test_tool_2" "Second result should be test_tool_2"
); );
// Test filtering by extension name
let results = db let results = db
.search_tools(query_vector.clone(), 2, Some("test_extension")) .search_tools(query_vector.clone(), 2, Some("test_extension"))
.await?; .await?;
@@ -460,60 +480,67 @@ mod tests {
Ok(()) Ok(())
} }
.await;
cleanup.await;
result
}
#[tokio::test] #[tokio::test]
#[serial_test::serial] #[serial_test::serial]
async fn test_empty_db() -> Result<()> { async fn test_empty_db() -> Result<()> {
// Create a new database instance with a unique table name let (db, cleanup) = ToolVectorDB::new_test_db("test_empty_db").await?;
let db = ToolVectorDB::new(Some("test_empty_db".to_string())).await?;
// Clear any existing tools let result = async {
db.clear_tools().await?; db.clear_tools().await?;
// Search in empty database
let query_vector = vec![0.1; 1536]; let query_vector = vec![0.1; 1536];
let results = db.search_tools(query_vector, 2, None).await?; let results = db.search_tools(query_vector, 2, None).await?;
// Verify no results returned
assert_eq!(results.len(), 0, "Empty database should return no results"); assert_eq!(results.len(), 0, "Empty database should return no results");
Ok(()) Ok(())
} }
.await;
cleanup.await;
result
}
#[tokio::test] #[tokio::test]
#[serial_test::serial] #[serial_test::serial]
async fn test_tool_deletion() -> Result<()> { async fn test_tool_deletion() -> Result<()> {
// Create a new database instance with a unique table name let (db, cleanup) = ToolVectorDB::new_test_db("test_tool_deletion").await?;
let db = ToolVectorDB::new(Some("test_tool_deletion".to_string())).await?;
// Clear any existing tools let result = async {
db.clear_tools().await?; db.clear_tools().await?;
// Create and index a test tool
let test_tool = ToolRecord { let test_tool = ToolRecord {
tool_name: "test_tool_to_delete".to_string(), tool_name: "test_tool_to_delete".to_string(),
description: "A test tool that will be deleted".to_string(), description: "A test tool that will be deleted".to_string(),
schema: r#"{"type": "object", "properties": {"path": {"type": "string"}}}"#.to_string(), schema: r#"{"type": "object", "properties": {"path": {"type": "string"}}}"#
.to_string(),
vector: vec![0.1; 1536], vector: vec![0.1; 1536],
extension_name: "test_extension".to_string(), extension_name: "test_extension".to_string(),
}; };
db.index_tools(vec![test_tool]).await?; db.index_tools(vec![test_tool]).await?;
// Verify tool exists
let query_vector = vec![0.1; 1536]; let query_vector = vec![0.1; 1536];
let results = db.search_tools(query_vector.clone(), 1, None).await?; let results = db.search_tools(query_vector.clone(), 1, None).await?;
assert_eq!(results.len(), 1, "Tool should exist before deletion"); assert_eq!(results.len(), 1, "Tool should exist before deletion");
// Delete the tool
db.remove_tool("test_tool_to_delete").await?; db.remove_tool("test_tool_to_delete").await?;
// Verify tool is gone
let results = db.search_tools(query_vector.clone(), 1, None).await?; let results = db.search_tools(query_vector.clone(), 1, None).await?;
assert_eq!(results.len(), 0, "Tool should be deleted"); assert_eq!(results.len(), 0, "Tool should be deleted");
Ok(()) Ok(())
} }
.await;
cleanup.await;
result
}
#[test] #[test]
#[serial_test::serial] #[serial_test::serial]
@@ -521,20 +548,15 @@ mod tests {
use std::env; use std::env;
use tempfile::TempDir; use tempfile::TempDir;
// Create a temporary directory for testing
let temp_dir = TempDir::new().unwrap(); let temp_dir = TempDir::new().unwrap();
let custom_path = temp_dir.path().join("custom_vector_db"); let custom_path = temp_dir.path().join("custom_vector_db");
// Set the environment variable
env::set_var("GOOSE_VECTOR_DB_PATH", custom_path.to_str().unwrap()); env::set_var("GOOSE_VECTOR_DB_PATH", custom_path.to_str().unwrap());
// Test that get_db_path returns the custom path
let db_path = ToolVectorDB::get_db_path()?; let db_path = ToolVectorDB::get_db_path()?;
assert_eq!(db_path, custom_path); assert_eq!(db_path, custom_path);
// Clean up
env::remove_var("GOOSE_VECTOR_DB_PATH"); env::remove_var("GOOSE_VECTOR_DB_PATH");
Ok(()) Ok(())
} }
@@ -543,7 +565,6 @@ mod tests {
fn test_custom_db_path_validation() { fn test_custom_db_path_validation() {
use std::env; use std::env;
// Test that relative paths are rejected
env::set_var("GOOSE_VECTOR_DB_PATH", "relative/path"); env::set_var("GOOSE_VECTOR_DB_PATH", "relative/path");
let result = ToolVectorDB::get_db_path(); let result = ToolVectorDB::get_db_path();
@@ -557,7 +578,6 @@ mod tests {
.to_string() .to_string()
.contains("must be an absolute path")); .contains("must be an absolute path"));
// Clean up
env::remove_var("GOOSE_VECTOR_DB_PATH"); env::remove_var("GOOSE_VECTOR_DB_PATH");
} }
@@ -566,10 +586,8 @@ mod tests {
fn test_fallback_to_default_path() -> Result<()> { fn test_fallback_to_default_path() -> Result<()> {
use std::env; use std::env;
// Ensure no custom path is set
env::remove_var("GOOSE_VECTOR_DB_PATH"); env::remove_var("GOOSE_VECTOR_DB_PATH");
// Test that it falls back to default XDG path
let db_path = ToolVectorDB::get_db_path()?; let db_path = ToolVectorDB::get_db_path()?;
assert!( assert!(
db_path.to_string_lossy().contains("goose"), db_path.to_string_lossy().contains("goose"),
-1
View File
@@ -432,7 +432,6 @@ impl RecipeBuilder {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use std::fs;
#[test] #[test]
fn test_from_content_with_json() { fn test_from_content_with_json() {