locus_mcp/tools/
get_context.rs1use 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}