Skip to main content

locus_mcp/tools/
get_context.rs

1use locus_sdk::application::memory_recall::MemoryRecallService;
2use locus_sdk::domain::memory::{MemoryFilter, MemoryPage, MemoryRecallRequest, MemoryScope, MemoryScoring};
3use serde_json::json;
4use tracing::error;
5
6use crate::{
7    GetContextRequest, SttpMcpServer, normalize_context_keywords, normalize_tiers,
8    parse_utc_optional, sttp_node_to_json, to_json_string, tool_error, validate_limit,
9};
10
11pub(crate) async fn execute(server: &SttpMcpServer, request: GetContextRequest) -> String {
12    let from_utc = match parse_utc_optional(request.from_utc.as_deref(), "from_utc") {
13        Ok(value) => value,
14        Err(message) => return tool_error("InvalidDate", &message),
15    };
16    let to_utc = match parse_utc_optional(request.to_utc.as_deref(), "to_utc") {
17        Ok(value) => value,
18        Err(message) => return tool_error("InvalidDate", &message),
19    };
20
21    let tiers = request
22        .tiers
23        .as_ref()
24        .map(|values| normalize_tiers(values.as_slice()));
25
26    let limit = match validate_limit(request.limit, "limit") {
27        Ok(value) => value,
28        Err(message) => return tool_error("InvalidArgument", &message),
29    };
30    let context_keywords = normalize_context_keywords(request.context_keywords.as_deref());
31
32    if let Some(alpha) = request.alpha {
33        if !(0.0..=1.0).contains(&alpha) {
34            return tool_error("InvalidArgument", "alpha must be between 0.0 and 1.0");
35        }
36    }
37    if let Some(beta) = request.beta {
38        if !(0.0..=1.0).contains(&beta) {
39            return tool_error("InvalidArgument", "beta must be between 0.0 and 1.0");
40        }
41    }
42    if let Some(gamma) = request.gamma {
43        if !(0.0..=1.0).contains(&gamma) {
44            return tool_error("InvalidArgument", "gamma must be between 0.0 and 1.0");
45        }
46    }
47
48    let alpha = request.alpha.unwrap_or(0.7);
49    let beta = request.beta.unwrap_or(0.3);
50    let gamma = request.gamma.unwrap_or(0.0);
51    let query_text = if context_keywords.is_empty() {
52        None
53    } else {
54        Some(context_keywords.join(" "))
55    };
56    let query_embedding = if context_keywords.is_empty() {
57        None
58    } else {
59        server.embed_context_keywords(&context_keywords).await
60    };
61    let query_tag_embedding = if context_keywords.is_empty() {
62        None
63    } else if gamma > 0.0 {
64        server.embed_context_keywords(&context_keywords).await
65    } else {
66        None
67    };
68
69    let recall_service = MemoryRecallService::new(server.node_store.clone())
70        .with_semantic_index(server.semantic_index.clone());
71    let recall_result = match recall_service
72        .execute(&MemoryRecallRequest {
73            scope: MemoryScope {
74                tenant_id: None,
75                session_ids: request.session_id.map(|session| vec![session]),
76                tiers,
77                from_utc,
78                to_utc,
79            },
80            page: MemoryPage {
81                limit,
82                cursor: None,
83            },
84            scoring: MemoryScoring {
85                alpha,
86                beta,
87                gamma,
88                ..Default::default()
89            },
90            filter: MemoryFilter {
91                indexed_tags: request.semantic_tags,
92                link_rel: request.link_rel,
93                link_target: request.link_target,
94                ..Default::default()
95            },
96            current_avec: Some(locus_core_rs::AvecState {
97                stability: request.stability,
98                friction: request.friction,
99                logic: request.logic,
100                autonomy: request.autonomy,
101            }),
102            query_text,
103            query_embedding,
104            query_tag_embedding,
105            ..Default::default()
106        })
107        .await
108    {
109        Ok(result) => result,
110        Err(err) => {
111            error!(error = %err, "get_context failed");
112            return tool_error("GetContextFailure", &err.to_string());
113        }
114    };
115
116    to_json_string(json!({
117        "retrieved": recall_result.retrieved,
118        "psi_range": {
119            "min": recall_result.psi_range.min,
120            "max": recall_result.psi_range.max,
121            "average": recall_result.psi_range.average,
122        },
123        "nodes": recall_result
124            .nodes
125            .iter()
126            .map(sttp_node_to_json)
127            .collect::<Vec<_>>(),
128    }))
129}