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(
96 &args.namespace,
97 &args.key,
98 &vector,
99 &args.text,
100 self.provider.name(),
101 )
102 .await?;
103 Ok(ToolResult::success(format!(
104 "Stored embedding for {}/{} with {} dimensions",
105 args.namespace,
106 args.key,
107 vector.len()
108 )))
109 }
110}
111
112pub struct EmbeddingSearchTool {
113 provider: Arc<dyn EmbeddingProvider>,
114 memory: Arc<dyn MemoryBackend>,
115}
116
117impl EmbeddingSearchTool {
118 pub fn new(provider: Arc<dyn EmbeddingProvider>, memory: Arc<dyn MemoryBackend>) -> Self {
119 Self { provider, memory }
120 }
121}
122
123#[derive(Deserialize)]
124struct EmbeddingSearchArgs {
125 namespace: String,
126 query: String,
127 limit: Option<usize>,
128}
129
130#[async_trait]
131impl Tool for EmbeddingSearchTool {
132 fn name(&self) -> &str {
133 "embedding_search"
134 }
135
136 fn spec(&self) -> ToolSpec {
137 ToolSpec {
138 name: "embedding_search".to_string(),
139 description: "Embed a query and run semantic search over stored vectors.".to_string(),
140 parameters: json!({
141 "type": "object",
142 "properties": {
143 "namespace": { "type": "string" },
144 "query": { "type": "string" },
145 "limit": { "type": "integer" }
146 },
147 "required": ["namespace", "query"]
148 }),
149 }
150 }
151
152 async fn execute(&self, arguments: &str) -> anyhow::Result<ToolResult> {
153 let args: EmbeddingSearchArgs = serde_json::from_str(arguments)?;
154 let limit = args.limit.unwrap_or(5);
155 let query_vector = self.provider.embed_one(&args.query).await?;
156 let results = self
157 .memory
158 .search_embeddings(&args.namespace, &query_vector, limit)
159 .await?;
160
161 if results.is_empty() {
162 return Ok(ToolResult::success("No embedding matches found."));
163 }
164
165 let output = results
166 .iter()
167 .map(|entry| format!("{}:{} — {}", entry.namespace, entry.key, entry.text))
168 .collect::<Vec<_>>()
169 .join("\n");
170 Ok(ToolResult::success(output))
171 }
172}