use std::path::{Path, PathBuf};
use std::time::Duration;
use serde::de::Error as _;
use serde_json::json;
pub use crate::residency::{ACTIVE_PROJECT_SET_SCHEMA, ActiveProject, ActiveProjectSet};
use crate::uds::UdsRpcError;
use crate::uds::send_framed_request;
use crate::uds::server::{JSONRPC_VERSION, RpcResponse};
pub const TRUSTY_MPM_SOCKET_ENV: &str = "TRUSTY_MPM_SOCKET";
const MPM_APP_NAME: &str = "trusty-mpm";
const UNREACHABLE_PLACEHOLDER: &str = "/nonexistent/trusty-mpm/trusty-mpm.sock";
pub const METHOD_RESIDENCY_ACTIVE: &str = "mpm.residency.active";
pub fn mpm_socket() -> PathBuf {
if let Ok(raw) = std::env::var(TRUSTY_MPM_SOCKET_ENV) {
let trimmed = raw.trim();
if !trimmed.is_empty() {
return PathBuf::from(trimmed);
}
}
crate::daemon_addr::daemon_socket_path(MPM_APP_NAME).unwrap_or_else(|e| {
eprintln!("trusty-mpm: {e}");
PathBuf::from(UNREACHABLE_PLACEHOLDER)
})
}
pub async fn fetch_active_projects(timeout: Duration) -> Result<ActiveProjectSet, UdsRpcError> {
fetch_active_projects_at(&mpm_socket(), timeout).await
}
pub async fn fetch_active_projects_at(
socket: &Path,
timeout: Duration,
) -> Result<ActiveProjectSet, UdsRpcError> {
let request = json!({
"jsonrpc": JSONRPC_VERSION,
"id": 1,
"method": METHOD_RESIDENCY_ACTIVE,
"params": {},
});
let response: RpcResponse = send_framed_request(socket, &request, timeout).await?;
decode_active_project_set(socket, response)
}
fn decode_active_project_set(
socket: &Path,
response: RpcResponse,
) -> Result<ActiveProjectSet, UdsRpcError> {
match (response.result, response.error) {
(Some(result), _) => serde_json::from_value(result).map_err(|source| UdsRpcError::Decode {
path: socket.to_path_buf(),
source,
}),
(None, Some(e)) => Err(UdsRpcError::Decode {
path: socket.to_path_buf(),
source: serde_json::Error::custom(format!(
"{METHOD_RESIDENCY_ACTIVE} failed: {} ({})",
e.message, e.code
)),
}),
(None, None) => Err(UdsRpcError::Decode {
path: socket.to_path_buf(),
source: serde_json::Error::custom(format!(
"{METHOD_RESIDENCY_ACTIVE} answered with neither a result nor an error"
)),
}),
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use tempfile::TempDir;
use tokio::sync::oneshot;
use super::*;
use crate::data_dir::ENV_LOCK;
use crate::uds::server::{RpcRouter, RpcServeOptions, serve_until};
struct FakeProducer {
socket: PathBuf,
_dir: TempDir,
stop: Option<oneshot::Sender<()>>,
}
impl Drop for FakeProducer {
fn drop(&mut self) {
if let Some(tx) = self.stop.take() {
let _ = tx.send(());
}
}
}
async fn fake_producer(result: serde_json::Value) -> FakeProducer {
let dir = TempDir::new().expect("tempdir");
let socket = dir.path().join("trusty-mpm.sock");
let listener = crate::uds::bind_hardened(&socket).expect("bind");
let router = RpcRouter::new().typed::<serde_json::Value, serde_json::Value, _, _>(
METHOD_RESIDENCY_ACTIVE,
move |_params: serde_json::Value| {
let result = result.clone();
async move { Ok(result) }
},
);
let (stop, shutdown) = oneshot::channel::<()>();
tokio::spawn(async move {
serve_until(
&listener,
Arc::new(router),
RpcServeOptions::default(),
async {
let _ = shutdown.await;
},
)
.await;
});
FakeProducer {
socket,
_dir: dir,
stop: Some(stop),
}
}
fn sample_set() -> ActiveProjectSet {
ActiveProjectSet {
schema: ACTIVE_PROJECT_SET_SCHEMA,
generation: 7,
published_at_unix: 1_700_000_000,
projects: vec![ActiveProject {
root: PathBuf::from("/repo"),
palace_id: Some("palace-1".to_string()),
index_ids: vec!["idx-1".to_string()],
session_ids: vec!["sess-1".to_string()],
last_activity_unix: Some(1_700_000_500),
}],
}
}
#[test]
fn mpm_socket_honours_the_env_override() {
let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
unsafe {
std::env::set_var(TRUSTY_MPM_SOCKET_ENV, "/tmp/example-mpm.sock");
}
let resolved = mpm_socket();
unsafe {
std::env::remove_var(TRUSTY_MPM_SOCKET_ENV);
}
assert_eq!(resolved, PathBuf::from("/tmp/example-mpm.sock"));
}
#[tokio::test]
async fn fetch_active_projects_at_round_trips_a_set() {
let set = sample_set();
let daemon = fake_producer(serde_json::to_value(&set).expect("serialize")).await;
let fetched = fetch_active_projects_at(&daemon.socket, Duration::from_secs(5))
.await
.expect("the daemon answers");
assert_eq!(fetched, set);
}
#[tokio::test]
async fn fetch_active_projects_at_reports_a_dead_socket_rather_than_hanging() {
let dir = TempDir::new().expect("tempdir");
let socket = dir.path().join("absent.sock");
let started = std::time::Instant::now();
let err = fetch_active_projects_at(&socket, Duration::from_secs(5))
.await
.expect_err("nothing is listening");
assert!(
matches!(err, UdsRpcError::Dial { .. }),
"a refused dial is a Dial error, not something else: {err:?}"
);
assert!(
started.elapsed() < Duration::from_secs(5),
"a refused dial must not wait out the budget: {:?}",
started.elapsed()
);
}
#[tokio::test]
async fn fetch_active_projects_at_reports_a_malformed_payload() {
let daemon = fake_producer(serde_json::json!("not an active project set")).await;
let err = fetch_active_projects_at(&daemon.socket, Duration::from_secs(5))
.await
.expect_err("a bare string does not decode into ActiveProjectSet");
assert!(
matches!(err, UdsRpcError::Decode { .. }),
"a malformed payload is a Decode error, not something else: {err:?}"
);
}
}