use std::path::{Path, PathBuf};
use axum::body::Bytes;
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, warn};
use trusty_common::uds::UdsRpcError;
use super::map::{Call, map_request};
use super::{
MAX_FRAME_BYTES, STREAM_OPEN_TIMEOUT, SearchRpcError, call, json_response, open_stream,
unary_timeout,
};
use crate::server::AppState;
use trusty_common::uds::sse::sse_response;
pub async fn search_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 body_limit = usize::try_from(MAX_FRAME_BYTES).unwrap_or(usize::MAX);
let body_bytes: Bytes = match axum::body::to_bytes(body, body_limit).await {
Ok(b) => b,
Err(e) => {
warn!("search_uds: could not read the request body: {e}");
return (
StatusCode::PAYLOAD_TOO_LARGE,
format!("request body exceeds the {MAX_FRAME_BYTES}-byte frame budget"),
)
.into_response();
}
};
let mapped = match map_request(&parts.method, &path, query.as_deref(), &body_bytes) {
Ok(call) => call,
Err(reason) => {
warn!(
"search_uds: {} /{path} is not mapped: {reason}",
parts.method
);
return (
StatusCode::NOT_IMPLEMENTED,
axum::Json(json!({ "error": reason, "service": super::SEARCH_SERVICE })),
)
.into_response();
}
};
let socket = match state.search_socket_path() {
Ok(p) => p,
Err(reason) => return SearchRpcError::Unresolved(reason).into_response(),
};
match mapped {
Call::Unary { method, params } => {
debug!("search_uds: {} /{path} → {method}", parts.method);
match call(&socket, method, params, unary_timeout(method)).await {
Ok(result) => json_response(&result),
Err(e) => e.into_response(),
}
}
Call::Stream { method, params } => {
debug!("search_uds: {} /{path} → {method} (stream)", parts.method);
stream_response(&socket, method, params, STREAM_OPEN_TIMEOUT).await
}
}
}
pub async fn deprecated_search_api_handler(
state: State<AppState>,
path: AxumPath<String>,
req: Request,
) -> Response {
tracing::trace!("search_uds: DEPRECATED /proxy/search/… — use /api/search/… instead (#1849)");
search_api_handler(state, path, req).await
}
async fn stream_response(
socket: &Path,
method: &'static str,
params: Value,
open_timeout: std::time::Duration,
) -> Response {
let deadline = tokio::time::Instant::now() + open_timeout;
let mut stream = match open_stream(socket, method, params, open_timeout).await {
Ok(s) => s,
Err(e) => return e.into_response(),
};
let peeked = match tokio::time::timeout_at(deadline, stream.next_frame()).await {
Ok(frame) => frame,
Err(_) => {
return SearchRpcError::Unreachable(format!(
"{} did not answer {method} within {}s",
super::SEARCH_SERVICE,
open_timeout.as_secs_f32()
))
.into_response();
}
};
let first = match peeked {
Some(Ok(item)) => Some(item),
Some(Err(e)) => return stream_error(method, e).into_response(),
None => None,
};
sse_response(first, stream, method)
}
fn stream_error(method: &str, e: UdsRpcError) -> SearchRpcError {
match e {
UdsRpcError::Stream { error, .. } => SearchRpcError::Refused {
code: error.code,
message: error.message,
},
other => SearchRpcError::Unreachable(format!(
"{} did not stream {method}: {other}",
super::SEARCH_SERVICE
)),
}
}
impl AppState {
pub(crate) fn search_socket_path(&self) -> Result<PathBuf, String> {
match &self.search_socket {
Some(p) => Ok(p.as_ref().clone()),
None => super::socket_path(),
}
}
#[must_use]
pub fn with_search_socket(mut self, socket: PathBuf) -> Self {
self.search_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(
"search.index.reindex.stream",
UdsRpcError::Stream {
path: PathBuf::from("/tmp/x.sock"),
error: RpcError::new(-32004, "no reindex in progress for 'ghost'"),
},
);
assert_eq!(err.status(), StatusCode::NOT_FOUND);
assert!(err.message().contains("ghost"), "{err:?}");
}
#[tokio::test(flavor = "multi_thread")]
async fn a_socket_that_never_answers_is_a_prompt_bad_gateway() {
let tmp = tempfile::TempDir::new().expect("tempdir");
let socket = tmp.path().join("silent.sock");
let listener = trusty_common::uds::bind_hardened(&socket).expect("bind");
let _accepting = tokio::spawn(async move {
let Ok((mut conn, _)) = listener.accept().await else {
return;
};
let mut sink = Vec::new();
let _ = tokio::io::AsyncReadExt::read_to_end(&mut conn, &mut sink).await;
std::future::pending::<()>().await;
});
let started = std::time::Instant::now();
let response = stream_response(
&socket,
super::super::METHOD_STATUS_STREAM,
json!({}),
std::time::Duration::from_millis(300),
)
.await;
assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
assert!(
started.elapsed() < std::time::Duration::from_secs(5),
"the open must be bounded, not left to the per-frame budget: {:?}",
started.elapsed()
);
}
fn stalls_then_drains(socket: PathBuf, drain_after: std::time::Duration) -> PathBuf {
let listener = trusty_common::uds::bind_hardened(&socket)
.expect("bind")
.into_std()
.expect("listener into std");
let accepted = tokio::task::spawn_blocking(move || accept_within(&listener));
tokio::spawn(async move {
let Ok(Ok(conn)) = accepted.await else {
return;
};
tokio::time::sleep(drain_after).await;
let drained = tokio::task::spawn_blocking(move || {
let mut conn = conn;
let mut sink = Vec::new();
let _ = std::io::Read::read_to_end(&mut conn, &mut sink);
conn
})
.await;
let _held = drained;
std::future::pending::<()>().await;
});
socket
}
const STUB_IO_BOUND: std::time::Duration = std::time::Duration::from_secs(10);
fn accept_within(
listener: &std::os::unix::net::UnixListener,
) -> std::io::Result<std::os::unix::net::UnixStream> {
let give_up = std::time::Instant::now() + STUB_IO_BOUND;
loop {
match listener.accept() {
Ok((conn, _)) => {
conn.set_nonblocking(false)?;
conn.set_read_timeout(Some(STUB_IO_BOUND))?;
return Ok(conn);
}
Err(e)
if e.kind() == std::io::ErrorKind::WouldBlock
&& std::time::Instant::now() < give_up =>
{
std::thread::sleep(std::time::Duration::from_millis(1));
}
Err(e) => return Err(e),
}
}
}
fn bulky_params() -> Value {
json!({ "blob": "x".repeat(4 * trusty_common::uds::SOCKET_BUFFER_BYTES) })
}
#[tokio::test(start_paused = true)]
async fn a_slow_open_and_a_silent_first_frame_share_one_budget() {
const BUDGET: std::time::Duration = std::time::Duration::from_millis(1000);
const DRAIN_AFTER: std::time::Duration = std::time::Duration::from_millis(600);
const CEILING: std::time::Duration = std::time::Duration::from_millis(1300);
let tmp = tempfile::TempDir::new().expect("tempdir");
let slow = stalls_then_drains(tmp.path().join("slow-open.sock"), DRAIN_AFTER);
let started = tokio::time::Instant::now();
let opened = open_stream(
&slow,
super::super::METHOD_STATUS_STREAM,
bulky_params(),
BUDGET,
)
.await;
let open_took = started.elapsed();
assert!(opened.is_ok(), "the open must succeed, slowly: {opened:?}");
assert!(
open_took >= DRAIN_AFTER,
"the write must park until the peer drains, or this test proves nothing: {open_took:?}"
);
drop(opened);
let silent = stalls_then_drains(tmp.path().join("silent-frame.sock"), DRAIN_AFTER);
let started = tokio::time::Instant::now();
let response = stream_response(
&silent,
super::super::METHOD_STATUS_STREAM,
bulky_params(),
BUDGET,
)
.await;
let elapsed = started.elapsed();
assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
assert!(
elapsed >= BUDGET,
"a shared deadline still spends the whole budget: {elapsed:?}"
);
assert!(
elapsed < CEILING,
"the dial, the write and the first frame read must share ONE {BUDGET:?} deadline, \
not take one each: {elapsed:?}"
);
}
#[test]
fn stream_error_reports_a_transport_failure_as_unreachable() {
let err = stream_error(
"search.status.stream",
UdsRpcError::NoResponse {
path: PathBuf::from("/tmp/x.sock"),
},
);
assert_eq!(err.status(), StatusCode::BAD_GATEWAY);
}
}