Skip to main content

locus_mcp/tools/
run_embedding_migration.rs

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}