1use akar_common::types::Value;
26use serde::{Deserialize, Serialize};
27use std::collections::HashMap;
28use std::fmt;
29use std::io::{ErrorKind, Read, Write};
30use std::net::TcpStream;
31use std::sync::Mutex;
32use std::sync::atomic::{AtomicBool, Ordering};
33use std::time::Duration;
34
35pub const DEFAULT_PORT: u16 = 9876;
37
38pub const MAX_FRAME_SIZE: usize = 128 * 1024 * 1024;
43
44#[derive(Debug, Clone, Serialize, Deserialize)]
46pub struct WireRequest {
47 #[serde(default, skip_serializing_if = "String::is_empty")]
49 pub query: String,
50 #[serde(default, skip_serializing_if = "Option::is_none")]
52 pub client_name: Option<String>,
53 #[serde(default, skip_serializing_if = "Option::is_none")]
63 pub op: Option<String>,
64 #[serde(default, skip_serializing_if = "Option::is_none")]
67 pub token: Option<String>,
68 #[serde(default, skip_serializing_if = "Option::is_none")]
70 pub path: Option<String>,
71 #[serde(default, skip_serializing_if = "Option::is_none")]
78 pub params: Option<HashMap<String, serde_json::Value>>,
79 #[serde(default, skip_serializing_if = "Option::is_none")]
81 pub action: Option<String>,
82}
83
84#[derive(Debug, Clone, Serialize, Deserialize)]
88pub struct WireResponse {
89 pub success: bool,
90 #[serde(default, skip_serializing_if = "Option::is_none")]
91 pub message: Option<String>,
92 #[serde(default, skip_serializing_if = "Option::is_none")]
93 pub error_message: Option<String>,
94 pub column_names: Vec<String>,
95 pub rows: Vec<Vec<Option<Value>>>,
96 #[serde(default, skip_serializing_if = "Option::is_none")]
98 pub stats: Option<ServerStats>,
99}
100
101#[derive(Debug, Clone, Serialize, Deserialize)]
103pub struct ServerStats {
104 pub num_clients: usize,
106 pub total_queries: u64,
108 pub uptime_secs: u64,
110 pub db_path: String,
112 pub pid: u32,
114}
115
116impl WireResponse {
117 pub fn success_message(msg: String) -> Self {
119 Self {
120 success: true,
121 message: Some(msg),
122 error_message: None,
123 column_names: Vec::new(),
124 rows: Vec::new(),
125 stats: None,
126 }
127 }
128
129 pub fn error(msg: String) -> Self {
131 Self {
132 success: false,
133 message: None,
134 error_message: Some(msg),
135 column_names: Vec::new(),
136 rows: Vec::new(),
137 stats: None,
138 }
139 }
140
141 pub fn num_rows(&self) -> usize {
143 self.rows.len()
144 }
145
146 pub fn num_columns(&self) -> usize {
148 self.column_names.len()
149 }
150
151 pub fn cell(&self, row: usize, col: usize) -> Option<&Value> {
153 self.rows.get(row)?.get(col)?.as_ref()
154 }
155
156 pub fn column_values(&self, col: usize) -> Vec<Value> {
158 self.rows.iter().filter_map(|r| r.get(col).cloned().flatten()).collect()
159 }
160
161 pub fn result_summary(&self) -> String {
164 if let Some(ref stats) = self.stats {
165 return format!(
166 "Server stats: {} clients, {} queries, uptime {}s, pid {}",
167 stats.num_clients, stats.total_queries, stats.uptime_secs, stats.pid,
168 );
169 }
170 if let Some(head) = crate::query_result::result_summary_head(
171 self.message.as_deref(),
172 self.success,
173 self.error_message.as_deref(),
174 !self.rows.is_empty(),
175 ) {
176 return head;
177 }
178 format!(
179 "Returned {} rows in {} columns",
180 self.rows.len(),
181 self.column_names.len()
182 )
183 }
184}
185
186impl fmt::Display for WireResponse {
187 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
188 if let Some(ref stats) = self.stats {
189 return write!(
190 f,
191 "Server: {} clients, {} queries, uptime {}s, pid {}",
192 stats.num_clients, stats.total_queries, stats.uptime_secs, stats.pid,
193 );
194 }
195 if let Some(head) = crate::query_result::result_summary_head(
196 self.message.as_deref(),
197 self.success,
198 self.error_message.as_deref(),
199 !self.rows.is_empty(),
200 ) {
201 return write!(f, "{head}");
202 }
203 for (i, row) in self.rows.iter().enumerate() {
204 if i > 0 {
205 writeln!(f)?;
206 }
207 write!(f, "Row {}: ", i)?;
208 for (col, cell) in row.iter().enumerate() {
209 if col > 0 {
210 write!(f, ", ")?;
211 }
212 match cell {
213 Some(v) => write!(f, "{v:?}")?,
214 None => write!(f, "null")?,
215 }
216 }
217 }
218 Ok(())
219 }
220}
221
222#[derive(Debug)]
228pub enum PartialFrame {
229 Header([u8; 4], usize),
231 Payload { len: usize, buf: Vec<u8>, filled: usize },
233}
234
235#[derive(Debug, Clone, Copy, PartialEq, Eq)]
237enum DrainOutcome {
238 FrameConsumed,
240 ConnectionClosed,
242 NoFrameWithinGrace,
244}
245
246const GRACE_TIMEOUT: Duration = Duration::from_millis(250);
248
249const MAX_GRACE_READS: usize = 8; fn drain_stale_frames<R: Read>(reader: &mut R, partial: &mut Option<PartialFrame>) -> DrainOutcome {
258 for _ in 0..MAX_GRACE_READS {
259 match read_frame(reader, partial) {
260 Ok(Some(_frame)) => return DrainOutcome::FrameConsumed,
261 Ok(None) => return DrainOutcome::ConnectionClosed,
262 Err(e) if e.kind() == ErrorKind::TimedOut || e.kind() == ErrorKind::WouldBlock => continue,
263 Err(_) => return DrainOutcome::NoFrameWithinGrace,
264 }
265 }
266 DrainOutcome::NoFrameWithinGrace
267}
268
269pub fn write_frame<W: Write>(writer: &mut W, payload: &[u8]) -> std::io::Result<()> {
271 let len = payload.len();
272 if len > MAX_FRAME_SIZE {
273 return Err(std::io::Error::new(
274 ErrorKind::InvalidData,
275 format!("Frame too large: {len} bytes"),
276 ));
277 }
278 writer.write_all(&(len as u32).to_le_bytes())?;
279 writer.write_all(payload)
280}
281
282pub fn read_frame<R: Read>(reader: &mut R, partial: &mut Option<PartialFrame>) -> std::io::Result<Option<Vec<u8>>> {
291 let mut header: ([u8; 4], usize) = match partial.take() {
292 Some(PartialFrame::Payload {
293 len,
294 mut buf,
295 mut filled,
296 }) => {
297 if len > MAX_FRAME_SIZE {
298 return Err(std::io::Error::new(ErrorKind::InvalidData, "Frame too large"));
299 }
300 loop {
301 match reader.read(&mut buf[filled..]) {
302 Ok(0) => {
303 return Err(std::io::Error::new(
304 ErrorKind::UnexpectedEof,
305 "Unexpected EOF in frame payload",
306 ));
307 }
308 Ok(n) => {
309 filled += n;
310 if filled == len {
311 return Ok(Some(buf));
312 }
313 }
314 Err(e) if e.kind() == ErrorKind::Interrupted => continue,
315 Err(e) => {
316 *partial = Some(PartialFrame::Payload { len, buf, filled });
317 return Err(e);
318 }
319 }
320 }
321 }
322 Some(PartialFrame::Header(h, filled)) => (h, filled),
323 None => ([0u8; 4], 0),
324 };
325
326 loop {
327 match reader.read(&mut header.0[header.1..]) {
328 Ok(0) => {
329 if header.1 == 0 {
330 return Ok(None);
331 }
332 return Err(std::io::Error::new(
333 ErrorKind::UnexpectedEof,
334 "Unexpected EOF in frame header",
335 ));
336 }
337 Ok(n) => {
338 header.1 += n;
339 if header.1 == 4 {
340 break;
341 }
342 }
343 Err(e) if e.kind() == ErrorKind::Interrupted => continue,
344 Err(e) => {
345 *partial = Some(PartialFrame::Header(header.0, header.1));
346 return Err(e);
347 }
348 }
349 }
350
351 let len = u32::from_le_bytes(header.0) as usize;
352 if len > MAX_FRAME_SIZE {
353 return Err(std::io::Error::new(
354 ErrorKind::InvalidData,
355 format!("Frame too large: {len} bytes"),
356 ));
357 }
358 if len == 0 {
359 return Ok(Some(Vec::new()));
360 }
361 let mut buf = vec![0u8; len];
362 let mut filled = 0;
363 loop {
364 match reader.read(&mut buf[filled..]) {
365 Ok(0) => {
366 return Err(std::io::Error::new(
367 ErrorKind::UnexpectedEof,
368 "Unexpected EOF in frame payload",
369 ));
370 }
371 Ok(n) => {
372 filled += n;
373 if filled == len {
374 return Ok(Some(buf));
375 }
376 }
377 Err(e) if e.kind() == ErrorKind::Interrupted => continue,
378 Err(e) => {
379 *partial = Some(PartialFrame::Payload { len, buf, filled });
380 return Err(e);
381 }
382 }
383 }
384}
385
386pub struct RemoteDatabase {
397 stream: TcpStream,
398 address: String,
399 partial: Mutex<Option<PartialFrame>>,
400 desynced: AtomicBool,
405 token: Option<String>,
407}
408
409impl RemoteDatabase {
410 pub fn connect_tcp(addr: impl Into<String>) -> Result<Self, String> {
412 let addr = addr.into();
413 let stream =
414 TcpStream::connect(&addr).map_err(|e| format!("Failed to connect to Akar server at '{addr}': {e}"))?;
415 let _ = stream.set_nodelay(true);
416 let _ = stream.set_read_timeout(Some(Duration::from_secs(30)));
417 let _ = stream.set_write_timeout(Some(Duration::from_secs(30)));
418 Ok(Self {
419 stream,
420 address: addr,
421 partial: Mutex::new(None),
422 desynced: AtomicBool::new(false),
423 token: None,
424 })
425 }
426
427 pub fn connect_with_token(addr: impl Into<String>, token: String) -> Result<Self, String> {
430 let mut client = Self::connect_tcp(addr)?;
431 client.token = Some(token);
432 Ok(client)
433 }
434
435 pub fn address(&self) -> &str {
437 &self.address
438 }
439
440 pub fn set_token(&mut self, token: String) {
442 self.token = Some(token);
443 }
444
445 pub fn query(&self, query_str: &str) -> Result<WireResponse, String> {
450 self.send_request(WireRequest {
451 query: query_str.to_string(),
452 client_name: None,
453 op: None,
454 token: self.token.clone(),
455 path: None,
456 params: None,
457 action: None,
458 })
459 }
460
461 pub fn query_with_params(
466 &self,
467 query_str: &str,
468 params: HashMap<String, serde_json::Value>,
469 ) -> Result<WireResponse, String> {
470 self.send_request(WireRequest {
471 query: query_str.to_string(),
472 client_name: None,
473 op: None,
474 token: self.token.clone(),
475 path: None,
476 params: Some(params),
477 action: None,
478 })
479 }
480
481 pub fn ping_op(&self) -> Result<WireResponse, String> {
483 self.send_request(WireRequest {
484 query: String::new(),
485 client_name: None,
486 op: Some("ping".to_string()),
487 token: self.token.clone(),
488 path: None,
489 params: None,
490 action: None,
491 })
492 }
493
494 pub fn flush(&self) -> Result<WireResponse, String> {
496 self.send_request(WireRequest {
497 query: String::new(),
498 client_name: None,
499 op: Some("flush".to_string()),
500 token: self.token.clone(),
501 path: None,
502 params: None,
503 action: None,
504 })
505 }
506
507 pub fn stats(&self) -> Result<WireResponse, String> {
509 self.send_request(WireRequest {
510 query: String::new(),
511 client_name: None,
512 op: Some("stats".to_string()),
513 token: self.token.clone(),
514 path: None,
515 params: None,
516 action: None,
517 })
518 }
519
520 pub fn export_db(&self, path: &str) -> Result<WireResponse, String> {
522 self.send_request(WireRequest {
523 query: String::new(),
524 client_name: None,
525 op: Some("export".to_string()),
526 token: self.token.clone(),
527 path: Some(path.to_string()),
528 params: None,
529 action: None,
530 })
531 }
532
533 pub fn shutdown_server(&self) -> Result<WireResponse, String> {
535 self.send_request(WireRequest {
536 query: String::new(),
537 client_name: None,
538 op: Some("shutdown".to_string()),
539 token: self.token.clone(),
540 path: None,
541 params: None,
542 action: None,
543 })
544 }
545
546 pub fn dream_control(&self, action: &str) -> Result<WireResponse, String> {
548 self.send_request(WireRequest {
549 query: String::new(),
550 client_name: None,
551 op: Some("dream_control".to_string()),
552 token: self.token.clone(),
553 path: None,
554 params: None,
555 action: Some(action.to_string()),
556 })
557 }
558
559 fn send_request(&self, request: WireRequest) -> Result<WireResponse, String> {
561 if self.desynced.load(Ordering::Acquire) {
562 return Err("Connection is desynchronized after a previous read timeout; \
563 reconnect before sending further queries"
564 .into());
565 }
566
567 let payload = serde_json::to_vec(&request).map_err(|e| format!("Failed to serialize request: {e}"))?;
568
569 let mut partial = self.partial.lock().map_err(|e| format!("Lock poisoned: {e}"))?;
574
575 {
576 let mut writer = &self.stream;
577 write_frame(&mut writer, &payload).map_err(|e| format!("Failed to send request: {e}"))?;
578 writer.flush().map_err(|e| format!("Failed to flush request: {e}"))?;
579 }
580
581 let frame = match read_frame(&mut &self.stream, &mut partial) {
582 Ok(Some(f)) => f,
583 Ok(None) => return Err("Connection closed by server".to_string()),
584 Err(e) => {
585 if e.kind() != ErrorKind::TimedOut && e.kind() != ErrorKind::WouldBlock {
589 return Err(format!("Failed to read response: {e}"));
590 }
591 match self.drain_pending_frame(&mut partial) {
595 DrainOutcome::FrameConsumed => {
596 return Err(format!("Failed to read response (query timed out): {e}"));
598 }
599 DrainOutcome::ConnectionClosed => return Err("Connection closed by server".to_string()),
600 DrainOutcome::NoFrameWithinGrace => {
601 self.desynced.store(true, Ordering::Release);
602 return Err(format!(
603 "Failed to read response: {e} (no stale frame arrived to re-synchronize; \
604 the connection has been marked desynchronized — reconnect before continuing)"
605 ));
606 }
607 }
608 }
609 };
610 drop(partial);
611
612 let response: WireResponse =
613 serde_json::from_slice(&frame).map_err(|e| format!("Failed to parse response: {e}"))?;
614 if response.success {
615 Ok(response)
616 } else {
617 Err(response
618 .error_message
619 .clone()
620 .unwrap_or_else(|| "Unknown server error".to_string()))
621 }
622 }
623
624 fn drain_pending_frame(&self, partial: &mut Option<PartialFrame>) -> DrainOutcome {
630 let _ = self.stream.set_read_timeout(Some(GRACE_TIMEOUT));
631 let outcome = drain_stale_frames(&mut &self.stream, partial);
632 let _ = self.stream.set_read_timeout(Some(Duration::from_secs(30)));
633 outcome
634 }
635
636 pub fn ping(&self) -> Result<(), String> {
638 self.query("RETURN 1").map(|_| ())
639 }
640
641 pub fn close(&self) {
644 let _ = self.stream.shutdown(std::net::Shutdown::Both);
645 }
646}
647
648#[cfg(test)]
649mod tests {
650 use super::*;
651
652 #[test]
653 fn test_frame_roundtrip() {
654 let payloads = [b"".to_vec(), b"hello".to_vec(), vec![0u8; 4096], b"{}".to_vec()];
655 for payload in payloads {
656 let mut buf = Vec::new();
657 write_frame(&mut buf, &payload).unwrap();
658 let mut cursor = &buf[..];
659 let mut partial = None;
660 let read = read_frame(&mut cursor, &mut partial).unwrap();
661 assert_eq!(read.as_deref(), Some(payload.as_slice()));
662 }
663 }
664
665 struct ChunkedReader<'a> {
668 inner: &'a mut &'a [u8],
669 first_read: bool,
670 }
671
672 impl<'a> ChunkedReader<'a> {
673 fn new(inner: &'a mut &'a [u8]) -> Self {
674 Self {
675 inner,
676 first_read: true,
677 }
678 }
679 }
680
681 impl<'a> Read for ChunkedReader<'a> {
682 fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
683 if self.first_read {
684 self.first_read = false;
685 return Err(std::io::Error::new(ErrorKind::WouldBlock, "no data yet"));
688 }
689 if self.inner.is_empty() {
690 return Ok(0);
691 }
692 let n = 3.min(buf.len()).min(self.inner.len());
693 buf[..n].copy_from_slice(&self.inner[..n]);
694 *self.inner = &self.inner[n..];
695 Ok(n)
696 }
697 }
698
699 #[test]
700 fn test_frame_partial_reads() {
701 let mut buf = Vec::new();
702 write_frame(&mut buf, b"partial-frame-test").unwrap();
703
704 let mut cursor = &buf[..];
705 let mut partial = None;
706 let mut reader = ChunkedReader::new(&mut cursor);
707
708 let err = read_frame(&mut reader, &mut partial).unwrap_err();
710 assert_eq!(err.kind(), ErrorKind::WouldBlock);
711 assert!(partial.is_some(), "partial state must be retained across timeouts");
712
713 let result = read_frame(&mut reader, &mut partial).unwrap();
715 assert_eq!(result.as_deref(), Some(b"partial-frame-test".as_slice()));
716 }
717
718 #[test]
722 fn test_drain_stale_frames_consumes_stale_response() {
723 let stale_response = serde_json::to_vec(&WireResponse::success_message("slow query".into())).unwrap();
724 let mut buf = Vec::new();
725 write_frame(&mut buf, &stale_response).unwrap();
726
727 let mut cursor = &buf[..];
728 let mut partial = None;
729 let mut reader = ChunkedReader::new(&mut cursor);
730
731 let err = read_frame(&mut reader, &mut partial).unwrap_err();
733 assert_eq!(err.kind(), ErrorKind::WouldBlock);
734
735 let outcome = drain_stale_frames(&mut reader, &mut partial);
738 assert_eq!(outcome, DrainOutcome::FrameConsumed);
739 let mut tail = Vec::new();
741 let _ = reader.read_to_end(&mut tail);
742 assert_eq!(tail.len(), 0);
743 }
744
745 #[test]
746 fn test_drain_stale_frames_no_frame_returns_grace() {
747 struct AlwaysBlocking;
750 impl Read for AlwaysBlocking {
751 fn read(&mut self, _buf: &mut [u8]) -> std::io::Result<usize> {
752 Err(std::io::Error::new(ErrorKind::WouldBlock, "no data"))
753 }
754 }
755
756 let mut reader = AlwaysBlocking;
757 let mut partial = None;
758 let outcome = drain_stale_frames(&mut reader, &mut partial);
759 assert_eq!(outcome, DrainOutcome::NoFrameWithinGrace);
760 }
761
762 #[test]
763 fn test_drain_stale_frames_eof_reports_closed() {
764 let mut cursor = &b""[..];
765 let mut partial = None;
766 let outcome = drain_stale_frames(&mut &mut cursor, &mut partial);
767 assert_eq!(outcome, DrainOutcome::ConnectionClosed);
768 }
769
770 #[test]
771 fn test_eof_returns_none() {
772 let mut cursor = &b""[..];
773 let mut partial = None;
774 assert!(read_frame(&mut cursor, &mut partial).unwrap().is_none());
775 }
776
777 #[test]
778 fn test_frame_too_large_rejected() {
779 let mut buf = Vec::new();
780 buf.extend_from_slice(&(MAX_FRAME_SIZE as u32 + 1).to_le_bytes());
782 let mut cursor = &buf[..];
783 let mut partial = None;
784 assert!(read_frame(&mut cursor, &mut partial).is_err());
785 }
786
787 #[test]
788 fn test_wire_response_accessors() {
789 let resp = WireResponse {
790 success: true,
791 message: None,
792 error_message: None,
793 column_names: vec!["name".into(), "age".into()],
794 rows: vec![
795 vec![Some(Value::String("alice".into())), Some(Value::Int64(30))],
796 vec![None, Some(Value::Int64(25))],
797 ],
798 stats: None,
799 };
800 assert_eq!(resp.num_rows(), 2);
801 assert_eq!(resp.num_columns(), 2);
802 assert_eq!(resp.cell(0, 0), Some(&Value::String("alice".into())));
803 assert_eq!(resp.cell(1, 0), None);
804 assert_eq!(resp.cell(9, 9), None);
805 assert_eq!(resp.column_values(1), vec![Value::Int64(30), Value::Int64(25)]);
806 assert_eq!(resp.result_summary(), "Returned 2 rows in 2 columns");
807 }
808}