use akar_common::types::Value;
use serde::{Deserialize, Serialize};
use std::fmt;
use std::io::{ErrorKind, Read, Write};
use std::net::TcpStream;
use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
pub const DEFAULT_PORT: u16 = 9876;
pub const MAX_FRAME_SIZE: usize = 128 * 1024 * 1024;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WireRequest {
pub query: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub client_name: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct WireResponse {
pub success: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub message: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error_message: Option<String>,
pub column_names: Vec<String>,
pub rows: Vec<Vec<Option<Value>>>,
}
impl WireResponse {
pub fn success_message(msg: String) -> Self {
Self {
success: true,
message: Some(msg),
error_message: None,
column_names: Vec::new(),
rows: Vec::new(),
}
}
pub fn error(msg: String) -> Self {
Self {
success: false,
message: None,
error_message: Some(msg),
column_names: Vec::new(),
rows: Vec::new(),
}
}
pub fn num_rows(&self) -> usize {
self.rows.len()
}
pub fn num_columns(&self) -> usize {
self.column_names.len()
}
pub fn cell(&self, row: usize, col: usize) -> Option<&Value> {
self.rows.get(row)?.get(col)?.as_ref()
}
pub fn column_values(&self, col: usize) -> Vec<Value> {
self.rows.iter().filter_map(|r| r.get(col).cloned().flatten()).collect()
}
pub fn result_summary(&self) -> String {
if let Some(head) = crate::query_result::result_summary_head(
self.message.as_deref(),
self.success,
self.error_message.as_deref(),
!self.rows.is_empty(),
) {
return head;
}
format!(
"Returned {} rows in {} columns",
self.rows.len(),
self.column_names.len()
)
}
}
impl fmt::Display for WireResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if let Some(head) = crate::query_result::result_summary_head(
self.message.as_deref(),
self.success,
self.error_message.as_deref(),
!self.rows.is_empty(),
) {
return write!(f, "{head}");
}
for (i, row) in self.rows.iter().enumerate() {
if i > 0 {
writeln!(f)?;
}
write!(f, "Row {}: ", i)?;
for (col, cell) in row.iter().enumerate() {
if col > 0 {
write!(f, ", ")?;
}
match cell {
Some(v) => write!(f, "{v:?}")?,
None => write!(f, "null")?,
}
}
}
Ok(())
}
}
#[derive(Debug)]
pub enum PartialFrame {
Header([u8; 4], usize),
Payload { len: usize, buf: Vec<u8>, filled: usize },
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum DrainOutcome {
FrameConsumed,
ConnectionClosed,
NoFrameWithinGrace,
}
const GRACE_TIMEOUT: Duration = Duration::from_millis(250);
const MAX_GRACE_READS: usize = 8;
fn drain_stale_frames<R: Read>(reader: &mut R, partial: &mut Option<PartialFrame>) -> DrainOutcome {
for _ in 0..MAX_GRACE_READS {
match read_frame(reader, partial) {
Ok(Some(_frame)) => return DrainOutcome::FrameConsumed,
Ok(None) => return DrainOutcome::ConnectionClosed,
Err(e) if e.kind() == ErrorKind::TimedOut || e.kind() == ErrorKind::WouldBlock => continue,
Err(_) => return DrainOutcome::NoFrameWithinGrace,
}
}
DrainOutcome::NoFrameWithinGrace
}
pub fn write_frame<W: Write>(writer: &mut W, payload: &[u8]) -> std::io::Result<()> {
let len = payload.len();
if len > MAX_FRAME_SIZE {
return Err(std::io::Error::new(
ErrorKind::InvalidData,
format!("Frame too large: {len} bytes"),
));
}
writer.write_all(&(len as u32).to_le_bytes())?;
writer.write_all(payload)
}
pub fn read_frame<R: Read>(reader: &mut R, partial: &mut Option<PartialFrame>) -> std::io::Result<Option<Vec<u8>>> {
let mut header: ([u8; 4], usize) = match partial.take() {
Some(PartialFrame::Payload {
len,
mut buf,
mut filled,
}) => {
if len > MAX_FRAME_SIZE {
return Err(std::io::Error::new(ErrorKind::InvalidData, "Frame too large"));
}
loop {
match reader.read(&mut buf[filled..]) {
Ok(0) => {
return Err(std::io::Error::new(
ErrorKind::UnexpectedEof,
"Unexpected EOF in frame payload",
));
}
Ok(n) => {
filled += n;
if filled == len {
return Ok(Some(buf));
}
}
Err(e) if e.kind() == ErrorKind::Interrupted => continue,
Err(e) => {
*partial = Some(PartialFrame::Payload { len, buf, filled });
return Err(e);
}
}
}
}
Some(PartialFrame::Header(h, filled)) => (h, filled),
None => ([0u8; 4], 0),
};
loop {
match reader.read(&mut header.0[header.1..]) {
Ok(0) => {
if header.1 == 0 {
return Ok(None);
}
return Err(std::io::Error::new(
ErrorKind::UnexpectedEof,
"Unexpected EOF in frame header",
));
}
Ok(n) => {
header.1 += n;
if header.1 == 4 {
break;
}
}
Err(e) if e.kind() == ErrorKind::Interrupted => continue,
Err(e) => {
*partial = Some(PartialFrame::Header(header.0, header.1));
return Err(e);
}
}
}
let len = u32::from_le_bytes(header.0) as usize;
if len > MAX_FRAME_SIZE {
return Err(std::io::Error::new(
ErrorKind::InvalidData,
format!("Frame too large: {len} bytes"),
));
}
if len == 0 {
return Ok(Some(Vec::new()));
}
let mut buf = vec![0u8; len];
let mut filled = 0;
loop {
match reader.read(&mut buf[filled..]) {
Ok(0) => {
return Err(std::io::Error::new(
ErrorKind::UnexpectedEof,
"Unexpected EOF in frame payload",
));
}
Ok(n) => {
filled += n;
if filled == len {
return Ok(Some(buf));
}
}
Err(e) if e.kind() == ErrorKind::Interrupted => continue,
Err(e) => {
*partial = Some(PartialFrame::Payload { len, buf, filled });
return Err(e);
}
}
}
}
pub struct RemoteDatabase {
stream: TcpStream,
address: String,
partial: Mutex<Option<PartialFrame>>,
desynced: AtomicBool,
}
impl RemoteDatabase {
pub fn connect_tcp(addr: impl Into<String>) -> Result<Self, String> {
let addr = addr.into();
let stream =
TcpStream::connect(&addr).map_err(|e| format!("Failed to connect to Akar server at '{addr}': {e}"))?;
let _ = stream.set_nodelay(true);
let _ = stream.set_read_timeout(Some(Duration::from_secs(30)));
let _ = stream.set_write_timeout(Some(Duration::from_secs(30)));
Ok(Self {
stream,
address: addr,
partial: Mutex::new(None),
desynced: AtomicBool::new(false),
})
}
pub fn address(&self) -> &str {
&self.address
}
pub fn query(&self, query_str: &str) -> Result<WireResponse, String> {
if self.desynced.load(Ordering::Acquire) {
return Err("Connection is desynchronized after a previous read timeout; \
reconnect before sending further queries"
.into());
}
let request = WireRequest {
query: query_str.to_string(),
client_name: None,
};
let payload = serde_json::to_vec(&request).map_err(|e| format!("Failed to serialize request: {e}"))?;
let mut partial = self.partial.lock().map_err(|e| format!("Lock poisoned: {e}"))?;
{
let mut writer = &self.stream;
write_frame(&mut writer, &payload).map_err(|e| format!("Failed to send request: {e}"))?;
writer.flush().map_err(|e| format!("Failed to flush request: {e}"))?;
}
let frame = match read_frame(&mut &self.stream, &mut partial) {
Ok(Some(f)) => f,
Ok(None) => return Err("Connection closed by server".to_string()),
Err(e) => {
if e.kind() != ErrorKind::TimedOut && e.kind() != ErrorKind::WouldBlock {
return Err(format!("Failed to read response: {e}"));
}
match self.drain_pending_frame(&mut partial) {
DrainOutcome::FrameConsumed => {
return Err(format!("Failed to read response (query timed out): {e}"));
}
DrainOutcome::ConnectionClosed => return Err("Connection closed by server".to_string()),
DrainOutcome::NoFrameWithinGrace => {
self.desynced.store(true, Ordering::Release);
return Err(format!(
"Failed to read response: {e} (no stale frame arrived to re-synchronize; \
the connection has been marked desynchronized — reconnect before continuing)"
));
}
}
}
};
drop(partial);
let response: WireResponse =
serde_json::from_slice(&frame).map_err(|e| format!("Failed to parse response: {e}"))?;
if response.success {
Ok(response)
} else {
Err(response
.error_message
.clone()
.unwrap_or_else(|| "Unknown server error".to_string()))
}
}
fn drain_pending_frame(&self, partial: &mut Option<PartialFrame>) -> DrainOutcome {
let _ = self.stream.set_read_timeout(Some(GRACE_TIMEOUT));
let outcome = drain_stale_frames(&mut &self.stream, partial);
let _ = self.stream.set_read_timeout(Some(Duration::from_secs(30)));
outcome
}
pub fn ping(&self) -> Result<(), String> {
self.query("RETURN 1").map(|_| ())
}
pub fn close(&self) {
let _ = self.stream.shutdown(std::net::Shutdown::Both);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_frame_roundtrip() {
let payloads = [b"".to_vec(), b"hello".to_vec(), vec![0u8; 4096], b"{}".to_vec()];
for payload in payloads {
let mut buf = Vec::new();
write_frame(&mut buf, &payload).unwrap();
let mut cursor = &buf[..];
let mut partial = None;
let read = read_frame(&mut cursor, &mut partial).unwrap();
assert_eq!(read.as_deref(), Some(payload.as_slice()));
}
}
struct ChunkedReader<'a> {
inner: &'a mut &'a [u8],
first_read: bool,
}
impl<'a> ChunkedReader<'a> {
fn new(inner: &'a mut &'a [u8]) -> Self {
Self {
inner,
first_read: true,
}
}
}
impl<'a> Read for ChunkedReader<'a> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
if self.first_read {
self.first_read = false;
return Err(std::io::Error::new(ErrorKind::WouldBlock, "no data yet"));
}
if self.inner.is_empty() {
return Ok(0);
}
let n = 3.min(buf.len()).min(self.inner.len());
buf[..n].copy_from_slice(&self.inner[..n]);
*self.inner = &self.inner[n..];
Ok(n)
}
}
#[test]
fn test_frame_partial_reads() {
let mut buf = Vec::new();
write_frame(&mut buf, b"partial-frame-test").unwrap();
let mut cursor = &buf[..];
let mut partial = None;
let mut reader = ChunkedReader::new(&mut cursor);
let err = read_frame(&mut reader, &mut partial).unwrap_err();
assert_eq!(err.kind(), ErrorKind::WouldBlock);
assert!(partial.is_some(), "partial state must be retained across timeouts");
let result = read_frame(&mut reader, &mut partial).unwrap();
assert_eq!(result.as_deref(), Some(b"partial-frame-test".as_slice()));
}
#[test]
fn test_drain_stale_frames_consumes_stale_response() {
let stale_response = serde_json::to_vec(&WireResponse::success_message("slow query".into())).unwrap();
let mut buf = Vec::new();
write_frame(&mut buf, &stale_response).unwrap();
let mut cursor = &buf[..];
let mut partial = None;
let mut reader = ChunkedReader::new(&mut cursor);
let err = read_frame(&mut reader, &mut partial).unwrap_err();
assert_eq!(err.kind(), ErrorKind::WouldBlock);
let outcome = drain_stale_frames(&mut reader, &mut partial);
assert_eq!(outcome, DrainOutcome::FrameConsumed);
let mut tail = Vec::new();
let _ = reader.read_to_end(&mut tail);
assert_eq!(tail.len(), 0);
}
#[test]
fn test_drain_stale_frames_no_frame_returns_grace() {
struct AlwaysBlocking;
impl Read for AlwaysBlocking {
fn read(&mut self, _buf: &mut [u8]) -> std::io::Result<usize> {
Err(std::io::Error::new(ErrorKind::WouldBlock, "no data"))
}
}
let mut reader = AlwaysBlocking;
let mut partial = None;
let outcome = drain_stale_frames(&mut reader, &mut partial);
assert_eq!(outcome, DrainOutcome::NoFrameWithinGrace);
}
#[test]
fn test_drain_stale_frames_eof_reports_closed() {
let mut cursor = &b""[..];
let mut partial = None;
let outcome = drain_stale_frames(&mut &mut cursor, &mut partial);
assert_eq!(outcome, DrainOutcome::ConnectionClosed);
}
#[test]
fn test_eof_returns_none() {
let mut cursor = &b""[..];
let mut partial = None;
assert!(read_frame(&mut cursor, &mut partial).unwrap().is_none());
}
#[test]
fn test_frame_too_large_rejected() {
let mut buf = Vec::new();
buf.extend_from_slice(&(MAX_FRAME_SIZE as u32 + 1).to_le_bytes());
let mut cursor = &buf[..];
let mut partial = None;
assert!(read_frame(&mut cursor, &mut partial).is_err());
}
#[test]
fn test_wire_response_accessors() {
let resp = WireResponse {
success: true,
message: None,
error_message: None,
column_names: vec!["name".into(), "age".into()],
rows: vec![
vec![Some(Value::String("alice".into())), Some(Value::Int64(30))],
vec![None, Some(Value::Int64(25))],
],
};
assert_eq!(resp.num_rows(), 2);
assert_eq!(resp.num_columns(), 2);
assert_eq!(resp.cell(0, 0), Some(&Value::String("alice".into())));
assert_eq!(resp.cell(1, 0), None);
assert_eq!(resp.cell(9, 9), None);
assert_eq!(resp.column_values(1), vec![Value::Int64(30), Value::Int64(25)]);
assert_eq!(resp.result_summary(), "Returned 2 rows in 2 columns");
}
}