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