use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use wispers_connect::ServingHandle;
#[cfg(unix)]
use tokio::net::{UnixListener, UnixStream};
#[cfg(windows)]
use tokio::net::{TcpListener, TcpStream};
#[cfg(unix)]
pub type IpcStream = UnixStream;
#[cfg(windows)]
pub type IpcStream = TcpStream;
#[cfg(unix)]
type ReadHalf = tokio::net::unix::OwnedReadHalf;
#[cfg(unix)]
type WriteHalf = tokio::net::unix::OwnedWriteHalf;
#[cfg(windows)]
type ReadHalf = tokio::net::tcp::OwnedReadHalf;
#[cfg(windows)]
type WriteHalf = tokio::net::tcp::OwnedWriteHalf;
pub fn ipc_path(connectivity_group_id: &str, node_number: i32) -> PathBuf {
let base = dirs::home_dir().unwrap_or_else(std::env::temp_dir);
let dir = base.join(".wconnect").join("sockets");
#[cfg(unix)]
return dir.join(format!("{}-{}.sock", connectivity_group_id, node_number));
#[cfg(windows)]
return dir.join(format!("{}-{}.port", connectivity_group_id, node_number));
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, clap::ValueEnum, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum TtlProfile {
#[default]
Interactive,
Asynchronous,
}
impl TtlProfile {
fn to_lib(self) -> wispers_connect::TtlProfile {
match self {
TtlProfile::Interactive => wispers_connect::TtlProfile::Interactive,
TtlProfile::Asynchronous => wispers_connect::TtlProfile::Asynchronous,
}
}
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(tag = "cmd", rename_all = "snake_case")]
pub enum Request {
Status,
GetActivationCode {
#[serde(default)]
ttl_profile: TtlProfile,
},
Shutdown,
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(untagged)]
pub enum Response {
Success { ok: bool, data: ResponseData },
Error { ok: bool, error: String },
}
#[derive(Debug, Serialize, Deserialize)]
#[serde(untagged)]
pub enum ResponseData {
Status(StatusData),
ActivationCode(ActivationCodeData),
Empty,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct StatusData {
pub connected: bool,
pub node_number: i32,
pub cg_id: String,
pub endorsing: Option<EndorsingData>,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct EndorsingData {
pub codes_outstanding: usize,
pub nodes_awaiting_cosign: Vec<i32>,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ActivationCodeData {
pub activation_code: String,
}
impl Response {
pub fn success(data: ResponseData) -> Self {
Response::Success { ok: true, data }
}
pub fn error(msg: impl Into<String>) -> Self {
Response::Error {
ok: false,
error: msg.into(),
}
}
}
pub struct Server {
#[cfg(unix)]
listener: UnixListener,
#[cfg(windows)]
listener: TcpListener,
#[cfg(windows)]
windows_ipc_password: String,
connectivity_group_id: String,
node_number: i32,
}
impl Server {
#[cfg(unix)]
pub async fn bind(connectivity_group_id: &str, node_number: i32) -> Result<Self> {
let path = ipc_path(connectivity_group_id, node_number);
if let Some(parent) = path.parent() {
tokio::fs::create_dir_all(parent)
.await
.context("failed to create socket directory")?;
}
if path.exists() {
match UnixStream::connect(&path).await {
Ok(_) => {
anyhow::bail!("server already running at {:?}", path);
}
Err(_) => {
tokio::fs::remove_file(&path)
.await
.context("failed to remove stale socket")?;
}
}
}
let listener = UnixListener::bind(&path).context("failed to bind socket")?;
Ok(Self {
listener,
connectivity_group_id: connectivity_group_id.to_string(),
node_number,
})
}
#[cfg(windows)]
pub async fn bind(connectivity_group_id: &str, node_number: i32) -> Result<Self> {
use rand::Rng;
let path = ipc_path(connectivity_group_id, node_number);
if let Some(parent) = path.parent() {
tokio::fs::create_dir_all(parent)
.await
.context("failed to create socket directory")?;
}
if path.exists() {
if let Ok(contents) = tokio::fs::read_to_string(&path).await
&& let Some((port, _)) = parse_port_file(&contents)
&& TcpStream::connect(("127.0.0.1", port)).await.is_ok()
{
anyhow::bail!("server already running on port {}", port);
}
tokio::fs::remove_file(&path)
.await
.context("failed to remove stale port file")?;
}
let listener = TcpListener::bind("127.0.0.1:0")
.await
.context("failed to bind TCP listener")?;
let port = listener.local_addr()?.port();
let password: String = rand::rng()
.sample_iter(rand::distr::Alphanumeric)
.take(32)
.map(char::from)
.collect();
tokio::fs::write(&path, format!("{}:{}", port, password))
.await
.context("failed to write port file")?;
Ok(Self {
listener,
windows_ipc_password: password,
connectivity_group_id: connectivity_group_id.to_string(),
node_number,
})
}
#[cfg(unix)]
pub async fn accept(&self) -> Result<IpcStream> {
let (stream, _addr) = self.listener.accept().await?;
Ok(stream)
}
#[cfg(windows)]
pub async fn accept(&self) -> Result<IpcStream> {
loop {
let (stream, _addr) = self.listener.accept().await?;
let mut buf_stream = BufReader::new(stream);
let mut password_line = String::new();
match buf_stream.read_line(&mut password_line).await {
Ok(0) => continue,
Ok(_) if password_line.trim() == self.windows_ipc_password => {
return Ok(buf_stream.into_inner());
}
_ => continue,
}
}
}
pub fn path(&self) -> PathBuf {
ipc_path(&self.connectivity_group_id, self.node_number)
}
}
impl Drop for Server {
fn drop(&mut self) {
let _ = std::fs::remove_file(self.path());
}
}
#[allow(dead_code)]
pub async fn handle_client(stream: IpcStream, handle: ServingHandle) {
let (reader, mut writer) = stream.into_split();
let mut reader = BufReader::new(reader);
let mut line = String::new();
loop {
line.clear();
match reader.read_line(&mut line).await {
Ok(0) => break,
Ok(_) => {
let response = process_request(&line, &handle).await;
let response_json = serde_json::to_string(&response).unwrap_or_else(|e| {
serde_json::to_string(&Response::error(format!("serialization error: {}", e)))
.unwrap()
});
if let Err(e) = writer.write_all(response_json.as_bytes()).await {
eprintln!("Failed to write response: {}", e);
break;
}
if let Err(e) = writer.write_all(b"\n").await {
eprintln!("Failed to write newline: {}", e);
break;
}
if let Err(e) = writer.flush().await {
eprintln!("Failed to flush: {}", e);
break;
}
if matches!(
serde_json::from_str::<Request>(&line),
Ok(Request::Shutdown)
) {
break;
}
}
Err(e) => {
eprintln!("Failed to read from client: {}", e);
break;
}
}
}
}
pub async fn handle_client_with_optional_handle(
stream: IpcStream,
handle_state: std::sync::Arc<tokio::sync::RwLock<Option<ServingHandle>>>,
) {
let (reader, mut writer) = stream.into_split();
let mut reader = BufReader::new(reader);
let mut line = String::new();
loop {
line.clear();
match reader.read_line(&mut line).await {
Ok(0) => break,
Ok(_) => {
let response = {
let guard = handle_state.read().await;
match &*guard {
Some(handle) => process_request(&line, handle).await,
None => {
let request: Result<Request, _> = serde_json::from_str(&line);
match request {
Ok(Request::Status) => {
Response::success(ResponseData::Status(StatusData {
connected: false,
node_number: 0, cg_id: String::new(),
endorsing: None,
}))
}
Ok(_) => Response::error("hub not connected yet"),
Err(e) => Response::error(format!("invalid request: {}", e)),
}
}
}
};
let response_json = serde_json::to_string(&response).unwrap_or_else(|e| {
serde_json::to_string(&Response::error(format!("serialization error: {}", e)))
.unwrap()
});
if let Err(e) = writer.write_all(response_json.as_bytes()).await {
eprintln!("Failed to write response: {}", e);
break;
}
if let Err(e) = writer.write_all(b"\n").await {
eprintln!("Failed to write newline: {}", e);
break;
}
if let Err(e) = writer.flush().await {
eprintln!("Failed to flush: {}", e);
break;
}
if matches!(
serde_json::from_str::<Request>(&line),
Ok(Request::Shutdown)
) {
break;
}
}
Err(e) => {
eprintln!("Failed to read from client: {}", e);
break;
}
}
}
}
async fn process_request(line: &str, handle: &ServingHandle) -> Response {
let request: Request = match serde_json::from_str(line) {
Ok(r) => r,
Err(e) => return Response::error(format!("invalid request: {}", e)),
};
match request {
Request::Status => match handle.status().await {
Ok(status) => {
let endorsing = status.endorsing.map(|e| EndorsingData {
codes_outstanding: e.codes_outstanding,
nodes_awaiting_cosign: e.nodes_awaiting_cosign,
});
Response::success(ResponseData::Status(StatusData {
connected: status.connected,
node_number: status.node_number,
cg_id: status.connectivity_group_id.to_string(),
endorsing,
}))
}
Err(e) => Response::error(format!("status failed: {}", e)),
},
Request::GetActivationCode { ttl_profile } => {
match handle
.generate_activation_code_with_ttl(ttl_profile.to_lib())
.await
{
Ok(code) => Response::success(ResponseData::ActivationCode(ActivationCodeData {
activation_code: code.format(),
})),
Err(e) => Response::error(format!("{}", e)),
}
}
Request::Shutdown => {
let _ = handle.shutdown().await;
Response::success(ResponseData::Empty)
}
}
}
#[cfg(windows)]
fn parse_port_file(contents: &str) -> Option<(u16, &str)> {
let contents = contents.trim();
let colon = contents.find(':')?;
let port: u16 = contents[..colon].parse().ok()?;
let password = &contents[colon + 1..];
Some((port, password))
}
pub struct Client {
reader: BufReader<ReadHalf>,
writer: WriteHalf,
}
impl Client {
#[cfg(unix)]
pub async fn connect(connectivity_group_id: &str, node_number: i32) -> Result<Self> {
let path = ipc_path(connectivity_group_id, node_number);
let stream = UnixStream::connect(&path).await.with_context(|| {
format!("failed to connect to server at {:?} (is it running?)", path)
})?;
let (reader, writer) = stream.into_split();
Ok(Self {
reader: BufReader::new(reader),
writer,
})
}
#[cfg(windows)]
pub async fn connect(connectivity_group_id: &str, node_number: i32) -> Result<Self> {
let path = ipc_path(connectivity_group_id, node_number);
let contents = tokio::fs::read_to_string(&path)
.await
.with_context(|| format!("server not running (no port file {:?})", path))?;
let (port, password) = parse_port_file(&contents).context("invalid server port file")?;
let stream = TcpStream::connect(("127.0.0.1", port))
.await
.with_context(|| format!("server not running (port {})", port))?;
let (reader, mut writer) = stream.into_split();
writer.write_all(password.as_bytes()).await?;
writer.write_all(b"\n").await?;
writer.flush().await?;
Ok(Self {
reader: BufReader::new(reader),
writer,
})
}
pub async fn request(&mut self, req: &Request) -> Result<Response> {
let request_json = serde_json::to_string(req)?;
self.writer.write_all(request_json.as_bytes()).await?;
self.writer.write_all(b"\n").await?;
self.writer.flush().await?;
let mut line = String::new();
self.reader.read_line(&mut line).await?;
let response: Response = serde_json::from_str(&line)?;
Ok(response)
}
}