Fix tool vector tests (#3709)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -382,137 +382,164 @@ 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],
|
||||||
vector: vec![0.1; 1536], // Mock embedding vector
|
extension_name: "test_extension".to_string(),
|
||||||
extension_name: "test_extension".to_string(),
|
},
|
||||||
},
|
ToolRecord {
|
||||||
ToolRecord {
|
tool_name: "test_tool_2".to_string(),
|
||||||
tool_name: "test_tool_2".to_string(),
|
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],
|
||||||
vector: vec![0.2; 1536], // Different mock embedding vector
|
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",
|
"First result should be test_tool_1"
|
||||||
"First result should be test_tool_1"
|
);
|
||||||
);
|
assert_eq!(
|
||||||
assert_eq!(
|
results[1].tool_name, "test_tool_2",
|
||||||
results[1].tool_name, "test_tool_2",
|
"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?;
|
assert_eq!(
|
||||||
assert_eq!(
|
results.len(),
|
||||||
results.len(),
|
2,
|
||||||
2,
|
"Should find both tools with test_extension"
|
||||||
"Should find both tools with test_extension"
|
);
|
||||||
);
|
|
||||||
|
|
||||||
let results = db
|
let results = db
|
||||||
.search_tools(query_vector.clone(), 2, Some("nonexistent_extension"))
|
.search_tools(query_vector.clone(), 2, Some("nonexistent_extension"))
|
||||||
.await?;
|
.await?;
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
results.len(),
|
results.len(),
|
||||||
0,
|
0,
|
||||||
"Should find no tools with nonexistent_extension"
|
"Should find no tools with nonexistent_extension"
|
||||||
);
|
);
|
||||||
|
|
||||||
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(())
|
||||||
|
}
|
||||||
|
.await;
|
||||||
|
|
||||||
Ok(())
|
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"}}}"#
|
||||||
schema: r#"{"type": "object", "properties": {"path": {"type": "string"}}}"#.to_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]
|
||||||
@@ -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"),
|
||||||
|
|||||||
@@ -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() {
|
||||||
|
|||||||
Reference in New Issue
Block a user