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(
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}