use super::*;
pub const UNIX_SOCK_PATH_MAX: usize = 104;
const READ_TIMEOUT: Duration = Duration::from_secs(3);
pub const MAX_CONCURRENT_CONNECTIONS: usize = 64;
pub const AUTO_DRAIN_TIMEOUT: Duration = Duration::from_secs(10);
#[derive(Default)]
pub struct Shutdown {
flag: std::sync::atomic::AtomicBool,
notify: tokio::sync::Notify,
}
impl Shutdown {
pub fn new() -> Self {
Self::default()
}
pub fn signal(&self) {
self.flag.store(true, std::sync::atomic::Ordering::SeqCst);
self.notify.notify_waiters();
}
pub fn is_set(&self) -> bool {
self.flag.load(std::sync::atomic::Ordering::SeqCst)
}
pub async fn wait(&self) {
let notified = self.notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if self.is_set() {
return;
}
notified.await;
}
}
pub(super) const PROTOCOL_VERSION: u32 = 1;
#[derive(Debug, Deserialize)]
pub(crate) struct SocketRequest {
pub cmd: String,
#[allow(dead_code)] #[serde(default, rename = "v")]
pub version: Option<u32>,
#[serde(default)]
pub args: serde_json::Value,
}
#[derive(Debug, Serialize)]
pub(crate) struct SocketResponse {
pub(crate) ok: bool,
#[serde(rename = "v")]
version: u32,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) data: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub(crate) error: Option<String>,
}
impl SocketResponse {
pub(crate) fn ok(data: serde_json::Value) -> Self {
Self {
ok: true,
version: PROTOCOL_VERSION,
data: Some(data),
error: None,
}
}
pub(crate) fn err(msg: impl Into<String>) -> Self {
Self {
ok: false,
version: PROTOCOL_VERSION,
data: None,
error: Some(msg.into()),
}
}
}
pub async fn socket_handle_connection(
graph: Arc<tokio::sync::RwLock<Graph>>,
policy_matcher: Arc<tokio::sync::RwLock<crate::hooks::policy_match::PolicyMatcherSet>>,
repo_root: &Path,
stream: UnixStream,
peer: super::metadata::PeerContext,
daemon_session: uuid::Uuid,
) -> Result<()> {
use super::protocol::MAX_FRAME_SIZE;
use tokio::io::AsyncReadExt;
let (reader, mut writer) = stream.into_split();
let mut buf = String::new();
let limited = reader.take(MAX_FRAME_SIZE as u64 + 1);
let mut buf_reader = BufReader::new(limited);
match tokio::time::timeout(READ_TIMEOUT, buf_reader.read_line(&mut buf)).await {
Ok(Ok(0)) => return Ok(()),
Ok(Ok(_)) => {}
Ok(Err(e)) => anyhow::bail!("read error: {e}"),
Err(_) => anyhow::bail!("read timeout"),
}
if buf.len() > MAX_FRAME_SIZE {
let resp = super::protocol::Response::err(
uuid::Uuid::nil(),
super::protocol::ErrorCode::FrameTooLarge,
format!("request exceeds {MAX_FRAME_SIZE} byte limit"),
);
let json = serde_json::to_string(&resp)?;
writer.write_all(json.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
return Ok(());
}
let trimmed = buf.trim();
let v2_req = match serde_json::from_str::<super::protocol::Request>(trimmed) {
Ok(r) => r,
Err(e) => {
let resp = super::protocol::Response::err(
uuid::Uuid::nil(),
super::protocol::ErrorCode::MalformedRequest,
format!("invalid v2 request: {e}"),
);
let json = serde_json::to_string(&resp)?;
writer.write_all(json.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
return Ok(());
}
};
let ctx = super::dispatch_v2::RequestContext {
peer,
daemon_session,
repo_root: repo_root.to_path_buf(),
policy_matcher,
};
let resp = super::dispatch_v2::dispatch_v2(&graph, &ctx, v2_req).await;
let json = serde_json::to_string(&resp)?;
writer.write_all(json.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(())
}
pub(super) fn build_v1_dispatch_ctx(repo_root: &Path) -> super::dispatch_v2::RequestContext {
super::dispatch_v2::RequestContext {
peer: super::metadata::PeerContext {
uid: super::metadata::current_euid(),
pid: Some(std::process::id()),
},
daemon_session: uuid::Uuid::nil(),
repo_root: repo_root.to_path_buf(),
policy_matcher: Arc::new(tokio::sync::RwLock::new(
crate::hooks::policy_match::PolicyMatcherSet::empty(),
)),
}
}