apollo/tools/
embeddings.rs1use std::sync::Arc;
2
3use async_trait::async_trait;
4use serde::Deserialize;
5use serde_json::json;
6
7use crate::memory::embeddings::EmbeddingProvider;
8use crate::memory::MemoryBackend;
9
10use super::{Tool, ToolResult, ToolSpec};
11
12pub struct EmbeddingStatusTool {
13 provider: Arc<dyn EmbeddingProvider>,
14}
15
16impl EmbeddingStatusTool {
17 pub fn new(provider: Arc<dyn EmbeddingProvider>) -> Self {
18 Self { provider }
19 }
20}
21
22#[async_trait]
23impl Tool for EmbeddingStatusTool {
24 fn name(&self) -> &str {
25 "embedding_status"
26 }
27
28 fn spec(&self) -> ToolSpec {
29 ToolSpec {
30 name: "embedding_status".to_string(),
31 description: "Show the configured embedding provider and vector dimensions."
32 .to_string(),
33 parameters: json!({
34 "type": "object",
35 "properties": {}
36 }),
37 }
38 }
39
40 async fn execute(&self, _arguments: &str) -> anyhow::Result<ToolResult> {
41 Ok(ToolResult::success(format!(
42 "Embedding provider: {}\nDimensions: {}",
43 self.provider.name(),
44 self.provider.dimensions()
45 )))
46 }
47}
48
49pub struct EmbeddingStoreTool {
50 provider: Arc<dyn EmbeddingProvider>,
51 memory: Arc<dyn MemoryBackend>,
52}
53
54impl EmbeddingStoreTool {
55 pub fn new(provider: Arc<dyn EmbeddingProvider>, memory: Arc<dyn MemoryBackend>) -> Self {
56 Self { provider, memory }
57 }
58}
59
60#[derive(Deserialize)]
61struct EmbeddingStoreArgs {
62 namespace: String,
63 key: String,
64 text: String,
65}
66
67#[async_trait]
68impl Tool for EmbeddingStoreTool {
69 fn name(&self) -> &str {
70 "embedding_store"
71 }
72
73 fn spec(&self) -> ToolSpec {
74 ToolSpec {
75 name: "embedding_store".to_string(),
76 description:
77 "Generate an embedding for text and store it in the active memory backend."
78 .to_string(),
79 parameters: json!({
80 "type": "object",
81 "properties": {
82 "namespace": { "type": "string" },
83 "key": { "type": "string" },
84 "text": { "type": "string" }
85 },
86 "required": ["namespace", "key", "text"]
87 }),
88 }
89 }
90
91 async fn execute(&self, arguments: &str) -> anyhow::Result<ToolResult> {
92 let args: EmbeddingStoreArgs = serde_json::from_str(arguments)?;
93 let vector = self.provider.embed_one(&args.text).await?;
94 self.memory
95 .store_embedding(&args.namespace, &args.key, &vector, &args.text)
96 .await?;
97 Ok(ToolResult::success(format!(
98 "Stored embedding for {}/{} with {} dimensions",
99 args.namespace,
100 args.key,
101 vector.len()
102 )))
103 }
104}
105
106pub struct EmbeddingSearchTool {
107 provider: Arc<dyn EmbeddingProvider>,
108 memory: Arc<dyn MemoryBackend>,
109}
110
111impl EmbeddingSearchTool {
112 pub fn new(provider: Arc<dyn EmbeddingProvider>, memory: Arc<dyn MemoryBackend>) -> Self {
113 Self { provider, memory }
114 }
115}
116
117#[derive(Deserialize)]
118struct EmbeddingSearchArgs {
119 namespace: String,
120 query: String,
121 limit: Option<usize>,
122}
123
124#[async_trait]
125impl Tool for EmbeddingSearchTool {
126 fn name(&self) -> &str {
127 "embedding_search"
128 }
129
130 fn spec(&self) -> ToolSpec {
131 ToolSpec {
132 name: "embedding_search".to_string(),
133 description: "Embed a query and run semantic search over stored vectors.".to_string(),
134 parameters: json!({
135 "type": "object",
136 "properties": {
137 "namespace": { "type": "string" },
138 "query": { "type": "string" },
139 "limit": { "type": "integer" }
140 },
141 "required": ["namespace", "query"]
142 }),
143 }
144 }
145
146 async fn execute(&self, arguments: &str) -> anyhow::Result<ToolResult> {
147 let args: EmbeddingSearchArgs = serde_json::from_str(arguments)?;
148 let limit = args.limit.unwrap_or(5);
149 let query_vector = self.provider.embed_one(&args.query).await?;
150 let results = self
151 .memory
152 .search_embeddings(&args.namespace, &query_vector, limit)
153 .await?;
154
155 if results.is_empty() {
156 return Ok(ToolResult::success("No embedding matches found."));
157 }
158
159 let output = results
160 .iter()
161 .map(|entry| format!("{}:{} — {}", entry.namespace, entry.key, entry.text))
162 .collect::<Vec<_>>()
163 .join("\n");
164 Ok(ToolResult::success(output))
165 }
166}