Skip to main content

apollo/tools/
embeddings.rs

1use 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}