use std::path::{Path, PathBuf};
use std::time::Duration;
use anyhow::{Context, Result, anyhow};
use serde_json::{Value, json};
use crate::uds::send_framed_request_capped;
use crate::uds::server::RpcResponse;
pub const TRUSTY_MEMORY_SOCKET_ENV: &str = "TRUSTY_MEMORY_SOCKET";
const MEMORY_APP_NAME: &str = "trusty-memory";
pub const MAX_FRAME_BYTES: u64 = 32 * 1024 * 1024;
pub const DEFAULT_TIMEOUT: Duration = Duration::from_secs(5);
pub const CODE_NOT_FOUND: i64 = -32004;
#[derive(Debug, thiserror::Error)]
#[error("{method} failed: {message} ({code})")]
#[non_exhaustive]
pub struct MemoryRpcError {
pub method: String,
pub code: i64,
pub message: String,
pub data: Option<Value>,
}
impl MemoryRpcError {
pub fn is_not_found(&self) -> bool {
self.code == CODE_NOT_FOUND
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum MemoryHealthStatus {
Ok,
Degraded,
Wedged,
Unrecognised(String),
Missing,
}
impl MemoryHealthStatus {
pub fn from_health_body(body: &Value) -> Self {
match body.get("status").and_then(Value::as_str) {
Some("ok") => Self::Ok,
Some("degraded") => Self::Degraded,
Some("wedged") => Self::Wedged,
Some(other) => Self::Unrecognised(other.to_string()),
None => Self::Missing,
}
}
pub fn is_ok(&self) -> bool {
matches!(self, Self::Ok)
}
}
const UNREACHABLE_PLACEHOLDER: &str = "/nonexistent/trusty-memory/trusty-memory.sock";
pub fn resolve_memory_socket() -> Result<PathBuf> {
if let Ok(raw) = std::env::var(TRUSTY_MEMORY_SOCKET_ENV) {
let trimmed = raw.trim();
if !trimmed.is_empty() {
return Ok(PathBuf::from(trimmed));
}
}
crate::daemon_addr::daemon_socket_path(MEMORY_APP_NAME)
}
pub fn resolve_memory_socket_or_unreachable() -> PathBuf {
resolve_memory_socket().unwrap_or_else(|e| {
eprintln!("trusty-memory: {e}");
PathBuf::from(UNREACHABLE_PLACEHOLDER)
})
}
pub async fn call_memory_tool(method: &str, params: Value) -> Result<Value> {
let socket = resolve_memory_socket()?;
call_memory_tool_at(&socket, method, params).await
}
pub async fn call_memory_tool_at(socket: &Path, method: &str, params: Value) -> Result<Value> {
call_memory_tool_at_with_timeout(socket, method, params, DEFAULT_TIMEOUT).await
}
pub async fn call_memory_tool_at_with_timeout(
socket: &Path,
method: &str,
params: Value,
timeout: Duration,
) -> Result<Value> {
let request = json!({
"jsonrpc": "2.0",
"id": 1,
"method": method,
"params": params,
});
let response: RpcResponse =
send_framed_request_capped(socket, &request, timeout, MAX_FRAME_BYTES)
.await
.with_context(|| {
format!(
"call {method} on the trusty-memory daemon at {}",
socket.display()
)
})?;
match (response.result, response.error) {
(Some(result), _) => Ok(result),
(None, Some(e)) => Err(anyhow::Error::new(MemoryRpcError {
method: method.to_string(),
code: e.code,
message: e.message,
data: e.data,
})),
(None, None) => Err(anyhow!(
"{method} answered with neither a result nor an error"
)),
}
}
pub async fn memory_daemon_is_serving(socket: &Path, timeout: Duration) -> bool {
crate::uds::socket_is_serving(socket, timeout).await
}
pub const METHOD_PROTOCOL: &str = "memory.protocol";
pub const SUPPORTED_MEMORY_PROTOCOLS: std::ops::RangeInclusive<u64> = 1..=1;
const PROTOCOL_RECHECK_INTERVAL: Duration = Duration::from_secs(10);
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[non_exhaustive]
pub struct MemoryProtocolInfo {
pub protocol_version: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub daemon_version: Option<String>,
}
impl MemoryProtocolInfo {
pub fn new(protocol_version: u64, daemon_version: Option<String>) -> Self {
Self {
protocol_version,
daemon_version,
}
}
}
const PRE_HANDSHAKE_PROTOCOL: u64 = 1;
fn pre_handshake_verdict(
socket: &Path,
supported: &std::ops::RangeInclusive<u64>,
) -> Result<MemoryProtocol, MemoryProtocolError> {
if supported.contains(&PRE_HANDSHAKE_PROTOCOL) {
return Ok(MemoryProtocol::PreHandshake);
}
Err(MemoryProtocolError::Unsupported {
socket: socket.to_path_buf(),
daemon: PRE_HANDSHAKE_PROTOCOL,
daemon_version: "older than protocol 1".to_string(),
min: *supported.start(),
max: *supported.end(),
})
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum MemoryProtocol {
Supported(MemoryProtocolInfo),
PreHandshake,
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum MemoryProtocolError {
#[error(
"unsupported trusty-memory protocol: the daemon at {socket} (version {daemon_version}) \
speaks protocol {daemon}, this client supports {min}..={max}; restart the daemon so \
it runs the installed release"
)]
Unsupported {
socket: PathBuf,
daemon: u64,
daemon_version: String,
min: u64,
max: u64,
},
#[error("the trusty-memory daemon at {socket} answered {METHOD_PROTOCOL} unreadably: {reason}")]
MalformedHandshake {
socket: PathBuf,
reason: String,
},
#[error("the trusty-memory protocol check against {socket} failed: {cause}")]
HandshakeFailed {
socket: PathBuf,
cause: String,
},
}
pub async fn check_memory_protocol_at(
socket: &Path,
timeout: Duration,
) -> Result<MemoryProtocol, MemoryProtocolError> {
let body =
match call_memory_tool_at_with_timeout(socket, METHOD_PROTOCOL, json!({}), timeout).await {
Ok(body) => body,
Err(e) => {
let pre_handshake = e
.downcast_ref::<MemoryRpcError>()
.is_some_and(|rpc| rpc.code == crate::uds::server::CODE_METHOD_NOT_FOUND);
if pre_handshake {
return pre_handshake_verdict(socket, &SUPPORTED_MEMORY_PROTOCOLS);
}
return Err(MemoryProtocolError::HandshakeFailed {
socket: socket.to_path_buf(),
cause: format!("{e:#}"),
});
}
};
let info: MemoryProtocolInfo =
serde_json::from_value(body).map_err(|e| MemoryProtocolError::MalformedHandshake {
socket: socket.to_path_buf(),
reason: e.to_string(),
})?;
if !SUPPORTED_MEMORY_PROTOCOLS.contains(&info.protocol_version) {
return Err(MemoryProtocolError::Unsupported {
socket: socket.to_path_buf(),
daemon: info.protocol_version,
daemon_version: info.daemon_version.unwrap_or_else(|| "unknown".to_string()),
min: *SUPPORTED_MEMORY_PROTOCOLS.start(),
max: *SUPPORTED_MEMORY_PROTOCOLS.end(),
});
}
Ok(MemoryProtocol::Supported(info))
}
type ProtocolCache = std::collections::HashMap<PathBuf, (std::time::Instant, MemoryProtocol)>;
fn protocol_cache() -> &'static std::sync::Mutex<ProtocolCache> {
static CACHE: std::sync::OnceLock<std::sync::Mutex<ProtocolCache>> = std::sync::OnceLock::new();
CACHE.get_or_init(Default::default)
}
pub async fn ensure_memory_protocol_at(
socket: &Path,
timeout: Duration,
) -> Result<MemoryProtocol, MemoryProtocolError> {
let cached = protocol_cache()
.lock()
.unwrap_or_else(|e| e.into_inner())
.get(socket)
.filter(|(at, _)| at.elapsed() < PROTOCOL_RECHECK_INTERVAL)
.map(|(_, verdict)| verdict.clone());
if let Some(verdict) = cached {
return Ok(verdict);
}
let verdict = check_memory_protocol_at(socket, timeout).await?;
if verdict == MemoryProtocol::PreHandshake {
warn_pre_handshake_once(socket, timeout).await;
}
protocol_cache()
.lock()
.unwrap_or_else(|e| e.into_inner())
.insert(
socket.to_path_buf(),
(std::time::Instant::now(), verdict.clone()),
);
Ok(verdict)
}
static PRE_HANDSHAKE_WARNINGS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
async fn warn_pre_handshake_once(socket: &Path, timeout: Duration) {
use std::sync::atomic::Ordering;
if PRE_HANDSHAKE_WARNINGS
.compare_exchange(0, 1, Ordering::SeqCst, Ordering::SeqCst)
.is_err()
{
return;
}
let version = call_memory_tool_at_with_timeout(socket, "memory.health", json!({}), timeout)
.await
.ok()
.and_then(|body| {
body.get("version")
.and_then(Value::as_str)
.map(str::to_string)
});
let daemon = match version {
Some(v) => format!("the trusty-memory daemon {v}"),
None => "a trusty-memory daemon older than protocol 1".to_string(),
};
let warning = format!(
"{daemon} at {} predates the protocol handshake (#9288); calling it as protocol 1. \
Restart the trusty-memory daemon to pick up the protocol handshake.",
socket.display()
);
eprintln!("trusty-memory: {warning}");
tracing::warn!("{warning}");
}
#[cfg(test)]
#[path = "memory_rpc_protocol_tests.rs"]
mod protocol_tests;
#[cfg(test)]
mod tests {
use super::*;
use crate::data_dir::ENV_LOCK;
#[test]
fn resolve_memory_socket_honours_the_env_override() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
unsafe {
std::env::set_var(TRUSTY_MEMORY_SOCKET_ENV, "/tmp/example-memory.sock");
}
let resolved = resolve_memory_socket();
unsafe {
std::env::remove_var(TRUSTY_MEMORY_SOCKET_ENV);
}
assert_eq!(
resolved.expect("an override always resolves"),
PathBuf::from("/tmp/example-memory.sock")
);
}
#[test]
fn resolve_memory_socket_or_unreachable_falls_back() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let tmp = tempfile::NamedTempFile::new().expect("temp file");
unsafe {
std::env::remove_var(TRUSTY_MEMORY_SOCKET_ENV);
std::env::set_var(
crate::data_dir::DATA_DIR_OVERRIDE_ENV,
tmp.path().join("under-a-file"),
);
}
let resolved = resolve_memory_socket_or_unreachable();
unsafe {
std::env::remove_var(crate::data_dir::DATA_DIR_OVERRIDE_ENV);
}
assert_eq!(resolved, PathBuf::from(UNREACHABLE_PLACEHOLDER));
}
#[test]
fn health_status_reads_every_handler_string() {
for (status, expected, healthy) in [
("ok", MemoryHealthStatus::Ok, true),
("degraded", MemoryHealthStatus::Degraded, false),
("wedged", MemoryHealthStatus::Wedged, false),
] {
let read = MemoryHealthStatus::from_health_body(&json!({ "status": status }));
assert_eq!(read, expected, "status {status}");
assert_eq!(read.is_ok(), healthy, "status {status}");
}
}
#[test]
fn health_status_fails_closed_on_an_unexpected_body() {
for (body, expected) in [
(
json!({ "status": "warming" }),
MemoryHealthStatus::Unrecognised("warming".to_string()),
),
(json!({ "version": "1.0" }), MemoryHealthStatus::Missing),
(json!({ "status": 1 }), MemoryHealthStatus::Missing),
(json!("not a health body"), MemoryHealthStatus::Missing),
] {
let read = MemoryHealthStatus::from_health_body(&body);
assert_eq!(read, expected, "body {body}");
assert!(!read.is_ok(), "body {body} must not read as healthy");
}
}
#[tokio::test]
async fn call_memory_tool_at_reports_a_dead_socket_rather_than_hanging() {
let tmp = tempfile::tempdir().expect("tempdir");
let started = std::time::Instant::now();
let result = call_memory_tool_at(
&tmp.path().join("absent.sock"),
"memory_list",
json!({ "palace": "p" }),
)
.await;
assert!(result.is_err(), "an absent socket cannot answer");
assert!(
started.elapsed() < DEFAULT_TIMEOUT,
"a refused dial must not wait out the budget: {:?}",
started.elapsed()
);
}
}