use anyhow::{anyhow, Result};
use std::path::{Path, PathBuf};
use std::time::Duration;
use trusty_mcp as mcp;
const REQUEST_TIMEOUT: Duration = Duration::from_secs(60);
pub const STREAMING_METHODS: &[&str] = &["memory.chat", "memory.activity_stream"];
fn normalise_jsonrpc(mut req: serde_json::Value) -> serde_json::Value {
if let Some(obj) = req.as_object_mut() {
obj.insert(
"jsonrpc".to_string(),
serde_json::Value::String("2.0".to_string()),
);
}
req
}
fn effective_method(req: &serde_json::Value) -> Option<&str> {
let method = req.get("method")?.as_str()?;
if method == "tools/call" {
return req
.get("params")
.and_then(|p| p.get("name"))
.and_then(|n| n.as_str())
.or(Some(method));
}
Some(method)
}
pub(crate) async fn forward_rpc(
socket: &Path,
req: serde_json::Value,
) -> Result<serde_json::Value> {
if let Some(method) = effective_method(&req) {
if STREAMING_METHODS.contains(&method) {
return Ok(serde_json::json!({
"jsonrpc": "2.0",
"id": req.get("id").cloned().unwrap_or(serde_json::Value::Null),
"error": {
"code": mcp::error_codes::INVALID_REQUEST,
"message": format!(
"{method} answers as a stream, which MCP stdio cannot carry \
(one response per request). Dial {} directly with a framed \
streaming client to read it.",
socket.display()
),
},
}));
}
}
let req = normalise_jsonrpc(req);
let response: trusty_common::uds::server::RpcResponse =
trusty_common::uds::send_framed_request_capped(
socket,
&req,
REQUEST_TIMEOUT,
crate::transport::uds::MAX_FRAME_BYTES,
)
.await
.map_err(|e| anyhow!("connection to the trusty-memory daemon failed: {e}"))?;
serde_json::to_value(response).map_err(|e| anyhow!("re-encode the daemon response: {e}"))
}
pub(crate) async fn ensure_daemon_up_for_stdio() -> Result<PathBuf> {
let lock_path = crate::commands::start::start_lock_path()
.ok_or_else(|| anyhow!("could not resolve the trusty-memory data directory"))?;
let socket = crate::transport::uds::socket_path()?;
crate::commands::daemon_guard::ensure_daemon_running(&socket, &lock_path).await?;
Ok(socket)
}
fn is_notification(req: &mcp::Request) -> bool {
req.id.is_none() || req.method.starts_with("notifications/")
}
pub async fn run_stdio_bridge(palace: Option<String>) -> Result<()> {
let socket = ensure_daemon_up_for_stdio().await?;
let default_palace = palace;
let caller_cwd = std::env::current_dir()
.ok()
.map(|p| p.to_string_lossy().into_owned());
let caller_workstream = crate::attribution::resolve_own_workstream_name(caller_cwd.as_deref());
let result = mcp::run_stdio_loop(move |req| {
let socket = socket.clone();
let default_palace = default_palace.clone();
let caller_cwd = caller_cwd.clone();
let caller_workstream = caller_workstream.clone();
async move {
answer_one_request(
&socket,
req,
default_palace.as_deref(),
caller_workstream.as_deref(),
caller_cwd.as_deref(),
)
.await
}
})
.await;
result
}
pub(crate) async fn answer_one_request(
socket: &Path,
req: mcp::Request,
default_palace: Option<&str>,
caller_workstream: Option<&str>,
caller_cwd: Option<&str>,
) -> mcp::Response {
if is_notification(&req) {
return mcp::Response::suppressed();
}
let id = req.id.clone();
let req_value = inject_default_palace(req_to_value(&req), default_palace);
let req_value = inject_caller_context(req_value, caller_workstream, caller_cwd);
match forward_rpc(socket, req_value).await {
Ok(resp_value) => value_to_mcp_response(resp_value),
Err(e) => {
tracing::warn!("daemon bridge: transport error: {e:#}");
mcp::Response::err(
id,
mcp::error_codes::INTERNAL_ERROR,
format!("trusty-memory daemon unreachable: {e:#}"),
)
}
}
}
fn req_to_value(req: &mcp::Request) -> serde_json::Value {
serde_json::to_value(req).unwrap_or_else(|_| serde_json::json!({}))
}
fn inject_default_palace(
mut req: serde_json::Value,
default_palace: Option<&str>,
) -> serde_json::Value {
let Some(palace) = default_palace else {
return req;
};
let is_tools_call = req.get("method").and_then(|m| m.as_str()) == Some("tools/call");
let params = match req.get_mut("params") {
Some(p) if p.is_object() => p,
Some(p) if p.is_null() => {
*p = serde_json::json!({});
p
}
None => {
req["params"] = serde_json::json!({});
req.get_mut("params").expect("just inserted")
}
_ => return req,
};
let target = if is_tools_call {
match params.get_mut("arguments") {
Some(a) if a.is_object() => a,
Some(a) if a.is_null() => {
*a = serde_json::json!({});
a
}
None => {
params["arguments"] = serde_json::json!({});
params.get_mut("arguments").expect("just inserted")
}
_ => return req,
}
} else {
params
};
if target.get("palace").is_none() {
target["palace"] = serde_json::Value::String(palace.to_string());
}
req
}
fn inject_caller_context(
mut req: serde_json::Value,
workstream: Option<&str>,
cwd: Option<&str>,
) -> serde_json::Value {
if workstream.is_none() && cwd.is_none() {
return req;
}
let is_tools_call = req.get("method").and_then(|m| m.as_str()) == Some("tools/call");
let params = match req.get_mut("params") {
Some(p) if p.is_object() => p,
Some(p) if p.is_null() => {
*p = serde_json::json!({});
p
}
None => {
req["params"] = serde_json::json!({});
req.get_mut("params").expect("just inserted")
}
_ => return req,
};
let args = if is_tools_call {
match params.get_mut("arguments") {
Some(a) if a.is_object() => a,
Some(a) if a.is_null() => {
*a = serde_json::json!({});
a
}
None => {
params["arguments"] = serde_json::json!({});
params.get_mut("arguments").expect("just inserted")
}
_ => return req,
}
} else {
params
};
if let Some(ws) = workstream {
if args.get("workstream").is_none() {
args["workstream"] = serde_json::Value::String(ws.to_string());
}
}
if let Some(c) = cwd {
if args.get("cwd").is_none() {
args["cwd"] = serde_json::Value::String(c.to_string());
}
}
req
}
pub(crate) fn value_to_mcp_response(v: serde_json::Value) -> mcp::Response {
let id = v.get("id").cloned().filter(|id| !id.is_null());
if let Some(result) = v.get("result").cloned() {
return mcp::Response::ok(id, result);
}
if let Some(err) = v.get("error") {
let code = err
.get("code")
.and_then(|c| c.as_i64())
.map(|c| c as i32)
.unwrap_or(mcp::error_codes::INTERNAL_ERROR);
let message = err
.get("message")
.and_then(|m| m.as_str())
.unwrap_or("unknown daemon error")
.to_string();
return mcp::Response::err(id, code, &message);
}
mcp::Response::err(
id,
mcp::error_codes::INTERNAL_ERROR,
"daemon returned a response with neither result nor error",
)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn inject_default_palace_adds_when_absent() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "memory_remember",
"params": {"content": "hello"}
});
let out = inject_default_palace(req, Some("my-palace"));
assert_eq!(out["params"]["palace"], "my-palace");
assert_eq!(out["params"]["content"], "hello");
}
#[test]
fn inject_default_palace_preserves_existing() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "memory_remember",
"params": {"content": "hi", "palace": "caller-palace"}
});
let out = inject_default_palace(req, Some("default-palace"));
assert_eq!(out["params"]["palace"], "caller-palace");
}
#[test]
fn inject_default_palace_noop_when_none() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "memory_remember",
"params": {"content": "hi"}
});
let out = inject_default_palace(req.clone(), None);
assert_eq!(out, req);
}
#[test]
fn inject_default_palace_null_params_becomes_object() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "palace_list",
"params": null
});
let out = inject_default_palace(req, Some("my-palace"));
assert_eq!(out["params"]["palace"], "my-palace");
}
#[test]
fn inject_default_palace_tools_call_adds_when_absent() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": {
"name": "memory_recall",
"arguments": {"query": "hello"}
}
});
let out = inject_default_palace(req, Some("owner-profile"));
assert_eq!(out["params"]["arguments"]["palace"], "owner-profile");
assert_eq!(out["params"]["arguments"]["query"], "hello");
assert!(
out["params"].get("palace").is_none(),
"must not land as a sibling of name/arguments"
);
}
#[test]
fn inject_default_palace_tools_call_preserves_existing() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": {
"name": "memory_recall",
"arguments": {"query": "hi", "palace": "caller-palace"}
}
});
let out = inject_default_palace(req, Some("default-palace"));
assert_eq!(out["params"]["arguments"]["palace"], "caller-palace");
}
#[test]
fn inject_default_palace_legacy_direct_shape_still_injects_top_level() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "memory_recall",
"params": {"query": "hello"}
});
let out = inject_default_palace(req, Some("owner-profile"));
assert_eq!(out["params"]["palace"], "owner-profile");
assert_eq!(out["params"]["query"], "hello");
}
#[test]
fn inject_caller_context_direct_dispatch_shape() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "memory_remember",
"params": {"text": "hello"}
});
let out = inject_caller_context(req, Some("feat-x"), Some("/x/.worktrees/feat-x"));
assert_eq!(out["params"]["workstream"], "feat-x");
assert_eq!(out["params"]["cwd"], "/x/.worktrees/feat-x");
assert_eq!(out["params"]["text"], "hello");
}
#[test]
fn inject_caller_context_tools_call_shape() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": {
"name": "memory_remember",
"arguments": {"text": "hello"}
}
});
let out = inject_caller_context(req, Some("feat-x"), Some("/x/.worktrees/feat-x"));
assert_eq!(out["params"]["arguments"]["workstream"], "feat-x");
assert_eq!(out["params"]["arguments"]["cwd"], "/x/.worktrees/feat-x");
assert_eq!(out["params"]["arguments"]["text"], "hello");
assert!(
out["params"].get("workstream").is_none(),
"must not land as a sibling of name/arguments"
);
}
#[test]
fn inject_caller_context_preserves_existing_caller_values() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": {
"name": "memory_remember",
"arguments": {"text": "hello", "workstream": "caller-ws", "cwd": "/caller/cwd"}
}
});
let out = inject_caller_context(req, Some("bridge-ws"), Some("/bridge/cwd"));
assert_eq!(out["params"]["arguments"]["workstream"], "caller-ws");
assert_eq!(out["params"]["arguments"]["cwd"], "/caller/cwd");
}
#[test]
fn inject_caller_context_noop_when_nothing_resolved() {
let req = json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": {"name": "memory_remember", "arguments": {"text": "hello"}}
});
let out = inject_caller_context(req.clone(), None, None);
assert_eq!(out, req);
}
#[test]
fn value_to_mcp_response_variants() {
let ok = value_to_mcp_response(json!({"jsonrpc":"2.0","id":42,"result":{"tools":[]}}));
assert!(!ok.suppress);
assert_eq!(ok.id, Some(json!(42)));
assert!(ok.error.is_none());
let err = value_to_mcp_response(
json!({"jsonrpc":"2.0","id":7,"error":{"code":-32601,"message":"Not found"}}),
);
assert_eq!(err.error.unwrap().code, -32601);
let bad = value_to_mcp_response(json!({"jsonrpc":"2.0","id":1}));
assert_eq!(bad.error.unwrap().code, mcp::error_codes::INTERNAL_ERROR);
let null_id = value_to_mcp_response(json!({"jsonrpc":"2.0","id":null,"result":{}}));
assert_eq!(null_id.id, None);
}
#[test]
fn notification_requests_are_suppressed() {
let normal = mcp::Request {
jsonrpc: Some("2.0".to_string()),
id: Some(json!(1)),
method: "tools/list".to_string(),
params: None,
};
assert!(!is_notification(&normal));
let notif = mcp::Request {
jsonrpc: Some("2.0".to_string()),
id: None,
method: "notifications/initialized".to_string(),
params: None,
};
assert!(is_notification(¬if));
let notif_with_id = mcp::Request {
jsonrpc: Some("2.0".to_string()),
id: Some(json!(99)),
method: "notifications/cancelled".to_string(),
params: None,
};
assert!(is_notification(¬if_with_id));
}
#[test]
fn forwarded_request_carries_jsonrpc_two_point_zero() {
let req = mcp::Request {
jsonrpc: None,
id: Some(json!(1)),
method: "tools/list".to_string(),
params: None,
};
let raw = req_to_value(&req);
assert!(
raw["jsonrpc"].is_null(),
"the fixture must reproduce the null this fix exists for, got {raw}"
);
let forwarded = normalise_jsonrpc(raw);
assert_eq!(forwarded["jsonrpc"], "2.0");
assert_eq!(forwarded["method"], "tools/list", "nothing else changes");
assert_eq!(forwarded["id"], json!(1));
}
#[test]
fn forwarded_request_normalises_an_absent_jsonrpc() {
let absent = normalise_jsonrpc(json!({"id": 1, "method": "ping"}));
assert_eq!(absent["jsonrpc"], "2.0");
let wrong = normalise_jsonrpc(json!({"jsonrpc": "1.0", "id": 1, "method": "ping"}));
assert_eq!(wrong["jsonrpc"], "2.0");
}
#[tokio::test]
async fn streaming_method_is_refused_rather_than_half_answered() {
let socket = std::path::Path::new("/nonexistent/trusty-memory.sock");
for req in [
json!({"jsonrpc": "2.0", "id": 1, "method": "memory.chat", "params": {}}),
json!({
"jsonrpc": "2.0", "id": 2, "method": "tools/call",
"params": {"name": "memory.chat", "arguments": {}}
}),
] {
let answer = forward_rpc(socket, req)
.await
.expect("the refusal is an answer, not a transport failure");
let message = answer["error"]["message"]
.as_str()
.expect("the refusal must carry a message");
assert!(
message.contains("stream"),
"the refusal must say why: {message}"
);
}
}
#[tokio::test]
async fn forward_reports_a_dead_socket_rather_than_hanging() {
let tmp = tempfile::tempdir().expect("tempdir");
let started = std::time::Instant::now();
let result = forward_rpc(
&tmp.path().join("absent.sock"),
json!({"jsonrpc": "2.0", "id": 1, "method": "ping"}),
)
.await;
assert!(result.is_err(), "no listener means no answer");
assert!(
started.elapsed() < Duration::from_secs(5),
"a refused dial must not wait out REQUEST_TIMEOUT: {:?}",
started.elapsed()
);
}
fn a_request_with_id(id: i64) -> mcp::Request {
mcp::Request {
jsonrpc: Some("2.0".to_string()),
id: Some(json!(id)),
method: "tools/call".to_string(),
params: Some(json!({
"name": "memory_remember",
"arguments": {"text": "anything"}
})),
}
}
#[tokio::test]
async fn a_transport_failure_answers_the_request_that_caused_it() {
let tmp = tempfile::tempdir().expect("tempdir");
let resp = answer_one_request(
&tmp.path().join("vanished.sock"),
a_request_with_id(7),
None,
None,
None,
)
.await;
assert!(resp.error.is_some(), "an unreachable daemon is an error");
assert_eq!(
resp.id,
Some(json!(7)),
"the error must carry the id of the request it answers, or the \
client never matches it and waits forever"
);
assert!(!resp.suppress, "a request with an id always gets a reply");
}
#[tokio::test]
async fn a_transport_failure_names_the_endpoint_and_does_not_hang() {
let tmp = tempfile::tempdir().expect("tempdir");
let socket = tmp.path().join("vanished.sock");
let started = std::time::Instant::now();
let resp = answer_one_request(&socket, a_request_with_id(11), None, None, None).await;
let elapsed = started.elapsed();
let message = resp
.error
.expect("an unreachable daemon is an error")
.message;
assert!(
message.contains(&socket.display().to_string()),
"the error must name the endpoint it could not reach, got: {message}"
);
assert!(
elapsed < Duration::from_secs(30),
"a gone endpoint must answer inside the backoff cap: {elapsed:?}"
);
}
#[tokio::test]
async fn answer_one_request_suppresses_a_notification() {
let tmp = tempfile::tempdir().expect("tempdir");
let req = mcp::Request {
jsonrpc: Some("2.0".to_string()),
id: None,
method: "notifications/initialized".to_string(),
params: None,
};
let resp = answer_one_request(&tmp.path().join("absent.sock"), req, None, None, None).await;
assert!(resp.suppress, "a notification must not be answered");
}
}