Skip to main content

locus_sdk/application/
memory_transform.rs

1use std::sync::Arc;
2
3use anyhow::Result;
4use chrono::Utc;
5use locus_core_rs::domain::contracts::{NodeStore, SemanticIndexStore, TagEmbedding};
6use locus_core_rs::domain::models::{NodeQuery, NodeUpsertStatus, SemanticTagNodeRef, SemanticTagQueryFilter};
7use locus_core_rs::storage::derive_tenant_id_from_session;
8
9use crate::application::ai_router::route_embedding;
10use crate::application::memory_filters::{build_session_filter, node_matches_common_filters};
11use crate::domain::ai::{AiProviderRegistry, AiTask, EmbedRequest, ProviderPolicy};
12use crate::domain::memory::{
13    MemoryTransformOperation, MemoryTransformRequest, MemoryTransformResult, clamp_batch_size,
14    clamp_nodes,
15};
16
17pub struct MemoryTransformService {
18    store: Arc<dyn NodeStore>,
19    providers: Arc<dyn AiProviderRegistry>,
20    semantic_index: Option<Arc<dyn SemanticIndexStore>>,
21}
22
23impl MemoryTransformService {
24    /// Create a transform service with storage and provider registry dependencies.
25    pub fn new(store: Arc<dyn NodeStore>, providers: Arc<dyn AiProviderRegistry>) -> Self {
26        Self {
27            store,
28            providers,
29            semantic_index: None,
30        }
31    }
32
33    pub fn with_semantic_index(
34        mut self,
35        semantic_index: Arc<dyn SemanticIndexStore>,
36    ) -> Self {
37        self.semantic_index = Some(semantic_index);
38        self
39    }
40
41    /// Execute a bulk memory transform operation.
42    ///
43    /// The current implementation supports embedding backfill with optional
44    /// dry-run behavior, batch control, and bounded failure reporting.
45    pub async fn execute(&self, request: &MemoryTransformRequest) -> Result<MemoryTransformResult> {
46        if matches!(
47            request.operation,
48            MemoryTransformOperation::EmbedTagBackfill | MemoryTransformOperation::ReindexTagEmbeddings
49        ) {
50            return self.execute_tag_transform(request).await;
51        }
52
53        let started_at = Utc::now();
54        let max_nodes = clamp_nodes(if request.max_nodes == 0 {
55            5000
56        } else {
57            request.max_nodes
58        });
59        let batch_size = clamp_batch_size(if request.batch_size == 0 {
60            100
61        } else {
62            request.batch_size
63        });
64
65        let single_session = request
66            .scope
67            .session_ids
68            .as_deref()
69            .filter(|sessions| sessions.len() == 1)
70            .and_then(|sessions| sessions.first().cloned());
71
72        let nodes = self
73            .store
74            .query_nodes_async(NodeQuery {
75                limit: max_nodes,
76                session_id: single_session,
77                from_utc: request.scope.from_utc,
78                to_utc: request.scope.to_utc,
79                tiers: request.scope.tiers.clone(),
80            })
81            .await?;
82
83        let session_filter = build_session_filter(&request.scope);
84
85        let mut selected = nodes
86            .into_iter()
87            .filter(|node| {
88                node_matches_common_filters(node, &request.scope, &request.filter, session_filter.as_ref())
89            })
90            .collect::<Vec<_>>();
91
92        if request.operation == MemoryTransformOperation::EmbedBackfill {
93            selected.retain(|node| node.embedding.as_ref().is_none_or(|values| values.is_empty()));
94        }
95
96        if request.operation == MemoryTransformOperation::ReindexEmbeddings {
97            // Reindex all selected nodes.
98        }
99
100        let mut result = MemoryTransformResult {
101            scanned: selected.len(),
102            selected: selected.len(),
103            started_at,
104            completed_at: started_at,
105            ..Default::default()
106        };
107
108        if request.dry_run {
109            result.updated = result.selected;
110            result.completed_at = Utc::now();
111            return Ok(result);
112        }
113
114        for chunk in selected.chunks(batch_size) {
115            for mut node in chunk.iter().cloned() {
116                let Some(embedding_input) = build_embedding_input(node.context_summary.as_deref(), &node.session_id)
117                else {
118                    result.skipped += 1;
119                    continue;
120                };
121
122                let embed_request = EmbedRequest {
123                    text: embedding_input,
124                    task: AiTask::SemanticEmbedding,
125                    provider_id: request.provider_id.clone(),
126                    model: request.model.clone(),
127                    policy: if request.provider_id.is_some() {
128                        ProviderPolicy::Required
129                    } else {
130                        ProviderPolicy::Auto
131                    },
132                };
133
134                let vector = match route_embedding(self.providers.as_ref(), &embed_request).await {
135                    Ok(values) if !values.is_empty() => values,
136                    Ok(_) => {
137                        result.failed += 1;
138                        push_failure(
139                            &mut result.failures,
140                            format!("{}: embedding provider returned empty vector", node.sync_key),
141                        );
142                        continue;
143                    }
144                    Err(err) => {
145                        result.failed += 1;
146                        push_failure(
147                            &mut result.failures,
148                            format!("{}: embedding failed: {err}", node.sync_key),
149                        );
150                        continue;
151                    }
152                };
153
154                node.embedding_dimensions = Some(vector.len());
155                node.embedding_model = request
156                    .model
157                    .clone()
158                    .or_else(|| request.provider_id.clone())
159                    .or_else(|| Some("sdk-memory-transform".to_string()));
160                node.embedding = Some(vector);
161                node.embedded_at = Some(Utc::now());
162                node.updated_at = Utc::now();
163
164                match self.store.upsert_node_async(node).await {
165                    Ok(status) => match status.status {
166                        NodeUpsertStatus::Created | NodeUpsertStatus::Updated => result.updated += 1,
167                        NodeUpsertStatus::Duplicate => result.duplicate += 1,
168                        NodeUpsertStatus::Skipped => result.skipped += 1,
169                    },
170                    Err(err) => {
171                        result.failed += 1;
172                        push_failure(&mut result.failures, format!("store upsert failed: {err}"));
173                    }
174                }
175            }
176        }
177
178        result.completed_at = Utc::now();
179        Ok(result)
180    }
181
182    async fn execute_tag_transform(
183        &self,
184        request: &MemoryTransformRequest,
185    ) -> Result<MemoryTransformResult> {
186        let started_at = Utc::now();
187        let index = self
188            .semantic_index
189            .as_ref()
190            .ok_or_else(|| anyhow::anyhow!("semantic index store is not configured"))?;
191
192        let tenant_id = request
193            .scope
194            .tenant_id
195            .clone()
196            .or_else(|| {
197                request
198                    .scope
199                    .session_ids
200                    .as_ref()
201                    .and_then(|sessions| sessions.first().cloned())
202                    .map(|session| derive_tenant_id_from_session(&session))
203            })
204            .unwrap_or_else(|| "default".to_string());
205
206        let missing_only = request.operation == MemoryTransformOperation::EmbedTagBackfill;
207        let records = index
208            .query_tag_records_async(SemanticTagQueryFilter {
209                tenant_id: Some(tenant_id),
210                session_id: request
211                    .scope
212                    .session_ids
213                    .as_ref()
214                    .and_then(|sessions| sessions.first().cloned()),
215                tags: request.filter.indexed_tags.clone(),
216                tag_prefix: request.filter.tag_prefix.clone(),
217                has_embedding: if missing_only {
218                    None
219                } else {
220                    Some(true)
221                },
222                missing_embedding_only: missing_only,
223                limit: clamp_nodes(if request.max_nodes == 0 {
224                    5000
225                } else {
226                    request.max_nodes
227                }),
228            })
229            .await?;
230
231        let mut result = MemoryTransformResult {
232            scanned: records.len(),
233            selected: records.len(),
234            started_at,
235            completed_at: started_at,
236            ..Default::default()
237        };
238
239        if request.dry_run {
240            result.updated = result.selected;
241            result.completed_at = Utc::now();
242            return Ok(result);
243        }
244
245        let batch_size = clamp_batch_size(if request.batch_size == 0 {
246            100
247        } else {
248            request.batch_size
249        });
250
251        for chunk in records.chunks(batch_size) {
252            for record in chunk {
253                let embed_request = EmbedRequest {
254                    text: record.tag.clone(),
255                    task: AiTask::SemanticEmbedding,
256                    provider_id: request.provider_id.clone(),
257                    model: request.model.clone(),
258                    policy: if request.provider_id.is_some() {
259                        ProviderPolicy::Required
260                    } else {
261                        ProviderPolicy::Auto
262                    },
263                };
264
265                let vector = match route_embedding(self.providers.as_ref(), &embed_request).await {
266                    Ok(values) if !values.is_empty() => values,
267                    Ok(_) => {
268                        result.failed += 1;
269                        push_failure(
270                            &mut result.failures,
271                            format!(
272                                "{}:{}: embedding provider returned empty vector",
273                                record.sync_key, record.tag
274                            ),
275                        );
276                        continue;
277                    }
278                    Err(err) => {
279                        result.failed += 1;
280                        push_failure(
281                            &mut result.failures,
282                            format!(
283                                "{}:{}: embedding failed: {err}",
284                                record.sync_key, record.tag
285                            ),
286                        );
287                        continue;
288                    }
289                };
290
291                let model = request
292                    .model
293                    .clone()
294                    .or_else(|| request.provider_id.clone())
295                    .unwrap_or_else(|| "sdk-memory-transform".to_string());
296
297                let mut embeddings = std::collections::HashMap::new();
298                embeddings.insert(
299                    record.tag.clone(),
300                    TagEmbedding {
301                        vector,
302                        model,
303                    },
304                );
305
306                let node_ref = SemanticTagNodeRef {
307                    tenant_id: record.tenant_id.clone(),
308                    session_id: record.session_id.clone(),
309                    node_id: record.node_id.clone(),
310                    sync_key: record.sync_key.clone(),
311                };
312
313                match index
314                    .sync_node_tags_async(node_ref, &[record.tag.clone()], Some(&embeddings))
315                    .await
316                {
317                    Ok(()) => result.updated += 1,
318                    Err(err) => {
319                        result.failed += 1;
320                        push_failure(
321                            &mut result.failures,
322                            format!(
323                                "{}:{}: semantic index sync failed: {err}",
324                                record.sync_key, record.tag
325                            ),
326                        );
327                    }
328                }
329            }
330        }
331
332        result.completed_at = Utc::now();
333        Ok(result)
334    }
335}
336
337fn build_embedding_input(context_summary: Option<&str>, session_id: &str) -> Option<String> {
338    let summary = context_summary.and_then(|value| {
339        let trimmed = value.trim();
340        if trimmed.is_empty() {
341            None
342        } else {
343            Some(trimmed)
344        }
345    });
346
347    let session = {
348        let trimmed = session_id.trim();
349        if trimmed.is_empty() {
350            None
351        } else {
352            Some(trimmed)
353        }
354    };
355
356    match (summary, session) {
357        (Some(summary), Some(session)) => Some(format!("{summary}\nsession_id:{session}")),
358        (Some(summary), None) => Some(summary.to_string()),
359        (None, Some(session)) => Some(format!("session_id:{session}")),
360        (None, None) => None,
361    }
362}
363
364fn push_failure(failures: &mut Vec<String>, reason: String) {
365    if failures.len() < 100 {
366        failures.push(reason);
367    }
368}
369
370#[cfg(test)]
371mod tests {
372    use std::sync::Arc;
373
374    use anyhow::Result;
375    use async_trait::async_trait;
376    use chrono::Utc;
377    use locus_core_rs::{InMemoryNodeStore, NodeStore};
378    use locus_core_rs::domain::models::{AvecState, SttpNode};
379
380    use super::MemoryTransformService;
381    use crate::domain::ai::{
382        AiCapability, AiProvider, EmbedRequest, ScoreAvecRequest,
383    };
384    use crate::domain::memory::{MemoryTransformOperation, MemoryTransformRequest};
385    use crate::infrastructure::registry::InMemoryAiProviderRegistry;
386
387    struct MockEmbeddingProvider;
388
389    #[async_trait]
390    impl AiProvider for MockEmbeddingProvider {
391        fn provider_id(&self) -> &str {
392            "mock"
393        }
394
395        fn capabilities(&self) -> &'static [AiCapability] {
396            &[AiCapability::SemanticEmbedding]
397        }
398
399        async fn embed_semantic(&self, _request: &EmbedRequest) -> Result<Vec<f32>> {
400            Ok(vec![0.2, 0.3, 0.4])
401        }
402
403        async fn embed_avec(&self, _request: &EmbedRequest) -> Result<Vec<f32>> {
404            Ok(vec![0.2, 0.3, 0.4])
405        }
406
407        async fn score_avec(&self, _request: &ScoreAvecRequest) -> Result<AvecState> {
408            Ok(AvecState::zero())
409        }
410    }
411
412    #[tokio::test]
413    async fn dry_run_reports_selected_without_writes() {
414        let store: Arc<dyn NodeStore> = Arc::new(InMemoryNodeStore::new());
415        let node = test_node("dry-run", None);
416        store
417            .upsert_node_async(node)
418            .await
419            .expect("upsert should succeed");
420
421        let mut providers = InMemoryAiProviderRegistry::new();
422        providers.register(MockEmbeddingProvider);
423
424        let service = MemoryTransformService::new(store, Arc::new(providers));
425
426        let request = MemoryTransformRequest {
427            operation: MemoryTransformOperation::EmbedBackfill,
428            dry_run: true,
429            max_nodes: 100,
430            batch_size: 10,
431            ..Default::default()
432        };
433
434        let result = service.execute(&request).await.expect("transform should succeed");
435
436        assert_eq!(result.selected, 1);
437        assert_eq!(result.updated, 1);
438        assert_eq!(result.failed, 0);
439    }
440
441    #[tokio::test]
442    async fn embed_backfill_updates_missing_embedding_nodes() {
443        let store: Arc<dyn NodeStore> = Arc::new(InMemoryNodeStore::new());
444        let node = test_node("backfill", None);
445        store
446            .upsert_node_async(node)
447            .await
448            .expect("upsert should succeed");
449
450        let mut providers = InMemoryAiProviderRegistry::new();
451        providers.register(MockEmbeddingProvider);
452
453        let service = MemoryTransformService::new(store.clone(), Arc::new(providers));
454
455        let request = MemoryTransformRequest {
456            operation: MemoryTransformOperation::EmbedBackfill,
457            dry_run: false,
458            max_nodes: 100,
459            batch_size: 10,
460            ..Default::default()
461        };
462
463        let result = service.execute(&request).await.expect("transform should succeed");
464
465        assert_eq!(result.updated, 1);
466        assert_eq!(result.failed, 0);
467
468        let nodes = store
469            .query_nodes_async(locus_core_rs::domain::models::NodeQuery {
470                limit: 10,
471                session_id: Some("backfill".to_string()),
472                ..Default::default()
473            })
474            .await
475            .expect("query should succeed");
476
477        assert_eq!(nodes.len(), 1);
478        assert!(nodes[0].embedding.as_ref().is_some_and(|v| !v.is_empty()));
479    }
480
481    fn test_node(session_id: &str, embedding: Option<Vec<f32>>) -> SttpNode {
482        let now = Utc::now();
483        let user = AvecState {
484            stability: 0.6,
485            friction: 0.4,
486            logic: 0.8,
487            autonomy: 0.7,
488        };
489        let model = AvecState {
490            stability: 0.5,
491            friction: 0.3,
492            logic: 0.9,
493            autonomy: 0.6,
494        };
495
496        SttpNode {
497            raw: format!("raw:{session_id}"),
498            session_id: session_id.to_string(),
499            tier: "raw".to_string(),
500            timestamp: now,
501            compression_depth: 1,
502            parent_node_id: None,
503            sync_key: format!("{}:{}", session_id, now.timestamp_nanos_opt().unwrap_or_default()),
504            updated_at: now,
505            source_metadata: None,
506            context_summary: Some("summary".to_string()),
507            semantic_tags: None,
508            semantic_links: None,
509            embedding_dimensions: embedding.as_ref().map(|v| v.len()),
510            embedding_model: embedding.as_ref().map(|_| "existing".to_string()),
511            embedding,
512            embedded_at: None,
513            user_avec: user,
514            model_avec: model,
515            compression_avec: Some(model),
516            rho: 0.9,
517            kappa: 0.8,
518            psi: 2.5,
519        }
520    }
521}