#![cfg(unix)]
use crate::daemon::client;
use crate::daemon::frame;
use crate::daemon::protocol::{
Request, Response, WarmBackendIdentity, WarmBackendStatus, WIRE_VERSION,
};
use std::os::unix::fs::PermissionsExt;
use std::path::PathBuf;
use std::time::{Duration, Instant};
use tokio::io::{BufReader, BufWriter};
use tokio::net::UnixListener;
fn ready_warm_backend(detector_rules_digest: &str) -> WarmBackendStatus {
let WarmBackendIdentity {
engine,
binary_sha256,
config_digest,
..
} = client::current_warm_backend_identity(detector_rules_digest.to_string())
.expect("current daemon-client warm identity");
WarmBackendStatus {
ready: true,
daemon_generation: "mock-daemon-generation".into(),
identity: WarmBackendIdentity {
engine,
gpu_artifact: None,
binary_sha256,
detector_rules_digest: detector_rules_digest.to_string(),
config_digest,
},
required_backends: vec!["cpu-fallback".into()],
initialized_backends: vec!["cpu-fallback".into()],
reason: None,
repair_command: None,
}
}
async fn spawn_mock_daemon(socket: PathBuf, wire_version: u32, keyhog_version: String) {
let detector_rules_digest = keyhog_core::detector_digest().to_owned();
spawn_mock_daemon_identity(
socket,
wire_version,
keyhog_version,
keyhog_core::git_hash().to_string(),
detector_rules_digest,
)
.await;
}
async fn spawn_mock_daemon_identity(
socket: PathBuf,
wire_version: u32,
keyhog_version: String,
git_hash: String,
detector_rules_digest: String,
) {
spawn_mock_daemon_response(
socket,
Response::Hello {
wire_version,
keyhog_version,
git_hash,
detector_rules_digest: detector_rules_digest.clone(),
backend_policy: "autoroute".to_string(),
detector_count: 902,
uptime_secs: 1,
warm_backend: ready_warm_backend(&detector_rules_digest),
mass_service: false,
mass_gpu_primary_required: false,
},
)
.await;
}
async fn spawn_mock_daemon_response(socket: PathBuf, response: Response) {
if let Some(parent) = socket.parent() {
std::fs::set_permissions(parent, std::fs::Permissions::from_mode(0o700))
.expect("chmod mock daemon parent 0700");
}
let listener = UnixListener::bind(&socket).expect("bind mock daemon socket");
std::fs::set_permissions(&socket, std::fs::Permissions::from_mode(0o600))
.expect("chmod mock daemon socket 0600");
tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept mock client");
let (reader, writer) = stream.into_split();
let mut reader = BufReader::new(reader);
let mut writer = BufWriter::new(writer);
assert!(matches!(
frame::read_request(&mut reader).await,
Ok(Some(Request::Hello))
));
frame::write_response(&mut writer, &response)
.await
.expect("write mock daemon response");
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
});
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
}
async fn spawn_mock_daemon_backend_policy(socket: PathBuf, backend_policy: &str) {
let detector_rules_digest = keyhog_core::detector_digest().to_owned();
spawn_mock_daemon_response(
socket,
Response::Hello {
wire_version: WIRE_VERSION,
keyhog_version: env!("CARGO_PKG_VERSION").to_string(),
git_hash: keyhog_core::git_hash().to_string(),
detector_rules_digest: detector_rules_digest.clone(),
backend_policy: backend_policy.to_string(),
detector_count: keyhog_core::embedded_detector_count(),
uptime_secs: 1,
warm_backend: ready_warm_backend(&detector_rules_digest),
mass_service: false,
mass_gpu_primary_required: false,
},
)
.await;
}
async fn spawn_stuck_handshake_daemon(socket: PathBuf) -> tokio::task::JoinHandle<()> {
if let Some(parent) = socket.parent() {
std::fs::set_permissions(parent, std::fs::Permissions::from_mode(0o700))
.expect("chmod mock daemon parent 0700");
}
let listener = UnixListener::bind(&socket).expect("bind stuck mock daemon socket");
std::fs::set_permissions(&socket, std::fs::Permissions::from_mode(0o600))
.expect("chmod stuck mock daemon socket 0600");
let handle = tokio::spawn(async move {
let (stream, _) = listener
.accept()
.await
.expect("accept stuck-handshake client");
let (reader, _writer) = stream.into_split();
let mut reader = BufReader::new(reader);
assert!(matches!(
frame::read_request(&mut reader).await,
Ok(Some(Request::Hello))
));
std::future::pending::<()>().await;
});
tokio::time::sleep(Duration::from_millis(20)).await;
handle
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn connect_fails_closed_on_keyhog_version_mismatch() {
let dir = tempfile::tempdir().unwrap();
let socket = dir.path().join("stale.sock");
spawn_mock_daemon(socket.clone(), WIRE_VERSION, "0.0.1-stale".to_string()).await;
let res = client::connect(&socket).await;
assert!(
res.is_err(),
"connect must refuse a daemon running a different keyhog version"
);
let err = res.err().unwrap();
let msg = format!("{err:#}");
assert!(
msg.contains("identity mismatch") && msg.contains("package version"),
"error must name the version mismatch as the reason: {msg}"
);
assert!(
msg.contains("0.0.1-stale"),
"error must report the stale daemon's version so the operator can act: {msg}"
);
assert!(
msg.contains("daemon stop") && msg.contains("daemon start"),
"error must tell the operator how to clear the stale daemon: {msg}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn connect_fails_closed_on_same_version_different_detector_corpus() {
let dir = tempfile::tempdir().unwrap();
let socket = dir.path().join("different-corpus.sock");
spawn_mock_daemon_identity(
socket.clone(),
WIRE_VERSION,
env!("CARGO_PKG_VERSION").to_string(),
keyhog_core::git_hash().to_string(),
"different-detector-rules".to_string(),
)
.await;
let error = client::connect(&socket)
.await
.err()
.expect("different detector corpus must be rejected");
let message = format!("{error:#}");
assert!(message.contains("detector rules daemon=different-detector-rules"));
assert!(message.contains("daemon stop") && message.contains("daemon start"));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn connect_any_version_accepts_stale_daemon_so_stop_status_work() {
let dir = tempfile::tempdir().unwrap();
let socket = dir.path().join("stale2.sock");
spawn_mock_daemon(socket.clone(), WIRE_VERSION, "0.0.1-stale".to_string()).await;
let conn = client::connect_any_version(&socket)
.await
.expect("connect_any_version must tolerate a keyhog-version mismatch");
assert!(
conn.is_stale(),
"a daemon on a different keyhog version must be reported stale"
);
assert_eq!(
conn.daemon_version(),
"0.0.1-stale",
"the daemon's reported version must be exposed for the status warning"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn connect_succeeds_against_same_version_daemon() {
let dir = tempfile::tempdir().unwrap();
let socket = dir.path().join("fresh.sock");
spawn_mock_daemon(
socket.clone(),
WIRE_VERSION,
env!("CARGO_PKG_VERSION").to_string(),
)
.await;
let conn = client::connect(&socket)
.await
.expect("connect must succeed when the daemon runs the same keyhog version");
assert!(
!conn.is_stale(),
"a same-version daemon must NOT be reported stale"
);
assert_eq!(
conn.daemon_version(),
env!("CARGO_PKG_VERSION"),
"connect must record the daemon's reported version"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn connect_still_rejects_wire_version_mismatch() {
let dir = tempfile::tempdir().unwrap();
let socket = dir.path().join("badwire.sock");
spawn_mock_daemon(
socket.clone(),
WIRE_VERSION.wrapping_add(1),
env!("CARGO_PKG_VERSION").to_string(),
)
.await;
let res = client::connect(&socket).await;
assert!(
res.is_err(),
"connect must refuse an incompatible wire version"
);
let err = res.err().unwrap();
let msg = format!("{err:#}");
assert!(
msg.contains("wire version mismatch"),
"wire-version mismatch must be reported distinctly: {msg}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn connect_rejects_unrecognized_daemon_backend_policy() {
let dir = tempfile::tempdir().unwrap();
let socket = dir.path().join("bad-backend-policy.sock");
spawn_mock_daemon_backend_policy(socket.clone(), "prefer-gpu-then-fallback").await;
let error = client::connect(&socket)
.await
.err()
.expect("an undisclosed fallback policy must not be accepted");
let message = format!("{error:#}");
assert!(
message.contains("invalid backend policy")
&& message.contains("prefer-gpu-then-fallback")
&& message.contains("Restart it"),
"invalid daemon routing policy must fail with repair guidance: {message}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn connect_protocol_mismatch_does_not_dump_response_payload() {
let dir = tempfile::tempdir().unwrap();
let socket = dir.path().join("wrong-kind.sock");
spawn_mock_daemon_response(
socket.clone(),
Response::Error {
message: "daemon-controlled plaintext payload must stay hidden".to_string(),
},
)
.await;
let res = client::connect(&socket).await;
assert!(
res.is_err(),
"connect must reject a non-Hello handshake response"
);
let err = res.err().unwrap();
let msg = format!("{err:#}");
assert!(
msg.contains("expected Hello reply") && msg.contains("Error"),
"protocol mismatch should name only the response kind: {msg}"
);
assert!(
!msg.contains("plaintext payload") && !msg.contains("message:"),
"daemon client must not Debug-dump daemon-controlled response fields: {msg}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn connect_times_out_when_daemon_never_answers_hello() {
let dir = tempfile::tempdir().unwrap();
let socket = dir.path().join("stuck-hello.sock");
let stuck_daemon = spawn_stuck_handshake_daemon(socket.clone()).await;
let started = Instant::now();
let res = tokio::time::timeout(Duration::from_secs(3), client::connect(&socket))
.await
.expect("client::connect must return via its internal handshake timeout");
stuck_daemon.abort();
assert!(res.is_err(), "stuck daemon handshake must fail");
let msg = format!("{:#}", res.err().unwrap());
assert!(
msg.contains("handshake timeout waiting for Hello"),
"timeout error must name the stuck Hello handshake: {msg}"
);
assert!(
started.elapsed() < Duration::from_secs(3),
"internal handshake timeout should fire before the outer test guard"
);
}