use std::collections::HashSet;
use std::os::unix::net::UnixStream;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use nix::sys::socket::{getsockopt, sockopt};
use crate::config::Config;
use crate::process;
use crate::protocol::{GuestMessage, HostMessage, read_frame, write_frame};
use crate::systemd;
use super::handlers;
use super::monitor::monitor_pidfd;
use super::{MAX_NEGOTIATION_FAILURES, MAX_SESSIONS, PING_INTERVAL, SharedState};
pub(crate) fn handle_connection(
stream: &mut UnixStream,
config: &Config,
state: &Arc<SharedState>,
) -> anyhow::Result<()> {
stream.set_read_timeout(Some(PING_INTERVAL))?;
let mut last_ping = std::time::Instant::now();
let mut negotiated: Option<HashSet<String>> = None;
let mut failures: u32 = 0;
loop {
let msg_bytes = match read_frame(stream) {
Ok(Some(b)) => b,
Ok(None) => return Ok(()),
Err(e)
if e.kind() == std::io::ErrorKind::WouldBlock
|| e.kind() == std::io::ErrorKind::TimedOut =>
{
if last_ping.elapsed() >= PING_INTERVAL {
if write_frame(stream, &HostMessage::Ping).is_err() {
return Ok(());
}
last_ping = std::time::Instant::now();
}
continue;
}
Err(e) => return Err(e.into()),
};
last_ping = std::time::Instant::now();
let msg: GuestMessage = match serde_json::from_slice(&msg_bytes) {
Ok(m) => m,
Err(e) => {
tracing::warn!("malformed frame from peer: {e}");
if note_failure(stream, &mut failures, "malformed frame") {
return Ok(());
}
continue;
}
};
match msg {
GuestMessage::Hello {
protocol_version,
guest_version,
container,
capabilities,
} => {
let outcome = handlers::handle_hello(
stream,
&config.integration,
state.idle_timeout_secs,
protocol_version,
guest_version,
container,
capabilities,
)?;
let handlers::HelloOutcome::Accepted(accepted) = outcome else {
if note_failure(stream, &mut failures, "hello rejected") {
return Ok(());
}
continue;
};
negotiated = Some(accepted.into_iter().collect());
}
GuestMessage::RegisterSession => {
if !peer_is_in_host_userns(stream) {
tracing::warn!("rejecting RegisterSession from foreign user namespace");
let _ = write_frame(
stream,
&HostMessage::Error {
reason: "register_session is host-only".into(),
},
);
return Ok(());
}
if state.session_count.load(Ordering::SeqCst) >= MAX_SESSIONS {
tracing::warn!("rejecting RegisterSession: session cap reached");
let _ = write_frame(
stream,
&HostMessage::Error {
reason: "session limit reached".into(),
},
);
return Ok(());
}
let raw_fd = match process::recv_fd(stream) {
Ok(Some(fd)) => fd,
Ok(None) => return Ok(()),
Err(_) => return Ok(()),
};
let fd = process::adopt_scm_fd(raw_fd);
state.session_count.fetch_add(1, Ordering::SeqCst);
let s = Arc::clone(state);
std::thread::spawn(move || monitor_pidfd(fd, s));
return Ok(());
}
GuestMessage::Busy => {
if negotiated.is_none() {
if note_failure(stream, &mut failures, "hello required") {
return Ok(());
}
}
}
GuestMessage::IdleTimeout => {
if negotiated.is_none() {
if note_failure(stream, &mut failures, "hello required") {
return Ok(());
}
continue;
}
if state.idle_timeout_secs > 0 {
let name = &state.container_name;
tracing::info!("container '{}' idle — stopping", name);
let _ = systemd::stop_unit(name);
if state.was_socket_activated {
std::process::exit(0);
}
}
}
GuestMessage::Notify {
summary,
body,
urgency: _,
actions,
app_name: _,
} => {
if !has_cap(&negotiated, crate::protocol::CAP_NOTIFY) {
if note_failure(stream, &mut failures, "capability 'notify' not accepted") {
return Ok(());
}
continue;
}
handlers::handle_notify(stream, summary, body, actions)?
}
GuestMessage::XdgOpen { uri } => {
if !has_cap(&negotiated, crate::protocol::CAP_XDG_OPEN) {
if note_failure(stream, &mut failures, "capability 'xdg_open' not accepted") {
return Ok(());
}
continue;
}
handlers::handle_xdg_open(uri)?
}
GuestMessage::ClipboardSet { text } => {
if !has_cap(&negotiated, crate::protocol::CAP_CLIPBOARD) {
if note_failure(stream, &mut failures, "capability 'clipboard' not accepted") {
return Ok(());
}
continue;
}
handlers::handle_clipboard_set(text)?
}
GuestMessage::ClipboardGet => {
if !has_cap(&negotiated, crate::protocol::CAP_CLIPBOARD) {
if note_failure(stream, &mut failures, "capability 'clipboard' not accepted") {
return Ok(());
}
continue;
}
handlers::handle_clipboard_get(stream)?
}
GuestMessage::HostExec { cmd, args } => {
if !has_cap(&negotiated, crate::protocol::CAP_HOST_EXEC) {
if note_failure(stream, &mut failures, "capability 'host_exec' not accepted") {
return Ok(());
}
continue;
}
handlers::handle_host_exec(stream, &config.integration, cmd, args)?
}
GuestMessage::Info {
guest_version,
protocol_version,
} => {
tracing::info!(
"guest info: version {}, protocol {}",
guest_version,
protocol_version
);
}
}
}
}
pub(crate) fn note_failure(stream: &mut UnixStream, failures: &mut u32, reason: &str) -> bool {
*failures = failures.saturating_add(1);
let _ = write_frame(
stream,
&HostMessage::Error {
reason: reason.to_string(),
},
);
*failures >= MAX_NEGOTIATION_FAILURES
}
pub(crate) fn has_cap(negotiated: &Option<HashSet<String>>, cap: &str) -> bool {
negotiated
.as_ref()
.is_some_and(|caps| caps.iter().any(|c| c == cap))
}
fn peer_is_in_host_userns(stream: &UnixStream) -> bool {
let creds = match getsockopt(stream, sockopt::PeerCredentials) {
Ok(c) => c,
Err(_) => return false,
};
let self_ns = std::fs::read_link("/proc/self/ns/user").ok();
let peer_ns = std::fs::read_link(format!("/proc/{}/ns/user", creds.pid())).ok();
match (self_ns, peer_ns) {
(Some(a), Some(b)) => a == b,
_ => false,
}
}