use std::path::{Path, PathBuf};
use axum::extract::{Path as AxumPath, Request, State};
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use serde_json::{Value, json};
use tracing::debug;
use trusty_common::uds::UdsRpcError;
use super::map::{Call, map_request};
use super::{
CALL_TIMEOUT, MEMORY_SERVICE, MemoryRpcError, STREAM_FIRST_FRAME_PEEK, STREAM_OPEN_TIMEOUT,
call, json_response, open_stream,
};
use crate::server::AppState;
use trusty_common::uds::sse::sse_response;
pub async fn memory_api_handler(
State(state): State<AppState>,
AxumPath(path): AxumPath<String>,
req: Request,
) -> Response {
let (parts, _body) = req.into_parts();
let query = parts.uri.query().map(str::to_owned);
let mapped = match map_request(&parts.method, &path, query.as_deref()) {
Ok(call) => call,
Err(reason) => {
tracing::warn!(
"memory_uds: {} /{path} is not mapped: {reason}",
parts.method
);
return (
StatusCode::NOT_IMPLEMENTED,
axum::Json(json!({ "error": reason, "service": MEMORY_SERVICE })),
)
.into_response();
}
};
let socket = match state.memory_socket_path() {
Ok(p) => p,
Err(reason) => return MemoryRpcError::Unresolved(reason).into_response(),
};
match mapped {
Call::Unary { method, params } => {
debug!("memory_uds: {} /{path} → {method}", parts.method);
match call(&socket, method, params, CALL_TIMEOUT).await {
Ok(result) => json_response(&result),
Err(e) => e.into_response(),
}
}
Call::Stream { method, params } => {
debug!("memory_uds: {} /{path} → {method} (stream)", parts.method);
stream_response(
&socket,
method,
params,
STREAM_OPEN_TIMEOUT,
STREAM_FIRST_FRAME_PEEK,
)
.await
}
}
}
pub async fn deprecated_memory_api_handler(
state: State<AppState>,
path: AxumPath<String>,
req: Request,
) -> Response {
tracing::trace!("memory_uds: DEPRECATED /proxy/memory/… — use /api/memory/… instead (#1849)");
memory_api_handler(state, path, req).await
}
async fn stream_response(
socket: &Path,
method: &'static str,
params: Value,
open_timeout: std::time::Duration,
peek: std::time::Duration,
) -> Response {
let mut stream = match open_stream(socket, method, params, open_timeout).await {
Ok(s) => s,
Err(e) => return e.into_response(),
};
let first = match tokio::time::timeout(peek, stream.next_frame()).await {
Ok(Some(Ok(item))) => Some(item),
Ok(Some(Err(e))) => return stream_error(method, e).into_response(),
Ok(None) => None,
Err(_) => None,
};
sse_response(first, stream, method)
}
fn stream_error(method: &str, e: UdsRpcError) -> MemoryRpcError {
match e {
UdsRpcError::Stream { error, .. } => MemoryRpcError::Refused {
code: error.code,
message: error.message,
},
other => MemoryRpcError::Unreachable(format!(
"{MEMORY_SERVICE} did not stream {method}: {other}"
)),
}
}
impl AppState {
pub(crate) fn memory_socket_path(&self) -> Result<PathBuf, String> {
match &self.memory_socket {
Some(p) => Ok(p.as_ref().clone()),
None => super::socket_path(),
}
}
#[must_use]
pub fn with_memory_socket(mut self, socket: PathBuf) -> Self {
self.memory_socket = Some(std::sync::Arc::new(socket));
self
}
}
#[cfg(test)]
mod tests {
use super::*;
use trusty_common::uds::server::RpcError;
#[test]
fn stream_error_carries_the_daemon_code() {
let err = stream_error(
"memory.activity_stream",
UdsRpcError::Stream {
path: PathBuf::from("/tmp/memory.sock"),
error: RpcError::new(-32004, "no such palace".to_string()),
},
);
assert_eq!(err.status(), StatusCode::NOT_FOUND);
assert!(err.message().contains("no such palace"), "{err:?}");
}
#[test]
fn a_transport_failure_mid_stream_is_a_bad_gateway() {
let err = stream_error(
"memory.activity_stream",
UdsRpcError::NoResponse {
path: PathBuf::from("/tmp/memory.sock"),
},
);
assert_eq!(err.status(), StatusCode::BAD_GATEWAY);
}
}