use rmcp::handler::server::wrapper::{Json, Parameters};
use rmcp::model::ErrorCode;
use tempfile::TempDir;
use super::super::dto::RememberParams;
use super::*;
use crate::context::{
fragment_id, ContextAction, ContextFact, ContextFragment, MemoryScope, WorkingContext,
};
use crate::embedder::{DynEmbedder, HashEmbedder};
use crate::service::MemoryService;
fn server() -> (TempDir, McpServer) {
let dir = TempDir::new().expect("create tempdir");
let embedder: DynEmbedder = Box::new(HashEmbedder::new(crate::DEFAULT_DIMENSION));
let service = MemoryService::open(dir.path(), embedder).expect("open memory store");
(dir, McpServer::new(service))
}
fn fragment(content: &str) -> ContextFragment {
ContextFragment {
id: None,
content: content.to_owned(),
kind: None,
priority: None,
metadata: None,
media: None,
}
}
fn request(query: &str, fragments: Vec<ContextFragment>, budget: u64) -> CompileRequest {
CompileRequest {
query: query.to_owned(),
fragments,
project: None,
target_model: None,
token_budget: budget,
memory_scope: None,
policy: None,
}
}
fn compiled_context_of(value: serde_json::Value) -> CompiledContext {
serde_json::from_value(value).expect("valid CompiledContext wire value")
}
fn decision_of(value: serde_json::Value) -> ContextDecision {
serde_json::from_value(value).expect("valid ContextDecision wire value")
}
#[tokio::test]
async fn test_compile_context_tool_returns_compiled_context_and_insights() {
let (_dir, srv) = server();
let req = request(
"deploy",
vec![fragment("a fact"), fragment("a fact")],
10_000,
);
let Json(value) = srv
.compile_context(Parameters(req))
.await
.expect("compile_context");
let out = compiled_context_of(value);
assert!(out.content.contains("a fact"));
assert_eq!(out.decisions.len(), 2);
assert!(out.insights.tokens_saved > 0, "the duplicate saves tokens");
}
#[tokio::test]
async fn test_compile_context_tool_pulls_memory_scope() {
let (_dir, srv) = server();
srv.remember(Parameters(RememberParams {
fact: "the deploy pipeline runs clippy before tests".to_owned(),
links: Vec::new(),
metadata: None,
ttl_seconds: None,
}))
.await
.expect("remember");
let mut req = request("deploy pipeline checks", vec![fragment("note")], 10_000);
req.memory_scope = Some(MemoryScope {
k: Some(3),
..MemoryScope::default()
});
let Json(value) = srv
.compile_context(Parameters(req))
.await
.expect("compile_context");
let out = compiled_context_of(value);
assert!(out.content.contains("runs clippy before tests"));
assert!(out.decisions.iter().any(|d| d.memory_id.is_some()));
}
#[tokio::test]
async fn test_context_savings_tool_aggregates_by_project() {
let (_dir, srv) = server();
for _ in 0..2 {
let mut req = request("deploy", vec![fragment("x"), fragment("x")], 10_000);
req.project = Some("veles".to_owned());
srv.compile_context(Parameters(req))
.await
.expect("compile_context");
}
let Json(savings) = srv
.context_savings(Parameters(ContextSavingsParams {
project: Some("veles".to_owned()),
}))
.await
.expect("context_savings");
assert_eq!(savings.events, 2);
assert!(savings.tokens_saved > 0);
}
#[tokio::test]
async fn test_explain_compilation_tool_returns_decision_for_fragment() {
let (_dir, srv) = server();
let req = request(
"deploy",
vec![fragment("a fact"), fragment("other")],
10_000,
);
let wanted = fragment_id("a fact");
let Json(value) = srv
.explain_compilation(Parameters(ExplainCompilationParams {
request: req,
fragment_id: wanted,
fragment_index: None,
}))
.await
.expect("explain_compilation");
let decision = decision_of(value);
assert_eq!(decision.fragment_id, wanted);
assert!(matches!(decision.action, ContextAction::Preserve));
assert!(!decision.reason.is_empty());
}
#[tokio::test]
async fn test_explain_compilation_tool_unknown_fragment_is_invalid_params() {
let (_dir, srv) = server();
let req = request("deploy", vec![fragment("a fact")], 10_000);
let Err(err) = srv
.explain_compilation(Parameters(ExplainCompilationParams {
request: req,
fragment_id: 424_242,
fragment_index: None,
}))
.await
else {
panic!("no such fragment in the request — the tool must fail");
};
assert_eq!(err.code, ErrorCode::INVALID_PARAMS);
}
const ID_ABOVE_JS_SAFE_INTEGER: u64 = 9_007_199_254_740_993;
#[tokio::test]
async fn test_compile_context_tool_ids_as_strings_stringifies_response_ids() {
let (_dir, srv) = server();
let mut fragment = fragment("a fact above the safe integer range");
fragment.id = Some(ID_ABOVE_JS_SAFE_INTEGER);
let mut req = request("deploy", vec![fragment], 10_000);
req.policy = Some(CompilePolicy {
ids_as_strings: true,
..CompilePolicy::default()
});
let Json(value) = srv
.compile_context(Parameters(req))
.await
.expect("compile_context");
let decision_id = &value["decisions"][0]["fragment_id"];
assert_eq!(
decision_id.as_str(),
Some(ID_ABOVE_JS_SAFE_INTEGER.to_string().as_str()),
"fragment_id must be a JSON string when ids_as_strings is active: {value}"
);
assert!(
!decision_id.is_number(),
"fragment_id must not still be a JSON number: {value}"
);
}
#[tokio::test]
async fn test_compile_context_tool_ids_as_strings_default_false_is_byte_identical() {
let (_dir, srv) = server();
let fragment_a = {
let mut f = fragment("a fact above the safe integer range");
f.id = Some(ID_ABOVE_JS_SAFE_INTEGER);
f
};
let fragment_b = fragment_a.clone();
let req_default = request("deploy", vec![fragment_a], 10_000);
let mut req_explicit_false = request("deploy", vec![fragment_b], 10_000);
req_explicit_false.policy = Some(CompilePolicy {
ids_as_strings: false,
..CompilePolicy::default()
});
let Json(default_value) = srv
.compile_context(Parameters(req_default))
.await
.expect("compile_context (default policy)");
let Json(explicit_value) = srv
.compile_context(Parameters(req_explicit_false))
.await
.expect("compile_context (ids_as_strings: false)");
assert!(default_value["decisions"][0]["fragment_id"].is_number());
assert_eq!(default_value, explicit_value);
}
#[tokio::test]
async fn test_explain_compilation_tool_ids_as_strings_stringifies_response_ids() {
let (_dir, srv) = server();
let mut fragment = fragment("a fact above the safe integer range");
fragment.id = Some(ID_ABOVE_JS_SAFE_INTEGER);
let mut req = request("deploy", vec![fragment], 10_000);
req.policy = Some(CompilePolicy {
ids_as_strings: true,
..CompilePolicy::default()
});
let Json(value) = srv
.explain_compilation(Parameters(ExplainCompilationParams {
request: req,
fragment_id: ID_ABOVE_JS_SAFE_INTEGER,
fragment_index: None,
}))
.await
.expect("explain_compilation");
assert_eq!(
value["fragment_id"].as_str(),
Some(ID_ABOVE_JS_SAFE_INTEGER.to_string().as_str())
);
assert!(value["content_hash"].is_string());
}
#[tokio::test]
async fn test_compile_context_tool_accepts_fragment_id_as_decimal_string_on_input() {
let (_dir, srv) = server();
let mut req_value = serde_json::to_value(request(
"deploy",
vec![fragment("a fact above the safe integer range")],
10_000,
))
.expect("serialize request");
req_value["fragments"][0]["id"] =
serde_json::Value::String(ID_ABOVE_JS_SAFE_INTEGER.to_string());
let req: CompileRequest =
serde_json::from_value(req_value).expect("fragment id accepts a decimal string");
let Json(value) = srv
.compile_context(Parameters(req))
.await
.expect("compile_context");
assert_eq!(
value["decisions"][0]["fragment_id"].as_u64(),
Some(ID_ABOVE_JS_SAFE_INTEGER)
);
}
fn collect_id_property_types(
value: &serde_json::Value,
keys: &[&str],
found: &mut Vec<(String, serde_json::Value)>,
) {
match value {
serde_json::Value::Object(map) => {
if let Some(serde_json::Value::Object(properties)) = map.get("properties") {
for (name, subschema) in properties {
if keys.contains(&name.as_str()) {
let leaf = if subschema.get("type") == Some(&serde_json::json!("array")) {
subschema
.get("items")
.unwrap_or_else(|| panic!("array property {name} declares items"))
} else {
subschema
};
found.push((name.clone(), leaf["type"].clone()));
}
}
}
for entry in map.values() {
collect_id_property_types(entry, keys, found);
}
}
serde_json::Value::Array(items) => {
for item in items {
collect_id_property_types(item, keys, found);
}
}
_ => {}
}
}
fn assert_ids_widened(found: &[(String, serde_json::Value)]) {
for (name, type_value) in found {
let types = type_value
.as_array()
.unwrap_or_else(|| panic!("{name} must type a list of forms, got {type_value}"));
assert!(
types.contains(&serde_json::json!("integer"))
&& types.contains(&serde_json::json!("string")),
"{name} must advertise integer|string on the wire, got {type_value}"
);
}
}
#[test]
fn test_compile_context_output_schema_advertises_string_ids() {
let tool = McpServer::compile_context_tool_attr();
let schema = serde_json::to_value(
tool.output_schema
.expect("compile_context declares an output schema"),
)
.expect("schema serializes");
let mut found = Vec::new();
collect_id_property_types(&schema, crate::context::wire::ID_KEYS, &mut found);
let names: std::collections::BTreeSet<&str> =
found.iter().map(|(name, _)| name.as_str()).collect();
for expected in ["fragment_id", "content_hash", "memory_id", "fragment_ids"] {
assert!(
names.contains(expected),
"the compile_context output schema must carry {expected}; found {names:?}"
);
}
assert_ids_widened(&found);
}
#[test]
fn test_explain_compilation_output_schema_advertises_string_ids() {
let tool = McpServer::explain_compilation_tool_attr();
let schema = serde_json::to_value(
tool.output_schema
.expect("explain_compilation declares an output schema"),
)
.expect("schema serializes");
let mut found = Vec::new();
collect_id_property_types(&schema, crate::context::wire::ID_KEYS, &mut found);
let names: std::collections::BTreeSet<&str> =
found.iter().map(|(name, _)| name.as_str()).collect();
for expected in ["fragment_id", "content_hash", "memory_id"] {
assert!(
names.contains(expected),
"the explain_compilation output schema must carry {expected}; found {names:?}"
);
}
assert_ids_widened(&found);
}
#[test]
fn test_compile_context_input_schema_advertises_string_fragment_id() {
let tool = McpServer::compile_context_tool_attr();
let schema = serde_json::to_value(&tool.input_schema).expect("schema serializes");
let id_type = &schema["$defs"]["ContextFragment"]["properties"]["id"]["type"];
let types = id_type
.as_array()
.unwrap_or_else(|| panic!("fragments[].id must type a list of forms, got {id_type}"));
for expected in ["integer", "string", "null"] {
assert!(
types.contains(&serde_json::json!(expected)),
"fragments[].id must advertise {expected} on input, got {id_type}"
);
}
}
#[test]
fn test_explain_compilation_input_schema_keeps_top_level_fragment_id_strict() {
let tool = McpServer::explain_compilation_tool_attr();
let schema = serde_json::to_value(&tool.input_schema).expect("schema serializes");
assert_eq!(
schema["properties"]["fragment_id"]["type"],
serde_json::json!("integer"),
"top-level fragment_id stays integer-only"
);
let id_type = &schema["$defs"]["ContextFragment"]["properties"]["id"]["type"];
let types = id_type
.as_array()
.unwrap_or_else(|| panic!("fragments[].id must type a list of forms, got {id_type}"));
assert!(
types.contains(&serde_json::json!("string")),
"request.fragments[].id must advertise string on input, got {id_type}"
);
}
#[tokio::test]
async fn test_explain_compilation_tool_fragment_index_disambiguates_byte_identical_twins() {
let (_dir, srv) = server();
let req = request(
"deploy",
vec![fragment("duplicate payload"), fragment("duplicate payload")],
10_000,
);
let shared_id = fragment_id("duplicate payload");
let Json(survivor_value) = srv
.explain_compilation(Parameters(ExplainCompilationParams {
request: req.clone(),
fragment_id: shared_id,
fragment_index: None,
}))
.await
.expect("explain_compilation (by id)");
let survivor = decision_of(survivor_value);
let Json(twin_value) = srv
.explain_compilation(Parameters(ExplainCompilationParams {
request: req,
fragment_id: shared_id,
fragment_index: Some(1),
}))
.await
.expect("explain_compilation (by index)");
let twin = decision_of(twin_value);
assert!(matches!(survivor.action, ContextAction::Preserve));
assert!(matches!(twin.action, ContextAction::Drop));
assert_eq!(twin.rule_id, "drop.duplicate");
assert_eq!(twin.fragment_id, shared_id);
}
#[tokio::test]
async fn test_explain_compilation_tool_fragment_index_out_of_bounds_is_invalid_params() {
let (_dir, srv) = server();
let req = request("deploy", vec![fragment("a fact")], 10_000);
let wanted = fragment_id("a fact");
let Err(err) = srv
.explain_compilation(Parameters(ExplainCompilationParams {
request: req,
fragment_id: wanted,
fragment_index: Some(5),
}))
.await
else {
panic!("fragment_index 5 has no fragment — the tool must fail");
};
assert_eq!(err.code, ErrorCode::INVALID_PARAMS);
assert!(err.message.contains("fragment_index"));
}
#[tokio::test]
async fn test_retrieve_context_source_tool_round_trips_original() {
let (_dir, srv) = server();
let original = "Never restart the primary node during a rebalance.";
let req = request("rebalance", vec![fragment(original)], 10_000);
let Json(value) = srv
.compile_context(Parameters(req))
.await
.expect("compile_context");
let out = compiled_context_of(value);
let handle = out.sources[0].handle.clone();
let Json(retrieved) = srv
.retrieve_context_source(Parameters(RetrieveContextSourceParams {
handle: handle.clone(),
}))
.await
.expect("retrieve_context_source");
assert_eq!(retrieved.content, original);
assert_eq!(retrieved.handle, handle);
}
const PNG_B64: &str = "iVBORw0KGgoAAAANSUhEUgAAAEAAAAAwCAYAAAAAAAAA";
fn media_fragment(caption: &str) -> ContextFragment {
ContextFragment {
media: Some(crate::context::MediaRef {
mime: "image/png".to_owned(),
bytes_b64: PNG_B64.to_owned(),
}),
..fragment(caption)
}
}
#[tokio::test]
async fn test_retrieve_context_source_tool_round_trips_media_byte_identical() {
let (_dir, srv) = server();
let req = request("a screenshot", vec![media_fragment("a screenshot")], 1);
let Json(value) = srv
.compile_context(Parameters(req))
.await
.expect("compile_context");
let out = compiled_context_of(value);
let handle = out
.retrieval_handles
.first()
.expect("the oversized media fragment must externalize")
.handle
.clone();
let Json(retrieved) = srv
.retrieve_context_source(Parameters(RetrieveContextSourceParams {
handle: handle.clone(),
}))
.await
.expect("retrieve_context_source");
let media = retrieved
.media
.expect("a media source must carry its media back through the MCP tool");
assert_eq!(media.mime, "image/png");
assert_eq!(media.bytes_b64, PNG_B64);
}
#[tokio::test]
async fn test_retrieve_context_source_tool_text_only_carries_no_media() {
let (_dir, srv) = server();
let req = request("plain", vec![fragment("no picture here")], 10_000);
let Json(value) = srv
.compile_context(Parameters(req))
.await
.expect("compile_context");
let out = compiled_context_of(value);
let handle = out.sources[0].handle.clone();
let Json(retrieved) = srv
.retrieve_context_source(Parameters(RetrieveContextSourceParams { handle }))
.await
.expect("retrieve_context_source");
assert!(retrieved.media.is_none());
}
#[test]
fn test_compile_context_input_schema_advertises_fragment_media_field() {
let tool = McpServer::compile_context_tool_attr();
let schema = serde_json::to_value(&tool.input_schema).expect("schema serializes");
let media_property = &schema["$defs"]["ContextFragment"]["properties"]["media"];
assert!(
!media_property.is_null(),
"fragments[].media must be advertised on compile_context's input schema"
);
let media_schema = if let Some(reference) = media_property.get("$ref") {
let reference = reference
.as_str()
.expect("$ref is a string")
.trim_start_matches("#/$defs/");
&schema["$defs"][reference]
} else if let Some(one_of) = media_property
.get("anyOf")
.or_else(|| media_property.get("oneOf"))
{
one_of
.as_array()
.and_then(|variants| variants.iter().find(|v| v.get("$ref").is_some()))
.and_then(|v| v.get("$ref"))
.and_then(serde_json::Value::as_str)
.map(|r| r.trim_start_matches("#/$defs/"))
.map_or(media_property, |name| &schema["$defs"][name])
} else {
media_property
};
for expected in ["mime", "bytes_b64"] {
assert!(
!media_schema["properties"][expected].is_null(),
"MediaRef must advertise '{expected}'; media schema was {media_schema}"
);
}
}
#[test]
fn test_retrieve_context_source_output_schema_advertises_optional_media_field() {
let tool = McpServer::retrieve_context_source_tool_attr();
let schema = serde_json::to_value(
tool.output_schema
.expect("retrieve_context_source declares an output schema"),
)
.expect("schema serializes");
let media_property = &schema["properties"]["media"];
assert!(
!media_property.is_null(),
"retrieve_context_source's output schema must advertise 'media'; schema was {schema}"
);
let required = schema["required"]
.as_array()
.expect("output schema declares required properties");
assert!(
!required.contains(&serde_json::json!("media")),
"media must stay optional on the advertised schema, got required: {required:?}"
);
}
#[tokio::test]
async fn test_compile_context_tool_zero_budget_is_invalid_params() {
let (_dir, srv) = server();
let req = request("deploy", vec![fragment("anything")], 0);
let Err(err) = srv.compile_context(Parameters(req)).await else {
panic!("a zero budget cannot compile — the tool must fail");
};
assert_eq!(err.code, ErrorCode::INVALID_PARAMS);
}
#[tokio::test]
async fn test_retrieve_context_source_tool_unknown_handle_is_invalid_params() {
let (_dir, srv) = server();
let Err(err) = srv
.retrieve_context_source(Parameters(RetrieveContextSourceParams {
handle: "ctx://source/999999".to_owned(),
}))
.await
else {
panic!("nothing stored under this handle — the tool must fail");
};
assert_eq!(err.code, ErrorCode::INVALID_PARAMS);
}
fn working() -> WorkingContext {
WorkingContext {
goal: Some("ship EPIC-P-071 PR3".to_owned()),
active_constraints: vec![ContextFact {
text: "never merge without green gates".to_owned(),
source: None,
}],
verified_facts: vec![ContextFact {
text: "compile_context already ships on MCP+Node".to_owned(),
source: None,
}],
open_hypotheses: Vec::new(),
decisions: Vec::new(),
exact_evidence: Vec::new(),
pending_actions: vec!["wire save/load working-context tools".to_owned()],
}
}
#[tokio::test]
async fn test_save_working_context_tool_then_load_round_trips() {
let (_dir, srv) = server();
let saved = working();
let Json(save_result) = srv
.save_working_context(Parameters(SaveWorkingContextParams {
project: "veles".to_owned(),
session: "session-1".to_owned(),
working: saved.clone(),
}))
.await
.expect("save_working_context");
assert!(save_result.id > 0);
let Json(loaded) = srv
.load_working_context(Parameters(LoadWorkingContextParams {
project: "veles".to_owned(),
session: "session-1".to_owned(),
}))
.await
.expect("load_working_context");
let recovered = loaded
.working
.expect("a previously saved working context must load back");
assert_eq!(recovered.goal, saved.goal);
assert_eq!(recovered.pending_actions, saved.pending_actions);
assert_eq!(recovered.active_constraints.len(), 1);
}
#[tokio::test]
async fn test_load_working_context_tool_none_when_never_saved() {
let (_dir, srv) = server();
let Json(loaded) = srv
.load_working_context(Parameters(LoadWorkingContextParams {
project: "veles".to_owned(),
session: "never-saved".to_owned(),
}))
.await
.expect("load_working_context");
assert!(loaded.working.is_none());
}
#[tokio::test]
async fn test_save_working_context_tool_is_idempotent_upsert() {
let (_dir, srv) = server();
let mut state = working();
srv.save_working_context(Parameters(SaveWorkingContextParams {
project: "veles".to_owned(),
session: "session-2".to_owned(),
working: state.clone(),
}))
.await
.expect("save_working_context");
state.goal = Some("ship a follow-up PR".to_owned());
srv.save_working_context(Parameters(SaveWorkingContextParams {
project: "veles".to_owned(),
session: "session-2".to_owned(),
working: state.clone(),
}))
.await
.expect("save_working_context (replace)");
let Json(loaded) = srv
.load_working_context(Parameters(LoadWorkingContextParams {
project: "veles".to_owned(),
session: "session-2".to_owned(),
}))
.await
.expect("load_working_context");
assert_eq!(loaded.working.expect("saved").goal, state.goal);
}