use std::path::{Path, PathBuf};
use std::time::Duration;
use serde::{Deserialize, Serialize};
use super::{EmbedderClient, EmbedderError};
const EMBED_TIMEOUT: Duration = Duration::from_secs(600);
const EMBED_MAX_FRAME_BYTES: u64 = 256 * 1024 * 1024;
const METHOD_EMBED: &str = "embed";
const JSONRPC_VERSION: &str = "2.0";
#[derive(Debug, Serialize)]
struct RpcRequest<'a> {
jsonrpc: &'a str,
method: &'a str,
params: EmbedParams<'a>,
id: u64,
}
#[derive(Debug, Serialize)]
struct EmbedParams<'a> {
texts: &'a [String],
}
#[derive(Debug, Deserialize)]
struct RpcResponse {
#[serde(default)]
result: Option<EmbedResult>,
#[serde(default)]
error: Option<RpcError>,
}
#[derive(Debug, Deserialize)]
struct EmbedResult {
embeddings: Vec<Vec<f32>>,
}
#[derive(Debug, Deserialize)]
struct RpcError {
code: i32,
message: String,
}
#[derive(Debug, Clone)]
pub struct UdsEmbedderClient {
socket_path: PathBuf,
}
impl UdsEmbedderClient {
pub fn new(socket_path: impl Into<PathBuf>) -> Self {
Self {
socket_path: socket_path.into(),
}
}
pub fn default_path() -> PathBuf {
crate::uds::scratch_socket_dir().join(SOCKET_FILENAME)
}
pub fn socket_path(&self) -> &Path {
&self.socket_path
}
}
pub const SOCKET_FILENAME: &str = "trusty-embedderd.sock";
#[async_trait::async_trait]
impl EmbedderClient for UdsEmbedderClient {
async fn embed_batch(&self, texts: Vec<String>) -> Result<Vec<Vec<f32>>, EmbedderError> {
if texts.is_empty() {
return Ok(vec![]);
}
let sent = texts.len();
tracing::debug!(
socket = %self.socket_path.display(),
n = sent,
"UdsEmbedderClient: sending batch"
);
let req = RpcRequest {
jsonrpc: JSONRPC_VERSION,
method: METHOD_EMBED,
params: EmbedParams { texts: &texts },
id: 1,
};
let resp: RpcResponse = crate::uds::rpc::send_framed_request_capped(
&self.socket_path,
&req,
EMBED_TIMEOUT,
EMBED_MAX_FRAME_BYTES,
)
.await
.map_err(|e| EmbedderError::Uds(e.to_string()))?;
if let Some(err) = resp.error {
return Err(EmbedderError::ModelError(format!(
"daemon RPC error {}: {}",
err.code, err.message
)));
}
let result = resp.result.ok_or_else(|| {
EmbedderError::Uds("response missing both result and error fields".to_owned())
})?;
if result.embeddings.len() != sent {
return Err(EmbedderError::DimensionMismatch {
sent,
got: result.embeddings.len(),
});
}
tracing::debug!(
socket = %self.socket_path.display(),
n = sent,
"UdsEmbedderClient: batch complete"
);
Ok(result.embeddings)
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
#[tokio::test]
async fn embed_batch_sends_one_newline_framed_jsonrpc_frame() {
let tmp = tempfile::tempdir().expect("tempdir");
let sock = tmp.path().join("sockets").join("embedderd-stub.sock");
let listener = crate::uds::bind_hardened(&sock).expect("bind stub socket");
let served = tokio::spawn(async move {
let (mut conn, _) = listener.accept().await.expect("accept");
let mut raw = Vec::new();
conn.read_to_end(&mut raw).await.expect("drain request");
conn.write_all(
b"{\"jsonrpc\":\"2.0\",\"result\":{\"embeddings\":[[0.5,0.25]]},\"id\":1}\n",
)
.await
.expect("write reply");
conn.flush().await.expect("flush reply");
raw
});
let vectors = UdsEmbedderClient::new(sock)
.embed_batch(vec!["hello".to_string()])
.await
.expect("embed round trip");
assert_eq!(vectors, vec![vec![0.5_f32, 0.25_f32]]);
let raw = String::from_utf8(served.await.expect("join")).expect("utf8");
assert!(
raw.ends_with('\n'),
"the request frame must be newline-terminated: {raw:?}"
);
assert_eq!(
raw.matches('\n').count(),
1,
"exactly one frame, one terminator: {raw:?}"
);
let sent: serde_json::Value =
serde_json::from_str(raw.trim_end_matches('\n')).expect("the frame is one JSON value");
assert_eq!(sent["jsonrpc"], "2.0");
assert_eq!(sent["method"], "embed");
assert_eq!(sent["params"]["texts"][0], "hello");
assert_eq!(sent["id"], 1);
}
#[test]
fn embed_bounds_are_generous_but_finite() {
assert!(EMBED_TIMEOUT >= Duration::from_secs(120));
assert!(EMBED_TIMEOUT <= Duration::from_secs(3600));
const {
assert!(
EMBED_MAX_FRAME_BYTES > crate::uds::MAX_FRAME_BYTES,
"an embed reply is bulk data; the control-plane default is too small"
);
}
const { assert!(EMBED_MAX_FRAME_BYTES >= 92_160_000) };
}
#[tokio::test]
async fn empty_batch_short_circuits() {
let client = UdsEmbedderClient::new("/nonexistent/socket/path");
let result = client
.embed_batch(vec![])
.await
.expect("empty batch must short-circuit");
assert!(result.is_empty());
}
#[test]
fn request_serialises_correctly() {
let texts = vec!["hello".to_string(), "world".to_string()];
let req = RpcRequest {
jsonrpc: JSONRPC_VERSION,
method: METHOD_EMBED,
params: EmbedParams { texts: &texts },
id: 1,
};
let s = serde_json::to_string(&req).unwrap();
assert!(s.contains("\"jsonrpc\":\"2.0\""), "must have jsonrpc 2.0");
assert!(s.contains("\"method\":\"embed\""), "must have embed method");
assert!(
s.contains("\"texts\":[\"hello\",\"world\"]"),
"must include texts"
);
assert!(s.contains("\"id\":1"), "must have id");
}
#[test]
fn error_response_maps_to_model_error() {
let json = r#"{"jsonrpc":"2.0","error":{"code":-32603,"message":"ort failed"},"id":1}"#;
let resp: RpcResponse = serde_json::from_str(json).unwrap();
assert!(resp.error.is_some());
assert!(resp.result.is_none());
let err = resp.error.unwrap();
assert_eq!(err.code, -32603);
assert!(err.message.contains("ort failed"));
}
#[test]
fn default_socket_path_uses_tmpdir() {
let p = UdsEmbedderClient::default_path();
assert_eq!(
p.file_name().and_then(|s| s.to_str()),
Some(SOCKET_FILENAME),
"default path must end with {SOCKET_FILENAME}"
);
assert_eq!(
p.parent(),
Some(crate::uds::scratch_socket_dir().as_path()),
"default path must live in the per-uid scratch socket directory"
);
}
#[test]
fn dimension_mismatch_detected() {
let resp = RpcResponse {
result: Some(EmbedResult {
embeddings: vec![vec![0.1_f32]],
}),
error: None,
};
let sent = 2;
let got = resp.result.unwrap().embeddings.len();
assert_ne!(sent, got);
let err = EmbedderError::DimensionMismatch { sent, got };
let s = err.to_string();
assert!(s.contains("2") && s.contains("1"));
}
}