1use locus_core_rs::{EmbeddingMigrationFilter, EmbeddingMigrationRunRequest};
2use locus_sdk::application::memory_transform::MemoryTransformService;
3use locus_sdk::domain::memory::{
4 MemoryFilter, MemoryScope, MemoryTransformOperation, MemoryTransformRequest,
5};
6use locus_sdk::infrastructure::registry::InMemoryAiProviderRegistry;
7use locus_sdk::infrastructure::sttp_native::embedding_provider_adapter::SttpEmbeddingProviderAdapter;
8use serde_json::json;
9use tracing::error;
10
11use crate::{
12 RunEmbeddingMigrationRequest, SttpMcpServer, mode_to_string, normalize_tiers,
13 parse_migration_mode, parse_utc_optional, to_json_string, tool_error, validate_batch_size,
14 validate_max_nodes,
15};
16
17pub(crate) async fn execute(
18 server: &SttpMcpServer,
19 request: RunEmbeddingMigrationRequest,
20) -> String {
21 let from_utc = match parse_utc_optional(request.from_utc.as_deref(), "from_utc") {
22 Ok(value) => value,
23 Err(message) => return tool_error("InvalidDate", &message),
24 };
25 let to_utc = match parse_utc_optional(request.to_utc.as_deref(), "to_utc") {
26 Ok(value) => value,
27 Err(message) => return tool_error("InvalidDate", &message),
28 };
29 let tiers = request
30 .tiers
31 .as_ref()
32 .map(|values| normalize_tiers(values.as_slice()));
33 let batch_size = match validate_batch_size(request.batch_size) {
34 Ok(value) => value,
35 Err(message) => return tool_error("InvalidArgument", &message),
36 };
37 let max_nodes = match validate_max_nodes(request.max_nodes) {
38 Ok(value) => value,
39 Err(message) => return tool_error("InvalidArgument", &message),
40 };
41
42 let mode_raw = request
43 .mode
44 .as_deref()
45 .unwrap_or("missing_only")
46 .trim()
47 .to_ascii_lowercase();
48
49 if matches!(mode_raw.as_str(), "tags" | "tag" | "embed_tag_backfill" | "both") {
50 return run_tag_transform(server, request, from_utc, to_utc, tiers, batch_size, max_nodes)
51 .await;
52 }
53
54 let mode = match parse_migration_mode(request.mode.as_deref()) {
55 Ok(value) => value,
56 Err(message) => return tool_error("InvalidArgument", &message),
57 };
58
59 let filter = EmbeddingMigrationFilter {
60 session_id: request.session_id,
61 from_utc,
62 to_utc,
63 tiers,
64 has_embedding: request.has_embedding,
65 embedding_model: request.embedding_model,
66 sync_keys: request.sync_keys,
67 };
68
69 match server
70 .embedding_migration
71 .run_async(EmbeddingMigrationRunRequest {
72 filter,
73 mode,
74 dry_run: request.dry_run,
75 batch_size,
76 max_nodes,
77 })
78 .await
79 {
80 Ok(result) => to_json_string(json!({
81 "scanned": result.scanned,
82 "selected": result.selected,
83 "updated": result.updated,
84 "skipped": result.skipped,
85 "failed": result.failed,
86 "duplicate": result.duplicate,
87 "started_at": result.started_at.to_rfc3339(),
88 "completed_at": result.completed_at.to_rfc3339(),
89 "provider_model": result.provider_model,
90 "dry_run": request.dry_run,
91 "mode": mode_to_string(mode),
92 "failure_reasons": result.failure_reasons,
93 })),
94 Err(err) => {
95 error!(error = %err, "run_embedding_migration failed");
96 tool_error("MigrationRunFailure", &err.to_string())
97 }
98 }
99}
100
101async fn run_tag_transform(
102 server: &SttpMcpServer,
103 request: RunEmbeddingMigrationRequest,
104 from_utc: Option<chrono::DateTime<chrono::Utc>>,
105 to_utc: Option<chrono::DateTime<chrono::Utc>>,
106 tiers: Option<Vec<String>>,
107 batch_size: usize,
108 max_nodes: usize,
109) -> String {
110 let mode_raw = request
111 .mode
112 .as_deref()
113 .unwrap_or("tags")
114 .trim()
115 .to_ascii_lowercase();
116 let operation = if mode_raw == "both" {
117 MemoryTransformOperation::ReindexTagEmbeddings
118 } else {
119 MemoryTransformOperation::EmbedTagBackfill
120 };
121
122 let mut registry = InMemoryAiProviderRegistry::new();
123 if let Some(provider) = server.embedding_provider.as_ref() {
124 registry.register(SttpEmbeddingProviderAdapter::new(
125 "mcp-embedding",
126 provider.clone(),
127 ));
128 }
129
130 let transform_service = MemoryTransformService::new(
131 server.node_store.clone(),
132 std::sync::Arc::new(registry),
133 )
134 .with_semantic_index(server.semantic_index.clone());
135
136 match transform_service
137 .execute(&MemoryTransformRequest {
138 scope: MemoryScope {
139 tenant_id: None,
140 session_ids: request.session_id.map(|session| vec![session]),
141 tiers,
142 from_utc,
143 to_utc,
144 },
145 filter: MemoryFilter::default(),
146 operation,
147 dry_run: request.dry_run,
148 batch_size,
149 max_nodes,
150 provider_id: server
151 .embedding_provider
152 .as_ref()
153 .map(|_| "mcp-embedding".to_string()),
154 model: server
155 .embedding_provider
156 .as_ref()
157 .map(|provider| provider.model_name().to_string()),
158 })
159 .await
160 {
161 Ok(result) => to_json_string(json!({
162 "scanned": result.scanned,
163 "selected": result.selected,
164 "updated": result.updated,
165 "skipped": result.skipped,
166 "failed": result.failed,
167 "duplicate": result.duplicate,
168 "started_at": result.started_at.to_rfc3339(),
169 "completed_at": result.completed_at.to_rfc3339(),
170 "provider_model": server.embedding_provider.as_ref().map(|provider| provider.model_name().to_string()),
171 "dry_run": request.dry_run,
172 "mode": mode_raw,
173 "failure_reasons": result.failures,
174 })),
175 Err(err) => {
176 error!(error = %err, "run_embedding_migration tag transform failed");
177 tool_error("MigrationRunFailure", &err.to_string())
178 }
179 }
180}