use anyhow::{anyhow, Result};
use std::path::PathBuf;
use std::time::Duration;
use trusty_mcp::{DaemonBridgeJsonRpc, UdsBridgeConfig};
const REQUEST_TIMEOUT: Duration = Duration::from_secs(60);
pub const STREAMING_METHODS: &[&str] = &["memory.chat", "memory.activity_stream"];
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)
}
pub(crate) fn build_bridge(
socket: PathBuf,
default_palace: Option<String>,
caller_workstream: Option<String>,
caller_cwd: Option<String>,
) -> DaemonBridgeJsonRpc {
let config = UdsBridgeConfig::new(socket, "trusty-memory")
.with_streaming_methods(STREAMING_METHODS.iter().copied())
.with_request_timeout(REQUEST_TIMEOUT)
.with_max_frame_bytes(crate::transport::uds::MAX_FRAME_BYTES)
.with_bridge_version(env!("CARGO_PKG_VERSION"));
let handshake_palace = default_palace.clone();
DaemonBridgeJsonRpc::new(config)
.with_local_handler(move |req| {
crate::commands::serve_stdio_local::local_answer(req, handshake_palace.as_deref())
})
.with_request_rewriter(move |envelope| {
let envelope = inject_default_palace(envelope, default_palace.as_deref());
inject_caller_context(
envelope,
caller_workstream.as_deref(),
caller_cwd.as_deref(),
)
})
}
pub async fn run_stdio_bridge(palace: Option<String>) -> Result<()> {
let socket = match ensure_daemon_up_for_stdio().await {
Ok(socket) => socket,
Err(cause) => {
eprintln!(
"trusty-memory: the daemon is not up ({cause:#}). Serving the MCP \
handshake from this process; tool calls answer with an error \
until the daemon returns, then succeed with no restart (#8351)."
);
crate::transport::uds::socket_path()
.unwrap_or_else(|_| PathBuf::from("trusty-memory.sock"))
}
};
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());
build_bridge(socket, palace, caller_workstream, caller_cwd)
.with_socket_resolver(crate::transport::uds::socket_path)
.run_stdio()
.await
}
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
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use trusty_mcp::Request;
fn dead_bridge(socket: PathBuf) -> DaemonBridgeJsonRpc {
build_bridge(socket, None, None, None)
}
fn a_request_with_id(id: i64) -> Request {
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"}
})),
}
}
#[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);
}
#[tokio::test]
async fn streaming_method_is_refused_rather_than_half_answered() {
let bridge = dead_bridge(PathBuf::from("/nonexistent/trusty-memory.sock"));
for req in [
Request {
jsonrpc: Some("2.0".to_string()),
id: Some(json!(1)),
method: "memory.chat".to_string(),
params: Some(json!({})),
},
Request {
jsonrpc: Some("2.0".to_string()),
id: Some(json!(2)),
method: "tools/call".to_string(),
params: Some(json!({"name": "memory.chat", "arguments": {}})),
},
] {
let resp = bridge.answer(req).await;
let message = resp
.error
.expect("the refusal is an answer, not a transport failure")
.message;
assert!(
message.contains("stream"),
"the refusal must say why: {message}"
);
}
}
#[tokio::test]
async fn the_handshake_is_answered_with_no_daemon_listening() {
let tmp = tempfile::tempdir().expect("tempdir");
let bridge = dead_bridge(tmp.path().join("nothing-here.sock"));
for (id, method) in [(1, "initialize"), (2, "tools/list")] {
let resp = bridge
.answer(Request {
jsonrpc: Some("2.0".to_string()),
id: Some(json!(id)),
method: method.to_string(),
params: None,
})
.await;
assert!(
resp.error.is_none(),
"{method} must not depend on the daemon: {:?}",
resp.error
);
assert!(resp.result.is_some(), "{method} must answer with a result");
assert_eq!(resp.id, Some(json!(id)), "{method} must be matchable");
}
}
#[tokio::test]
async fn the_local_tool_list_is_the_daemon_table() {
let tmp = tempfile::tempdir().expect("tempdir");
for palace in [None, Some("owner-profile".to_string())] {
let has_default = palace.is_some();
let bridge = build_bridge(tmp.path().join("nothing-here.sock"), palace, None, None);
let resp = bridge
.answer(Request {
jsonrpc: Some("2.0".to_string()),
id: Some(json!(1)),
method: "tools/list".to_string(),
params: None,
})
.await;
assert_eq!(
resp.result,
Some(crate::tools::tool_definitions_with(has_default)),
"the bridge must answer from the daemon's own table"
);
}
}
#[tokio::test]
async fn a_transport_failure_answers_the_request_that_caused_it() {
let tmp = tempfile::tempdir().expect("tempdir");
let resp = dead_bridge(tmp.path().join("vanished.sock"))
.answer(a_request_with_id(7))
.await;
assert!(resp.error.is_some(), "an unreachable daemon is an error");
assert!(
resp.result.is_none(),
"an unreachable daemon must never read as an empty result"
);
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 = dead_bridge(socket.clone())
.answer(a_request_with_id(11))
.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 a_notification_is_suppressed_without_dialling() {
let tmp = tempfile::tempdir().expect("tempdir");
let resp = dead_bridge(tmp.path().join("absent.sock"))
.answer(Request {
jsonrpc: Some("2.0".to_string()),
id: None,
method: "notifications/initialized".to_string(),
params: None,
})
.await;
assert!(resp.suppress, "a notification must not be answered");
}
#[tokio::test]
async fn the_rewriter_reaches_the_forwarded_envelope() {
let dir = tempfile::tempdir().expect("tempdir");
let socket = dir.path().join("sockets").join("echo.sock");
let listener = trusty_common::uds::bind_hardened(&socket).expect("bind the echo socket");
tokio::spawn(async move {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let Ok((mut conn, _)) = listener.accept().await else {
return;
};
let mut frame = Vec::new();
let _ = conn.read_to_end(&mut frame).await;
let seen: serde_json::Value = serde_json::from_slice(&frame).unwrap_or_default();
let body = json!({
"jsonrpc": "2.0",
"id": seen.get("id").cloned().unwrap_or_default(),
"result": {"seen": seen},
});
let mut bytes = serde_json::to_vec(&body).unwrap_or_default();
bytes.push(b'\n');
let _ = conn.write_all(&bytes).await;
let _ = conn.flush().await;
});
let bridge = build_bridge(
socket,
Some("owner-profile".to_string()),
Some("feat-x".to_string()),
Some("/x/.worktrees/feat-x".to_string()),
);
let resp = bridge.answer(a_request_with_id(3)).await;
let result = resp.result.expect("the echo server answered");
let seen = &result["seen"];
assert_eq!(
seen["jsonrpc"], "2.0",
"jsonrpc is re-stamped after the rewriter"
);
assert_eq!(seen["params"]["arguments"]["palace"], "owner-profile");
assert_eq!(seen["params"]["arguments"]["workstream"], "feat-x");
assert_eq!(seen["params"]["arguments"]["cwd"], "/x/.worktrees/feat-x");
assert_eq!(seen["params"]["arguments"]["text"], "anything");
}
}