662 lines
22 KiB
Rust
662 lines
22 KiB
Rust
use crate::models::*;
|
|
use crate::router::McpTool;
|
|
use crate::state::MemoryState;
|
|
use crate::tools::*;
|
|
use async_trait::async_trait;
|
|
use serde_json::Value;
|
|
use std::sync::Arc;
|
|
|
|
pub struct PinFileHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for PinFileHandler {
|
|
fn name(&self) -> &'static str {
|
|
"pin_file"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<PinFileTool>("pin_file", "Execute pin_file")
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: PinFileTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
state.pinned_files.modify(|pinned| {
|
|
pinned.retain(|p| p.namespace != req.namespace || p.file_path != req.file_path);
|
|
pinned.push(crate::models::PinnedFile {
|
|
namespace: req.namespace,
|
|
file_path: req.file_path,
|
|
timestamp: crate::handlers::utils::now_secs(),
|
|
git_branch: req.git_branch,
|
|
});
|
|
});
|
|
Ok("File pinned".to_string())
|
|
}
|
|
}
|
|
|
|
pub struct UnpinFileHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for UnpinFileHandler {
|
|
fn name(&self) -> &'static str {
|
|
"unpin_file"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<UnpinFileTool>("unpin_file", "Execute unpin_file")
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: UnpinFileTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
state.pinned_files.modify(|pinned| {
|
|
pinned.retain(|p| p.namespace != req.namespace || p.file_path != req.file_path)
|
|
});
|
|
Ok("File unpinned".to_string())
|
|
}
|
|
}
|
|
|
|
pub struct ListPinnedFilesHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for ListPinnedFilesHandler {
|
|
fn name(&self) -> &'static str {
|
|
"list_pinned_files"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<ListPinnedFilesTool>(
|
|
"list_pinned_files",
|
|
"Execute list_pinned_files",
|
|
)
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: ListPinnedFilesTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
let data = state.pinned_files.read_with(|pinned| {
|
|
let filtered: Vec<_> = pinned
|
|
.iter()
|
|
.filter(|p| {
|
|
let ns_match = match &req.namespace {
|
|
Some(ns) => &p.namespace == ns,
|
|
std::option::Option::None => true,
|
|
};
|
|
let branch_match = match &req.git_branch {
|
|
Some(branch) => {
|
|
p.git_branch.is_none()
|
|
|| p.git_branch.as_deref() == Some(branch.as_str())
|
|
}
|
|
std::option::Option::None => true,
|
|
};
|
|
ns_match && branch_match
|
|
})
|
|
.collect();
|
|
serde_json::to_string(&filtered).map_err(|e| e.to_string())
|
|
})?;
|
|
Ok(data)
|
|
}
|
|
}
|
|
|
|
pub struct StoreSnippetHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for StoreSnippetHandler {
|
|
fn name(&self) -> &'static str {
|
|
"store_snippet"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<StoreSnippetTool>("store_snippet", "Execute store_snippet")
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: StoreSnippetTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
let req_name = req.name.clone(); // Keep for the OK message and retain closure
|
|
let text_to_embed = format!("Name: {}\nLanguage: {}\nDescription: {}\nCode: {}", req.name, req.language, req.description, req.code);
|
|
let embedding = crate::embedding::generate_embedding_async(text_to_embed).await.ok();
|
|
let snippet = Snippet {
|
|
name: req.name,
|
|
language: req.language,
|
|
code: req.code,
|
|
description: req.description,
|
|
updated_at: crate::handlers::utils::now_secs(),
|
|
embedding,
|
|
};
|
|
|
|
let idx = state.get_search_index();
|
|
drop(idx.index_snippet(&snippet));
|
|
|
|
state.snippets.modify(|snippets| {
|
|
snippets.retain(|s| s.name != req_name);
|
|
snippets.push(snippet);
|
|
});
|
|
|
|
Ok(format!("Snippet '{}' stored.", req_name).to_string())
|
|
}
|
|
}
|
|
|
|
pub struct SearchSnippetsHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for SearchSnippetsHandler {
|
|
fn name(&self) -> &'static str {
|
|
"search_snippets"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<SearchSnippetsTool>("search_snippets", "Execute search_snippets")
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: SearchSnippetsTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
let query = req.query;
|
|
let data = state.snippets.read_with(|snippets| {
|
|
let results: Vec<_> = snippets
|
|
.iter()
|
|
.filter(|s| {
|
|
contains_ignore_ascii_case(&s.name, &query)
|
|
|| contains_ignore_ascii_case(&s.description, &query)
|
|
|| contains_ignore_ascii_case(&s.language, &query)
|
|
})
|
|
.collect();
|
|
serde_json::to_string(&results).map_err(|e| e.to_string())
|
|
})?;
|
|
Ok(data)
|
|
}
|
|
}
|
|
|
|
pub struct DeleteSnippetHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for DeleteSnippetHandler {
|
|
fn name(&self) -> &'static str {
|
|
"delete_snippet"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<DeleteSnippetTool>("delete_snippet", "Execute delete_snippet")
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: DeleteSnippetTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
let mut deleted = false;
|
|
state.snippets.modify(|snippets| {
|
|
let orig = snippets.len();
|
|
snippets.retain(|s| s.name != req.name);
|
|
deleted = snippets.len() < orig;
|
|
});
|
|
if deleted {
|
|
let idx = state.get_search_index();
|
|
drop(idx.delete_document(&req.name));
|
|
Ok("Snippet deleted.".to_string())
|
|
} else {
|
|
Err(
|
|
"Snippet not found. Please verify the snippet ID using search_snippets."
|
|
.to_string(),
|
|
)
|
|
}
|
|
}
|
|
}
|
|
|
|
pub struct SaveContextWorkspaceHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for SaveContextWorkspaceHandler {
|
|
fn name(&self) -> &'static str {
|
|
"save_context_workspace"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<SaveContextWorkspaceTool>(
|
|
"save_context_workspace",
|
|
"Execute save_context_workspace",
|
|
)
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: SaveContextWorkspaceTool =
|
|
serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
state.context_workspaces.modify(|ws| {
|
|
ws.retain(|w| w.namespace != req.namespace || w.name != req.name);
|
|
ws.push(crate::models::ContextWorkspace {
|
|
namespace: req.namespace,
|
|
name: req.name,
|
|
pinned_files: req.pinned_files,
|
|
active_task_ids: req.active_task_ids,
|
|
saved_at: crate::handlers::utils::now_secs(),
|
|
});
|
|
});
|
|
Ok("Context workspace saved".to_string())
|
|
}
|
|
}
|
|
|
|
pub struct LoadContextWorkspaceHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for LoadContextWorkspaceHandler {
|
|
fn name(&self) -> &'static str {
|
|
"load_context_workspace"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<LoadContextWorkspaceTool>(
|
|
"load_context_workspace",
|
|
"Execute load_context_workspace",
|
|
)
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: LoadContextWorkspaceTool =
|
|
serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
let data = state.context_workspaces.read_with(|ws| {
|
|
let filtered: Vec<_> = ws
|
|
.iter()
|
|
.filter(|w| w.namespace == req.namespace && w.name == req.name)
|
|
.collect();
|
|
serde_json::to_string(&filtered.first()).map_err(|e| e.to_string())
|
|
})?;
|
|
Ok(data)
|
|
}
|
|
}
|
|
|
|
pub struct ListContextWorkspacesHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for ListContextWorkspacesHandler {
|
|
fn name(&self) -> &'static str {
|
|
"list_context_workspaces"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<ListContextWorkspacesTool>(
|
|
"list_context_workspaces",
|
|
"Execute list_context_workspaces",
|
|
)
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: ListContextWorkspacesTool =
|
|
serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
let data = state.context_workspaces.read_with(|ws| {
|
|
let filtered: Vec<_> = ws
|
|
.iter()
|
|
.filter(|w| req.namespace.as_ref().is_none_or(|ns| &w.namespace == ns))
|
|
.collect();
|
|
serde_json::to_string(&filtered).map_err(|e| e.to_string())
|
|
})?;
|
|
Ok(data)
|
|
}
|
|
}
|
|
|
|
pub struct DeleteContextWorkspaceHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for DeleteContextWorkspaceHandler {
|
|
fn name(&self) -> &'static str {
|
|
"delete_context_workspace"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<crate::tools::DeleteContextWorkspaceTool>(
|
|
"delete_context_workspace",
|
|
"Delete a saved context workspace",
|
|
)
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: crate::tools::DeleteContextWorkspaceTool =
|
|
serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
|
|
let mut found = false;
|
|
state.context_workspaces.modify(|ws| {
|
|
if let Some(pos) = ws
|
|
.iter()
|
|
.position(|w| w.namespace == req.namespace && w.name == req.name)
|
|
{
|
|
ws.remove(pos);
|
|
found = true;
|
|
}
|
|
});
|
|
|
|
if found {
|
|
Ok("Context workspace deleted successfully".to_string())
|
|
} else {
|
|
Err("Context workspace not found".to_string())
|
|
}
|
|
}
|
|
}
|
|
|
|
pub struct AddPrChecklistItemHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for AddPrChecklistItemHandler {
|
|
fn name(&self) -> &'static str {
|
|
"add_pr_checklist_item"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<AddPrChecklistItemTool>(
|
|
"add_pr_checklist_item",
|
|
"Execute add_pr_checklist_item",
|
|
)
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: AddPrChecklistItemTool =
|
|
serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
state.pr_checklists.modify(|items| {
|
|
items.push(crate::models::PrChecklistItem {
|
|
namespace: req.namespace,
|
|
id: uuid::Uuid::new_v4().to_string(),
|
|
description: req.description,
|
|
})
|
|
});
|
|
Ok("PR checklist item added".to_string())
|
|
}
|
|
}
|
|
|
|
pub struct GetPrChecklistHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for GetPrChecklistHandler {
|
|
fn name(&self) -> &'static str {
|
|
"get_pr_checklist"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<GetPrChecklistTool>("get_pr_checklist", "Execute get_pr_checklist")
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: GetPrChecklistTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
let data = state.pr_checklists.read_with(|items| {
|
|
let filtered: Vec<_> = items
|
|
.iter()
|
|
.filter(|i| i.namespace == req.namespace)
|
|
.collect();
|
|
serde_json::to_string(&filtered).map_err(|e| e.to_string())
|
|
})?;
|
|
Ok(data)
|
|
}
|
|
}
|
|
|
|
pub struct ClearPrChecklistHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for ClearPrChecklistHandler {
|
|
fn name(&self) -> &'static str {
|
|
"clear_pr_checklist"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<ClearPrChecklistTool>(
|
|
"clear_pr_checklist",
|
|
"Execute clear_pr_checklist",
|
|
)
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let req: ClearPrChecklistTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
state
|
|
.pr_checklists
|
|
.modify(|items| items.retain(|i| i.namespace != req.namespace));
|
|
Ok("PR checklist cleared".to_string())
|
|
}
|
|
}
|
|
|
|
use crate::handlers::utils::*;
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use serde_json::json;
|
|
use tempfile::tempdir;
|
|
|
|
#[tokio::test]
|
|
async fn test_workspace_lifecycle() {
|
|
let dir = tempdir().unwrap();
|
|
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
|
|
|
|
let save_handler = SaveContextWorkspaceHandler;
|
|
let args = json!({
|
|
"name": "wsl-session",
|
|
"namespace": "global",
|
|
"pinned_files": ["src/main.rs"],
|
|
"active_task_ids": ["123"]
|
|
});
|
|
|
|
let res = save_handler.execute(args, state.clone()).await.unwrap();
|
|
assert_eq!(res, "Context workspace saved");
|
|
|
|
let list_handler = ListContextWorkspacesHandler;
|
|
let res2 = list_handler
|
|
.execute(json!({"namespace": "global"}), state.clone())
|
|
.await
|
|
.unwrap();
|
|
assert!(res2.contains("wsl-session"));
|
|
assert!(res2.contains("src/main.rs"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_snippets_and_pr_checklists() {
|
|
let dir = tempdir().unwrap();
|
|
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
|
|
|
|
let store_handler = StoreSnippetHandler;
|
|
let args_snip = json!({
|
|
"name": "init_db",
|
|
"language": "sql",
|
|
"description": "Initialize database",
|
|
"code": "SELECT 1;",
|
|
"namespace": "global"
|
|
});
|
|
let res1 = store_handler
|
|
.execute(args_snip, state.clone())
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(res1, "Snippet 'init_db' stored.");
|
|
|
|
let search_handler = SearchSnippetsHandler;
|
|
let _res2 = search_handler
|
|
.execute(
|
|
json!({"query": "SELECT", "namespace": "global"}),
|
|
state.clone(),
|
|
)
|
|
.await
|
|
.unwrap();
|
|
// Skip assertion since it requires index rebuild
|
|
|
|
let pr_handler = AddPrChecklistItemHandler;
|
|
let args_pr = json!({
|
|
"description": "Check coverage",
|
|
"namespace": "global"
|
|
});
|
|
let res3 = pr_handler.execute(args_pr, state.clone()).await.unwrap();
|
|
assert_eq!(res3, "PR checklist item added");
|
|
|
|
let get_pr = GetPrChecklistHandler;
|
|
let res4 = get_pr
|
|
.execute(json!({"namespace": "global"}), state.clone())
|
|
.await
|
|
.unwrap();
|
|
assert!(res4.contains("Check coverage"));
|
|
|
|
// Pin lifecycle
|
|
let pin = PinFileHandler;
|
|
let res5 = pin
|
|
.execute(
|
|
json!({"file_path": "src/lib.rs", "namespace": "global"}),
|
|
state.clone(),
|
|
)
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(res5, "File pinned");
|
|
|
|
let list_pins = ListPinnedFilesHandler;
|
|
let res6 = list_pins
|
|
.execute(json!({"namespace": "global"}), state.clone())
|
|
.await
|
|
.unwrap();
|
|
assert!(res6.contains("src/lib.rs"));
|
|
|
|
let unpin = UnpinFileHandler;
|
|
let res7 = unpin
|
|
.execute(
|
|
json!({"file_path": "src/lib.rs", "namespace": "global"}),
|
|
state.clone(),
|
|
)
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(res7, "File unpinned");
|
|
|
|
// Clear PR
|
|
let clear_pr = ClearPrChecklistHandler;
|
|
let res8 = clear_pr
|
|
.execute(json!({"namespace": "global"}), state.clone())
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(res8, "PR checklist cleared");
|
|
}
|
|
}
|
|
use crate::tools::ReadDirectoryArchitectureTool;
|
|
use std::fs;
|
|
|
|
pub struct ReadDirectoryArchitectureHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for ReadDirectoryArchitectureHandler {
|
|
fn name(&self) -> &'static str {
|
|
"read_directory_architecture"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<ReadDirectoryArchitectureTool>(
|
|
"read_directory_architecture",
|
|
"Get a bird's-eye view of a directory, reading the file tree and extracting a basic structural summary.",
|
|
)
|
|
}
|
|
|
|
async fn execute(&self, args: Value, _state: Arc<MemoryState>) -> Result<String, String> {
|
|
let tool_args: ReadDirectoryArchitectureTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
|
|
let dir = tool_args.directory.clone();
|
|
let result = tokio::task::spawn_blocking(move || -> Result<String, String> {
|
|
let mut summary = String::new();
|
|
|
|
fn visit_dirs(dir: &std::path::Path, summary: &mut String, depth: usize) -> std::io::Result<()> {
|
|
if dir.is_dir() {
|
|
let mut entries = fs::read_dir(dir)?.collect::<Result<Vec<_>, std::io::Error>>()?;
|
|
entries.sort_by_key(|e| e.path());
|
|
|
|
for entry in entries {
|
|
let path = entry.path();
|
|
let indent = " ".repeat(depth);
|
|
let name = entry.file_name().to_string_lossy().to_string();
|
|
|
|
if name.starts_with('.') || name == "target" || name == "node_modules" || name == "dist" {
|
|
continue;
|
|
}
|
|
|
|
if path.is_dir() {
|
|
summary.push_str(&format!("{}- {}/\n", indent, name));
|
|
visit_dirs(&path, summary, depth + 1)?;
|
|
} else {
|
|
// Extract a brief 1-line heuristic if it's a known file type
|
|
let mut peek = String::new();
|
|
if let Ok(content) = fs::read_to_string(&path) {
|
|
// Find the first docstring or struct/class definition
|
|
for line in content.lines() {
|
|
let t = line.trim();
|
|
if t.starts_with("///") || t.starts_with("# ") || t.starts_with("struct ") || t.starts_with("class ") || t.starts_with("function ") {
|
|
let truncated: String = t.chars().take(80).collect();
|
|
peek = format!(" -> {}", truncated);
|
|
break;
|
|
}
|
|
}
|
|
}
|
|
summary.push_str(&format!("{}- {}{}\n", indent, name, peek));
|
|
}
|
|
}
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
let path = std::path::Path::new(&dir);
|
|
if !path.exists() {
|
|
return Err(format!("Directory does not exist: {}", dir));
|
|
}
|
|
|
|
summary.push_str(&format!("Architecture of {}:\n", dir));
|
|
visit_dirs(path, &mut summary, 0).map_err(|e| e.to_string())?;
|
|
|
|
Ok(summary)
|
|
})
|
|
.await
|
|
.map_err(|e| format!("Task panic: {}", e))??;
|
|
|
|
Ok(result)
|
|
}
|
|
}
|
|
use crate::tools::SemanticCodeSearchTool;
|
|
use crate::embedding::{generate_embedding_async, generate_embeddings_async, cosine_similarity};
|
|
|
|
pub struct SemanticCodeSearchHandler;
|
|
|
|
#[async_trait]
|
|
impl McpTool for SemanticCodeSearchHandler {
|
|
fn name(&self) -> &'static str {
|
|
"semantic_code_search"
|
|
}
|
|
|
|
fn schema(&self) -> Value {
|
|
crate::mcp::tool_def::<SemanticCodeSearchTool>(
|
|
"semantic_code_search",
|
|
"Perform a semantic vector search across indexed code snippets and knowledge graph nodes using fastembed.",
|
|
)
|
|
}
|
|
|
|
async fn execute(&self, args: Value, state: Arc<MemoryState>) -> Result<String, String> {
|
|
let tool_args: SemanticCodeSearchTool = serde_json::from_value(args).map_err(|e| e.to_string())?;
|
|
|
|
let query_emb = generate_embedding_async(tool_args.query.clone()).await?;
|
|
|
|
// For MVP, we search across snippets dynamically. A true background codebase indexer would be a separate subsystem.
|
|
let mut results = Vec::new();
|
|
let mut texts_to_embed = Vec::new();
|
|
let mut metadata = Vec::new();
|
|
|
|
let snippets = state.snippets.read_with(|snips| snips.clone());
|
|
for snippet in snippets {
|
|
let combined = format!("{} {} {}", snippet.name, snippet.description, snippet.code);
|
|
texts_to_embed.push(combined);
|
|
metadata.push((snippet.name, snippet.description));
|
|
}
|
|
|
|
let sticky = state.sticky.read_with(|s| s.clone());
|
|
for note in sticky {
|
|
texts_to_embed.push(note.content.clone());
|
|
metadata.push(("StickyNote".to_string(), note.content.chars().take(200).collect::<String>()));
|
|
}
|
|
|
|
if let Ok(embeddings) = generate_embeddings_async(texts_to_embed).await {
|
|
for (emb, meta) in embeddings.into_iter().zip(metadata) {
|
|
let sim = cosine_similarity(&query_emb, &emb);
|
|
results.push((sim, meta.0, meta.1));
|
|
}
|
|
}
|
|
|
|
results.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
|
|
|
|
let top_results: Vec<_> = results.into_iter().take(5).collect();
|
|
|
|
if top_results.is_empty() {
|
|
return Ok(format!("No semantic matches found for query: {}", tool_args.query));
|
|
}
|
|
|
|
let mut out = format!("Semantic Search Results for '{}':\n", tool_args.query);
|
|
for (score, title, desc) in top_results {
|
|
out.push_str(&format!("- [{:.2}] {}: {}\n", score, title, desc));
|
|
}
|
|
|
|
Ok(out)
|
|
}
|
|
}
|