1use std::collections::HashMap;
142use std::sync::Arc;
143
144use serde_json::Value;
145use tokio::io::{AsyncReadExt, AsyncWriteExt};
146use tokio::net::{TcpListener, TcpStream};
147
148use crate::db::Db;
149
150const OID_BOOL: i32 = 16;
152const OID_INT8: i32 = 20;
153const OID_FLOAT8: i32 = 701;
154const OID_TEXT: i32 = 25;
155
156const PROTO_V3: i32 = 196_608; const SSL_REQUEST: i32 = 80_877_103;
158const GSS_REQUEST: i32 = 80_877_104;
159const CANCEL_REQUEST: i32 = 80_877_102;
160
161pub trait DbResolver: Send + Sync + 'static {
167 fn resolve(&self, name: &str) -> Option<Arc<Db>>;
175 fn token(&self) -> Option<String> {
177 None
178 }
179}
180
181struct Out(Vec<u8>);
184
185impl Out {
186 fn msg(tag: u8) -> Self {
187 Out(vec![tag, 0, 0, 0, 0])
189 }
190 fn i16(&mut self, v: i16) { self.0.extend_from_slice(&v.to_be_bytes()); }
191 fn i32(&mut self, v: i32) { self.0.extend_from_slice(&v.to_be_bytes()); }
192 fn cstr(&mut self, s: &str) {
193 self.0.extend_from_slice(s.replace('\0', "").as_bytes());
196 self.0.push(0);
197 }
198 fn bytes(&mut self, b: &[u8]) { self.0.extend_from_slice(b); }
199 fn finish(mut self) -> Vec<u8> {
201 let len = (self.0.len() - 1) as i32;
202 self.0[1..5].copy_from_slice(&len.to_be_bytes());
203 self.0
204 }
205}
206
207fn err_msg(code: &str, message: &str) -> Vec<u8> {
208 let mut m = Out::msg(b'E');
209 m.bytes(b"S"); m.cstr("ERROR");
210 m.bytes(b"C"); m.cstr(code);
211 m.bytes(b"M"); m.cstr(message);
212 m.0.push(0);
213 m.finish()
214}
215
216fn ready() -> Vec<u8> {
217 let mut m = Out::msg(b'Z');
218 m.bytes(b"I"); m.finish()
220}
221
222fn command_complete(tag: &str) -> Vec<u8> {
223 let mut m = Out::msg(b'C');
224 m.cstr(tag);
225 m.finish()
226}
227
228#[derive(Debug, PartialEq, Clone)]
238pub struct Col {
239 pub src: String,
240 pub out: String,
241}
242
243impl Col {
244 fn same(name: &str) -> Self {
245 Col { src: name.to_string(), out: name.to_string() }
246 }
247 fn renamed(src: &str, out: &str) -> Self {
248 Col { src: src.to_string(), out: out.to_string() }
249 }
250}
251
252#[derive(Debug, PartialEq)]
270pub enum Stmt {
271 Query { nql: String, project: Vec<Col> },
273 Insert { coll: String, rows: Vec<InsertRow>, returning: Vec<Col> },
275 Update { coll: String, set: Vec<(String, Value)>, nql: String, returning: Vec<Col> },
277 Delete { coll: String, nql: String, returning: Vec<Col> },
279 Canned { cols: Vec<String>, row: Vec<String> },
281 Ok(&'static str),
283}
284
285#[derive(Debug, PartialEq, Clone)]
288pub struct InsertRow {
289 pub id: Option<String>,
291 pub doc: serde_json::Map<String, Value>,
292 pub caused_by: Vec<String>,
295 pub valid_from: Option<String>,
296 pub valid_to: Option<String>,
297}
298
299fn normalise(sql: &str) -> String {
302 let mut out = String::with_capacity(sql.len());
303 let mut chars = sql.chars().peekable();
304 let mut in_s = false;
305 while let Some(c) = chars.next() {
306 if in_s {
307 out.push(c);
308 if c == '\'' { in_s = false; }
309 continue;
310 }
311 match c {
312 '\'' => { in_s = true; out.push(c); }
313 '-' if chars.peek() == Some(&'-') => {
314 for n in chars.by_ref() { if n == '\n' { break; } }
316 out.push(' ');
317 }
318 '/' if chars.peek() == Some(&'*') => {
319 chars.next();
320 let mut prev = ' ';
321 while let Some(n) = chars.next() {
322 if prev == '*' && n == '/' { break; }
323 prev = n;
324 }
325 out.push(' ');
326 }
327 _ => out.push(c),
328 }
329 }
330 out.split_whitespace().collect::<Vec<_>>().join(" ")
331}
332
333fn strip_column_qualifiers(
370 tail: &str,
371 coll: &str,
372 alias: Option<&str>,
373) -> Result<String, String> {
374 let bare = coll.rsplit('.').next().unwrap_or(coll);
375 let b: Vec<char> = tail.chars().collect();
376 let mut out = String::with_capacity(tail.len());
377 let mut i = 0usize;
378 let ident_start = |c: char| c.is_alphabetic() || c == '_';
379 let ident_char = |c: char| c.is_alphanumeric() || c == '_';
380
381 while i < b.len() {
382 if b[i] == '\'' {
384 out.push(b[i]);
385 i += 1;
386 while i < b.len() {
387 out.push(b[i]);
388 if b[i] == '\'' {
389 if b.get(i + 1) == Some(&'\'') {
391 out.push('\'');
392 i += 2;
393 continue;
394 }
395 i += 1;
396 break;
397 }
398 i += 1;
399 }
400 continue;
401 }
402 if b[i] == '"' {
408 out.push(b[i]);
409 i += 1;
410 while i < b.len() {
411 out.push(b[i]);
412 if b[i] == '"' { i += 1; break; }
413 i += 1;
414 }
415 continue;
416 }
417 if !ident_start(b[i]) {
418 out.push(b[i]);
421 i += 1;
422 continue;
423 }
424
425 let start = i;
426 while i < b.len() && ident_char(b[i]) {
427 i += 1;
428 }
429 let word: String = b[start..i].iter().collect();
430
431 if b.get(i) == Some(&'.') && b.get(i + 1).is_some_and(|c| ident_start(*c)) {
433 let fstart = i + 1;
434 let mut j = fstart;
435 while j < b.len() && ident_char(b[j]) {
436 j += 1;
437 }
438 let field: String = b[fstart..j].iter().collect();
439 let is_call = b[j..].iter().find(|c| !c.is_whitespace()) == Some(&'(');
444 if is_call {
445 out.push_str(&word);
446 out.push('.');
447 out.push_str(&field);
448 i = j;
449 continue;
450 }
451 let matches_alias = alias.is_some_and(|a| word.eq_ignore_ascii_case(a));
452 if matches_alias || word.eq_ignore_ascii_case(bare) || word.eq_ignore_ascii_case(coll) {
453 out.push_str(&field);
454 i = j;
455 continue;
456 }
457 return Err(format!(
458 "no table or alias named {:?} in this query — this statement reads \
459 {:?}{}, and a qualifier naming anything else would have to be \
460 answered from a relation that is not in its FROM clause",
461 word,
462 bare,
463 alias.map(|a| format!(" (aliased {:?})", a)).unwrap_or_default()
464 ));
465 }
466 out.push_str(&word);
467 }
468 Ok(out)
469}
470
471fn flatten_count_of_subquery(projection: &str, rest: &str) -> Option<String> {
479 let (outer_expr, outer_alias) = split_output_alias(projection.trim());
483 let ou = outer_expr.to_uppercase().replace(' ', "");
484 if ou != "COUNT(*)" {
485 return None;
486 }
487
488 let b: Vec<char> = rest.chars().collect();
491 let mut depth = 0i32;
492 let mut in_s = false;
493 let mut end = None;
494 for (i, &c) in b.iter().enumerate() {
495 match c {
496 '\'' => in_s = !in_s,
497 '(' if !in_s => depth += 1,
498 ')' if !in_s => {
499 depth -= 1;
500 if depth == 0 {
501 end = Some(i);
502 break;
503 }
504 }
505 _ => {}
506 }
507 }
508 let end = end?;
509 let inner = b[1..end].iter().collect::<String>().trim().to_string();
510
511 let trailing = b[end + 1..].iter().collect::<String>();
514 let (_alias, after) = split_table_alias(trailing.trim());
515 if !after.trim().is_empty() {
516 return None;
517 }
518
519 let iu = inner.to_uppercase();
520 if !iu.starts_with("SELECT") {
521 return None;
522 }
523 for kw in ["LIMIT", "OFFSET", "GROUP BY", "HAVING", "UNION", "INTERSECT", "EXCEPT", "JOIN"] {
525 if find_kw(&iu, kw).is_some() {
526 return None;
527 }
528 }
529 if find_kw(&iu, "DISTINCT").is_some() {
530 return None;
531 }
532 let inner_from = find_kw(&iu, "FROM")?;
534 let inner_list = inner[..inner_from].to_uppercase();
535 for agg in ["COUNT(", "SUM(", "AVG(", "MIN(", "MAX(", "ARRAY_AGG(", "STRING_AGG("] {
536 if inner_list.contains(agg) {
537 return None;
538 }
539 }
540 let inner_rest = inner[inner_from + 4..].trim();
542 if inner_rest.starts_with('(') {
543 return None;
544 }
545
546 let mut tail = inner_rest.to_string();
548 let tu = tail.to_uppercase();
549 if let Some(ob) = find_kw(&tu, "ORDER BY") {
550 tail = tail[..ob].trim_end().to_string();
551 }
552 Some(format!(
553 "SELECT count(*){} FROM {}",
554 outer_alias.map(|a| format!(" AS {}", a)).unwrap_or_default(),
555 tail
556 ))
557}
558
559const CLAUSE_WORDS: &[&str] = &[
564 "WHERE", "GROUP", "ORDER", "LIMIT", "OFFSET", "HAVING", "FOR", "VALID",
565 "TRACE", "TRAVERSE", "SEARCH", "RETURNING", "UNION", "INTERSECT", "EXCEPT",
566 "JOIN", "LEFT", "RIGHT", "INNER", "FULL", "CROSS", "ON", "USING", "SET",
567];
568
569fn split_table_alias(tail: &str) -> (Option<String>, &str) {
580 let t = tail.trim_start();
581 let first_end = t.find(char::is_whitespace).unwrap_or(t.len());
582 let first = &t[..first_end];
583 let fu = first.to_uppercase();
584
585 if fu == "AS" {
586 let rest = t[first_end..].trim_start();
587 let end = rest.find(char::is_whitespace).unwrap_or(rest.len());
588 let word = &rest[..end];
589 if word.eq_ignore_ascii_case("OF") {
590 return (None, t); }
592 if word.is_empty() {
593 return (None, t);
594 }
595 return (Some(word.trim_matches('"').to_string()), rest[end..].trim_start());
596 }
597 if first.is_empty() || CLAUSE_WORDS.contains(&fu.as_str()) {
598 return (None, t);
599 }
600 if first.chars().next().is_some_and(|c| c.is_alphabetic() || c == '_' || c == '"') {
603 return (Some(first.trim_matches('"').to_string()), t[first_end..].trim_start());
604 }
605 (None, t)
606}
607
608fn split_top_level(s: &str, delim: char) -> Vec<String> {
613 let mut out = vec![];
614 let mut cur = String::new();
615 let mut depth = 0i32;
616 let mut in_s = false;
617 let mut in_d = false;
618 for c in s.chars() {
619 match c {
620 '\'' if !in_d => { in_s = !in_s; cur.push(c); }
621 '"' if !in_s => { in_d = !in_d; cur.push(c); }
622 '(' if !in_s && !in_d => { depth += 1; cur.push(c); }
623 ')' if !in_s && !in_d => { depth -= 1; cur.push(c); }
624 c if c == delim && depth == 0 && !in_s && !in_d => {
625 out.push(std::mem::take(&mut cur));
626 }
627 _ => cur.push(c),
628 }
629 }
630 out.push(cur);
631 out
632}
633
634fn split_output_alias(p: &str) -> (&str, Option<&str>) {
640 let pu = p.to_uppercase();
641 if let Some(at) = find_kw(&pu, "AS") {
642 let alias = p[at + 2..].trim().trim_matches('"');
643 if !alias.is_empty() {
644 return (p[..at].trim(), Some(alias));
645 }
646 }
647 let b: Vec<char> = p.chars().collect();
651 let mut depth = 0i32;
652 let mut in_s = false;
653 let mut cut = None;
654 for (i, &c) in b.iter().enumerate() {
655 match c {
656 '\'' => in_s = !in_s,
657 '(' if !in_s => depth += 1,
658 ')' if !in_s => depth -= 1,
659 c if c.is_whitespace() && depth == 0 && !in_s => cut = Some(i),
660 _ => {}
661 }
662 }
663 match cut {
664 Some(i) => {
665 let alias = p[i..].trim().trim_matches('"');
666 if alias.is_empty() { (p, None) } else { (p[..i].trim(), Some(alias)) }
667 }
668 None => (p, None),
669 }
670}
671
672fn sql_literals_to_nql(s: &str) -> String {
673 let mut out = String::with_capacity(s.len());
674 let mut it = s.chars().peekable();
675 while let Some(c) = it.next() {
676 match c {
677 '\'' => {
678 out.push('"');
679 while let Some(ch) = it.next() {
680 if ch == '\'' {
681 if it.peek() == Some(&'\'') {
682 it.next();
683 out.push('\''); } else {
685 break;
686 }
687 } else if ch == '"' {
688 out.push('\\');
691 out.push('"');
692 } else {
693 out.push(ch);
694 }
695 }
696 out.push('"');
697 }
698 '<' if it.peek() == Some(&'>') => { it.next(); out.push_str("!="); }
699 _ => out.push(c),
700 }
701 }
702 out
703}
704
705fn strip_prefix_ci(s: &str, prefix: &str) -> Option<String> {
706 if s.len() >= prefix.len() && s[..prefix.len()].eq_ignore_ascii_case(prefix) {
707 Some(s[prefix.len()..].trim_start().to_string())
708 } else {
709 None
710 }
711}
712
713fn find_kw(s: &str, kw: &str) -> Option<usize> {
716 let bytes = s.as_bytes();
717 let k = kw.as_bytes();
718 let mut depth = 0i32;
719 let mut in_s = false;
720 let mut in_d = false;
721 let mut i = 0usize;
722 while i < bytes.len() {
723 let c = bytes[i];
724 if in_s { if c == b'\'' { in_s = false; } i += 1; continue; }
725 if in_d { if c == b'"' { in_d = false; } i += 1; continue; }
726 match c {
727 b'\'' => { in_s = true; i += 1; continue; }
728 b'"' => { in_d = true; i += 1; continue; }
729 b'(' => { depth += 1; i += 1; continue; }
730 b')' => { depth -= 1; i += 1; continue; }
731 _ => {}
732 }
733 if depth == 0 && i + k.len() <= bytes.len()
734 && bytes[i..i + k.len()].eq_ignore_ascii_case(k)
735 {
736 let before_ok = i == 0 || !(bytes[i - 1] as char).is_alphanumeric() && bytes[i - 1] != b'_';
737 let after = i + k.len();
738 let after_ok = after >= bytes.len()
739 || !(bytes[after] as char).is_alphanumeric() && bytes[after] != b'_';
740 if before_ok && after_ok {
741 return Some(i);
742 }
743 }
744 i += 1;
745 }
746 None
747}
748
749fn split_top(s: &str, sep: char) -> Vec<String> {
753 let mut out = vec![];
754 let mut cur = String::new();
755 let mut depth = 0i32;
756 let mut in_s = false;
757 let mut it = s.chars().peekable();
758 while let Some(c) = it.next() {
759 if in_s {
760 cur.push(c);
761 if c == '\'' {
762 if it.peek() == Some(&'\'') { cur.push(it.next().unwrap()); } else { in_s = false; }
764 }
765 continue;
766 }
767 match c {
768 '\'' => { in_s = true; cur.push(c); }
769 '(' => { depth += 1; cur.push(c); }
770 ')' => { depth -= 1; cur.push(c); }
771 x if x == sep && depth == 0 => { out.push(cur.trim().to_string()); cur.clear(); }
772 _ => cur.push(c),
773 }
774 }
775 if !cur.trim().is_empty() { out.push(cur.trim().to_string()); }
776 out
777}
778
779fn sql_value(raw: &str) -> Result<Value, String> {
785 let t = raw.trim();
786 if t.is_empty() {
787 return Err("empty value".into());
788 }
789 let up = t.to_uppercase();
790 if up == "NULL" { return Ok(Value::Null); }
791 if up == "TRUE" { return Ok(Value::Bool(true)); }
792 if up == "FALSE" { return Ok(Value::Bool(false)); }
793 if t.starts_with('\'') && t.ends_with('\'') && t.len() >= 2 {
794 let inner = &t[1..t.len() - 1];
796 return Ok(Value::String(inner.replace("''", "'")));
797 }
798 if let Ok(i) = t.parse::<i64>() { return Ok(Value::from(i)); }
799 if let Ok(f) = t.parse::<f64>() { return Ok(Value::from(f)); }
800 Err(format!(
801 "cannot use {:?} as a value — this endpoint accepts string literals, \
802 numbers, TRUE/FALSE and NULL. Expressions, casts and function calls \
803 are not evaluated, because storing an unevaluated expression as text \
804 would be worse than refusing it", t))
805}
806
807fn split_returning(tail: &str) -> (String, Vec<Col>) {
809 let tu = tail.to_uppercase();
810 match find_kw(&tu, "RETURNING") {
811 None => (tail.to_string(), vec![]),
812 Some(at) => {
813 let head = tail[..at].trim().to_string();
814 let list = tail[at + "RETURNING".len()..].trim();
815 if list == "*" {
816 return (head, vec![]); }
818 let cols = split_top(list, ',')
819 .into_iter()
820 .map(|p| {
821 let raw = p.split_whitespace().next().unwrap_or(&p).to_string();
822 let name = raw.rsplit('.').next().unwrap_or(&raw).trim_matches('"').to_string();
823 Col::same(&name)
824 })
825 .collect();
826 (head, cols)
827 }
828 }
829}
830
831fn take_reserved(doc: &mut serde_json::Map<String, Value>) -> (Option<String>, Vec<String>, Option<String>, Option<String>) {
833 let id = doc.remove("_id").or_else(|| doc.remove("id"))
834 .and_then(|v| match v {
835 Value::String(s) => Some(s),
836 Value::Null => None,
837 other => Some(other.to_string()), });
839 let caused_by = match doc.remove("_caused_by") {
840 Some(Value::String(s)) => vec![s],
841 Some(Value::Array(a)) => a.into_iter()
842 .filter_map(|v| v.as_str().map(str::to_string)).collect(),
843 _ => vec![],
844 };
845 let vf = doc.remove("_valid_from").and_then(|v| v.as_str().map(str::to_string));
846 let vt = doc.remove("_valid_to").and_then(|v| v.as_str().map(str::to_string));
847 (id, caused_by, vf, vt)
848}
849
850fn translate_insert(sql: &str) -> Result<Stmt, String> {
852 let rest = strip_prefix_ci(sql, "INSERT")
853 .and_then(|r| strip_prefix_ci(&r, "INTO"))
854 .ok_or("expected INSERT INTO")?;
855 let ru = rest.to_uppercase();
859 let values_at = find_kw(&ru, "VALUES").ok_or(
860 "expected VALUES — `INSERT … SELECT` is not supported on this endpoint")?;
861 let head = rest[..values_at].trim().to_string();
862 let open = head.find('(').ok_or(
863 "INSERT needs an explicit column list — `INSERT INTO t (a, b) VALUES (…)`. \
864 NEDB is schemaless, so there is no declared column order to infer from")?;
865 let coll = head[..open].trim().trim_matches('"');
866 let coll = coll.rsplit('.').next().unwrap_or(coll).to_string();
867 if coll.is_empty() {
868 return Err("expected a collection name after INSERT INTO".into());
869 }
870 let close = head.rfind(')').ok_or("unterminated column list")?;
871 if close < open {
872 return Err("malformed column list".into());
873 }
874 let tail_from_values = rest[values_at..].to_string();
875 let cols: Vec<String> = split_top(&head[open + 1..close], ',')
876 .into_iter()
877 .map(|c| c.trim().trim_matches('"').to_string())
878 .collect();
879 if cols.is_empty() {
880 return Err("the column list is empty".into());
881 }
882
883 let after = strip_prefix_ci(&tail_from_values, "VALUES")
884 .ok_or("expected VALUES after the column list")?;
885 let (values_part, returning) = split_returning(&after);
886
887 let mut rows = vec![];
888 for group in split_top(&values_part, ',') {
889 let g = group.trim();
890 if !(g.starts_with('(') && g.ends_with(')')) {
891 return Err(format!("expected a parenthesised row of values, got {:?}", g));
892 }
893 let vals = split_top(&g[1..g.len() - 1], ',');
894 if vals.len() != cols.len() {
895 return Err(format!(
896 "{} values for {} columns — every row must match the column list",
897 vals.len(), cols.len()));
898 }
899 let mut doc = serde_json::Map::new();
900 for (c, v) in cols.iter().zip(vals.iter()) {
901 doc.insert(c.clone(), sql_value(v)?);
902 }
903 let (id, caused_by, valid_from, valid_to) = take_reserved(&mut doc);
904 rows.push(InsertRow { id, doc, caused_by, valid_from, valid_to });
905 }
906 if rows.is_empty() {
907 return Err("INSERT with no rows".into());
908 }
909 Ok(Stmt::Insert { coll, rows, returning })
910}
911
912fn translate_update(sql: &str) -> Result<Stmt, String> {
914 let rest = strip_prefix_ci(sql, "UPDATE").ok_or("expected UPDATE")?;
915 let ru = rest.to_uppercase();
916 let set_at = find_kw(&ru, "SET").ok_or("expected SET in UPDATE")?;
917 let target = rest[..set_at].trim();
920 let mut parts = target.split_whitespace();
921 let coll = parts.next().unwrap_or("").trim_matches('"');
922 let coll = coll.rsplit('.').next().unwrap_or(coll).to_string();
923 let upd_alias: Option<String> = match parts.next() {
924 Some(w) if w.eq_ignore_ascii_case("AS") => {
925 parts.next().map(|a| a.trim_matches('"').to_string())
926 }
927 Some(w) => Some(w.trim_matches('"').to_string()),
928 None => None,
929 };
930 if coll.is_empty() {
931 return Err("expected a collection name after UPDATE".into());
932 }
933 let after_set = rest[set_at + 3..].trim().to_string();
934 let (after_set, returning) = split_returning(&after_set);
935
936 let au = after_set.to_uppercase();
938 let (assigns_raw, where_raw) = match find_kw(&au, "WHERE") {
939 Some(at) => (after_set[..at].to_string(), after_set[at..].to_string()),
940 None => (after_set.clone(), String::new()),
941 };
942
943 let mut set = vec![];
944 for a in split_top(&assigns_raw, ',') {
945 let eq = a.find('=').ok_or(format!("expected `col = value` in SET, got {:?}", a))?;
946 let col = a[..eq].trim().trim_matches('"').to_string();
947 if col.is_empty() {
948 return Err("empty column name in SET".into());
949 }
950 set.push((col, sql_value(&a[eq + 1..])?));
951 }
952 if set.is_empty() {
953 return Err("UPDATE with no assignments".into());
954 }
955 let where_raw = strip_column_qualifiers(where_raw.trim(), &coll, upd_alias.as_deref())?;
958 let nql = format!("FROM {} {}", coll, sql_literals_to_nql(&where_raw))
959 .trim().to_string();
960 Ok(Stmt::Update { coll, set, nql, returning })
961}
962
963fn translate_delete(sql: &str) -> Result<Stmt, String> {
965 let rest = strip_prefix_ci(sql, "DELETE")
966 .and_then(|r| strip_prefix_ci(&r, "FROM"))
967 .ok_or("expected DELETE FROM")?;
968 let (rest, returning) = split_returning(&rest);
969 let end = rest.find(' ').unwrap_or(rest.len());
970 let coll = rest[..end].trim().trim_matches('"');
971 let coll = coll.rsplit('.').next().unwrap_or(coll).to_string();
972 if coll.is_empty() {
973 return Err("expected a collection name after DELETE FROM".into());
974 }
975 let (del_alias, where_raw) = split_table_alias(rest[end..].trim());
976 let where_raw = strip_column_qualifiers(where_raw, &coll, del_alias.as_deref())?;
977 let nql = format!("FROM {} {}", coll, sql_literals_to_nql(&where_raw))
978 .trim().to_string();
979 Ok(Stmt::Delete { coll, nql, returning })
980}
981
982pub fn translate(sql_raw: &str) -> Result<Stmt, String> {
984 let sql = normalise(sql_raw);
985 let sql = sql.trim().trim_end_matches(';').trim();
986 if sql.is_empty() {
987 return Ok(Stmt::Ok(""));
988 }
989 let upper = sql.to_uppercase();
990
991 if upper.starts_with("SET ") || upper.starts_with("BEGIN") || upper.starts_with("COMMIT")
996 || upper.starts_with("ROLLBACK") || upper.starts_with("DISCARD")
997 || upper.starts_with("LISTEN ") || upper.starts_with("UNLISTEN ")
998 {
999 return Ok(Stmt::Ok(if upper.starts_with("SET") { "SET" } else { "OK" }));
1001 }
1002 if upper.starts_with("SHOW ") {
1003 let name = sql[5..].trim().to_lowercase();
1004 let val = match name.as_str() {
1005 "transaction_isolation" | "default_transaction_isolation" => "read committed",
1006 "server_version" => SERVER_VERSION,
1007 "server_encoding" | "client_encoding" => "UTF8",
1008 "standard_conforming_strings" => "on",
1009 "is_superuser" => "off",
1010 _ => "",
1011 };
1012 return Ok(Stmt::Canned { cols: vec![name], row: vec![val.to_string()] });
1013 }
1014 if upper == "SELECT VERSION()" {
1015 return Ok(Stmt::Canned {
1016 cols: vec!["version".into()],
1017 row: vec![full_version_string()],
1018 });
1019 }
1020 if upper == "SELECT 1" || upper == "SELECT 1;" {
1021 return Ok(Stmt::Canned { cols: vec!["?column?".into()], row: vec!["1".into()] });
1022 }
1023 if upper.starts_with("SELECT CURRENT_SCHEMA") {
1024 return Ok(Stmt::Canned { cols: vec!["current_schema".into()], row: vec!["public".into()] });
1025 }
1026 if upper.starts_with("SELECT CURRENT_DATABASE") {
1027 return Ok(Stmt::Canned { cols: vec!["current_database".into()], row: vec!["nedb".into()] });
1028 }
1029 if upper.starts_with("SELECT CURRENT_USER") || upper.starts_with("SELECT USER") {
1030 return Ok(Stmt::Canned { cols: vec!["current_user".into()], row: vec!["nedb".into()] });
1031 }
1032
1033 if upper.starts_with("INSERT") { return translate_insert(sql); }
1037 if upper.starts_with("UPDATE") { return translate_update(sql); }
1038 if upper.starts_with("DELETE") { return translate_delete(sql); }
1039
1040 for (kw, why) in [
1042 ("CREATE", "DDL is not supported — collections are created implicitly by the first write to them, because NEDB is schemaless"),
1043 ("ALTER", "DDL is not supported — there is no schema to alter"),
1044 ("DROP", "DDL is not supported; drop a database with DELETE /v1/databases/<db>"),
1045 ("TRUNCATE", "not supported, and not an oversight: NEDB is append-only so that history cannot be discarded. That is the product"),
1046 ("COPY", "not supported; use GET /v1/databases/<db>/since for bulk export"),
1047 ("GRANT", "there is no SQL-level privilege system; auth is the bearer token"),
1048 ("REVOKE", "there is no SQL-level privilege system; auth is the bearer token"),
1049 ] {
1050 if upper.starts_with(kw) {
1051 return Err(format!("{} is not supported — {}", kw, why));
1052 }
1053 }
1054 if !upper.starts_with("SELECT") {
1055 return Err(format!(
1056 "only SELECT, INSERT, UPDATE and DELETE are supported on the Postgres \
1057 endpoint (got {:?})",
1058 sql.split_whitespace().next().unwrap_or("")
1059 ));
1060 }
1061 for (kw, why) in [
1062 (" JOIN ", "JOIN is not supported — NQL is single-collection; join in your client or model the relation with LINK/TRAVERSE"),
1063 (" UNION ", "UNION is not supported"),
1064 (" INTERSECT ", "INTERSECT is not supported"),
1065 (" EXCEPT ", "EXCEPT is not supported"),
1066 (" OVER (", "window functions are not supported"),
1067 ("DISTINCT ", "DISTINCT is not supported — GROUP BY <col> gives the distinct values with counts"),
1068 ] {
1069 if upper.contains(kw) {
1070 return Err(why.to_string());
1071 }
1072 }
1073 if find_kw(&upper, "FROM").is_none() {
1074 return Err("SELECT without FROM is not supported on this endpoint".into());
1075 }
1076
1077 let after_select = strip_prefix_ci(sql, "SELECT").ok_or("expected SELECT")?;
1079 let from_at = find_kw(&after_select.to_uppercase(), "FROM")
1080 .ok_or("expected FROM after the select list")?;
1081 let projection = after_select[..from_at].trim().to_string();
1082 let rest = after_select[from_at + 4..].trim().to_string();
1083 if rest.is_empty() {
1084 return Err("expected a collection name after FROM".into());
1085 }
1086 if rest.starts_with('(') {
1105 if let Some(flat) = flatten_count_of_subquery(&projection, &rest) {
1106 return translate(&flat);
1110 }
1111 return Err("subqueries in FROM are not supported — except \
1112 `SELECT count(*) FROM (…)`, which is rewritten to a flat \
1113 count when the inner query has no LIMIT, OFFSET, DISTINCT, \
1114 GROUP BY or aggregate of its own (any of those would make the \
1115 two counts different numbers)".into());
1116 }
1117 let coll_end = rest.find(' ').unwrap_or(rest.len());
1118 let coll = &rest[..coll_end];
1119 if coll.contains(',') {
1120 return Err("selecting from more than one collection is not supported (no JOIN)".into());
1121 }
1122 let bare = coll.rsplit('.').next().unwrap_or(coll).trim_matches('"');
1129 let qualified = coll
1130 .split('.')
1131 .map(|p| p.trim_matches('"'))
1132 .collect::<Vec<_>>()
1133 .join(".");
1134 let coll = if qualified.starts_with("information_schema.") {
1135 qualified.as_str()
1136 } else {
1137 bare
1138 };
1139 let tail = rest[coll_end..].trim();
1140
1141 let mut agg_clause = String::new();
1156 let mut agg_srcs: Vec<String> = vec![];
1157 let mut project: Vec<Col> = vec![];
1158
1159 if projection == "*" {
1160 } else {
1162 for part in split_top_level(&projection, ',') {
1163 let p = part.trim();
1164 if p.is_empty() {
1165 return Err("empty column in the select list".into());
1166 }
1167 let (expr, alias) = split_output_alias(p);
1168 let eu = expr.to_uppercase();
1169
1170 if eu.starts_with("COUNT(") {
1173 if agg_clause.is_empty() {
1174 agg_clause = " COUNT".to_string();
1175 }
1176 agg_srcs.push("count".to_string());
1177 project.push(Col::renamed("count", alias.unwrap_or("count")));
1178 continue;
1179 }
1180 if let Some(agg) = ["SUM", "AVG", "MIN", "MAX"]
1181 .iter()
1182 .find(|a| eu.starts_with(&format!("{}(", a)))
1183 {
1184 let inner = expr[agg.len() + 1..].trim_end_matches(')').trim();
1185 if inner.is_empty() || inner == "*" {
1186 return Err(format!("{}() needs a column", agg));
1187 }
1188 let inner = inner.rsplit('.').next().unwrap_or(inner).trim_matches('"');
1189 let named = format!("{} {}", agg, inner);
1190 if !agg_clause.is_empty() && agg_clause.trim() != "COUNT" && agg_clause.trim() != named {
1191 return Err(format!(
1192 "only one of SUM/AVG/MIN/MAX is supported per statement \
1193 (already have {:?}, then {:?}) — NQL's grouped row carries \
1194 the group key, `count`, and ONE named aggregate",
1195 agg_clause.trim(), named));
1196 }
1197 agg_clause = format!(" {}", named);
1198 let src = format!("{}_{}", agg.to_lowercase(), inner);
1201 project.push(Col::renamed(&src, alias.unwrap_or(&agg.to_lowercase())));
1202 agg_srcs.push(src);
1203 continue;
1204 }
1205 if expr.contains('(') {
1206 return Err(format!(
1207 "expressions in the select list are not supported ({:?}) — \
1208 supported: *, a column list, COUNT(*), or SUM/AVG/MIN/MAX(col)", p));
1209 }
1210 let name = expr.rsplit('.').next().unwrap_or(expr).trim_matches('"');
1211 project.push(Col::renamed(name, alias.unwrap_or(name)));
1212 }
1213 }
1214
1215 let (alias, tail) = split_table_alias(tail);
1224 let mut tail = strip_column_qualifiers(tail, coll, alias.as_deref())?;
1225 let tu = tail.to_uppercase();
1226 if let Some(at) = find_kw(&tu, "AS OF SYSTEM TIME") {
1227 let before = tail[..at].to_string();
1228 let after = tail[at + "AS OF SYSTEM TIME".len()..].trim_start().to_string();
1229 let end = after.find(' ').unwrap_or(after.len());
1231 let seq = after[..end].trim().trim_matches('\'').trim_matches('"').to_string();
1232 if seq.parse::<u64>().is_err() {
1233 return Err(format!(
1234 "AS OF SYSTEM TIME takes a NEDB sequence number here, not a timestamp (got {:?}). \
1235 NEDB's history is sequence-addressed and never garbage-collected, so a seq is \
1236 exact where a wall-clock time would be approximate", seq));
1237 }
1238 tail = format!("{} AS OF {} {}", before.trim(), seq, after[end..].trim())
1239 .trim()
1240 .to_string();
1241 }
1242
1243 let tu_ord = tail.to_uppercase();
1256 if let Some(ob_at) = find_kw(&tu_ord, "ORDER BY") {
1257 let start = ob_at + "ORDER BY".len();
1258 let end = ["LIMIT", "OFFSET", "GROUP BY", "TRACE", "TRAVERSE", "SEARCH"]
1260 .iter()
1261 .filter_map(|k| find_kw(&tu_ord[start..], k).map(|at| start + at))
1262 .min()
1263 .unwrap_or(tail.len());
1264 let mut keys = vec![];
1265 for item in split_top_level(&tail[start..end], ',') {
1266 let item = item.trim();
1267 if item.is_empty() {
1268 continue;
1269 }
1270 let mut parts = item.split_whitespace();
1271 let first = parts.next().unwrap_or("");
1272 let rest: Vec<&str> = parts.collect();
1273 match first.parse::<usize>() {
1274 Ok(n) if n >= 1 => {
1275 let col = project.get(n - 1).ok_or_else(|| {
1276 if project.is_empty() {
1277 format!(
1278 "ORDER BY {} is a select-list POSITION, and `SELECT *` \
1279 has no list to index — name the column instead", n)
1280 } else {
1281 format!(
1282 "ORDER BY {} is out of range: the select list has {} \
1283 column(s)", n, project.len())
1284 }
1285 })?;
1286 keys.push(
1287 std::iter::once(col.src.as_str())
1288 .chain(rest.iter().copied())
1289 .collect::<Vec<_>>()
1290 .join(" "),
1291 );
1292 }
1293 _ => keys.push(item.to_string()),
1296 }
1297 }
1298 tail = format!("{} ORDER BY {} {}", &tail[..ob_at], keys.join(", "), &tail[end..])
1299 .split_whitespace()
1300 .collect::<Vec<_>>()
1301 .join(" ");
1302 }
1303
1304 let tu_all = tail.to_uppercase();
1312 if let Some(gb_at) = find_kw(&tu_all, "GROUP BY") {
1313 let head = tail[..gb_at].trim_end().to_string();
1314 let after = tail[gb_at + "GROUP BY".len()..].trim_start();
1315 let key_end = after.find(|c: char| c == ' ' || c == ',').unwrap_or(after.len());
1316 let group_key = after[..key_end].trim().trim_matches('"').to_string();
1317 let after_key = after[key_end..].trim_start();
1318
1319 if after_key.starts_with(',') {
1323 return Err(format!(
1324 "GROUP BY takes one key here (got {:?} and more) — NQL groups by a \
1325 single field, and grouping by only the first would aggregate over \
1326 rows the query meant to keep apart",
1327 group_key));
1328 }
1329
1330 for c in &project {
1331 let ok = c.src == group_key
1332 || c.src == "count"
1333 || agg_srcs.contains(&c.src);
1334 if !ok {
1335 return Err(format!(
1336 "column {:?} must appear in the GROUP BY clause or be used in an \
1337 aggregate function — a grouped row carries the group key, `count`, \
1338 and the aggregate, nothing else",
1339 c.src));
1340 }
1341 }
1342
1343 tail = format!("{} GROUP BY {}{} {}", head, group_key, agg_clause, after_key)
1354 .split_whitespace()
1355 .collect::<Vec<_>>()
1356 .join(" ");
1357 agg_clause.clear();
1358 }
1359
1360 let tail = sql_literals_to_nql(&tail);
1361 let nql = format!("FROM {}{}{}", coll,
1362 if agg_clause.is_empty() { String::new() } else { agg_clause },
1363 if tail.is_empty() { String::new() } else { format!(" {}", tail) });
1364
1365 Ok(Stmt::Query { nql: nql.trim().to_string(), project })
1366}
1367
1368const SERVER_VERSION: &str = "15.0";
1369
1370pub fn version_string() -> String {
1372 full_version_string()
1373}
1374
1375fn full_version_string() -> String {
1376 format!(
1377 "PostgreSQL {} (NEDB {}) — tamper-evident, append-only, permanent \
1378 history. SELECT + INSERT/UPDATE/DELETE; an UPDATE is a new version, \
1379 so prior values stay readable with AS OF SYSTEM TIME.",
1380 SERVER_VERSION,
1381 env!("CARGO_PKG_VERSION")
1382 )
1383}
1384
1385fn columns_for(rows: &[Value], project: &[Col]) -> Vec<Col> {
1393 if !project.is_empty() {
1394 return project.to_vec();
1395 }
1396 let mut plain: Vec<String> = vec![];
1397 let mut meta: Vec<String> = vec![];
1398 for r in rows {
1399 if let Value::Object(m) = r {
1400 for k in m.keys() {
1401 let target = if k.starts_with('_') { &mut meta } else { &mut plain };
1402 if !target.contains(k) {
1403 target.push(k.clone());
1404 }
1405 }
1406 }
1407 }
1408 plain.sort();
1409 meta.sort();
1410 plain.extend(meta);
1411 plain.into_iter().map(|k| Col::same(&k)).collect()
1412}
1413
1414fn oid_of_value(v: &Value) -> Option<i32> {
1416 match v {
1417 Value::Null => None,
1418 Value::Bool(_) => Some(OID_BOOL),
1419 Value::Number(n) => Some(if n.is_i64() || n.is_u64() { OID_INT8 } else { OID_FLOAT8 }),
1420 Value::String(_) => Some(OID_TEXT),
1421 _ => Some(OID_TEXT),
1423 }
1424}
1425
1426fn unify_oid(a: i32, b: i32) -> i32 {
1433 if a == b {
1434 return a;
1435 }
1436 match (a, b) {
1437 (OID_INT8, OID_FLOAT8) | (OID_FLOAT8, OID_INT8) => OID_FLOAT8,
1438 _ => OID_TEXT,
1439 }
1440}
1441
1442pub fn oid_for_column(rows: &[Value], col: &str) -> i32 {
1454 oid_for(rows, col)
1455}
1456
1457fn oid_for(rows: &[Value], col: &str) -> i32 {
1458 let mut acc: Option<i32> = None;
1459 for r in rows {
1460 if let Some(o) = r.get(col).and_then(oid_of_value) {
1461 acc = Some(match acc {
1462 None => o,
1463 Some(prev) => unify_oid(prev, o),
1464 });
1465 if acc == Some(OID_TEXT) {
1466 break; }
1468 }
1469 }
1470 acc.unwrap_or(OID_TEXT)
1471}
1472
1473fn cell(v: Option<&Value>) -> Option<String> {
1475 match v {
1476 None | Some(Value::Null) => None, Some(Value::String(s)) => Some(s.clone()),
1478 Some(Value::Bool(b)) => Some(if *b { "t".into() } else { "f".into() }),
1479 Some(other) => Some(other.to_string()),
1480 }
1481}
1482
1483fn cell_binary(v: Option<&Value>, oid: i32) -> Result<Option<Vec<u8>>, String> {
1495 let v = match v {
1496 None | Some(Value::Null) => return Ok(None),
1497 Some(v) => v,
1498 };
1499 let as_f64 = |n: &serde_json::Number| n.as_f64()
1500 .ok_or_else(|| "a number too large to send as float8".to_string());
1501 Ok(Some(match (oid, v) {
1502 (OID_BOOL, Value::Bool(b)) => vec![u8::from(*b)],
1503 (OID_INT2, Value::Number(n)) => {
1504 let i = n.as_i64().ok_or("not an integer")?;
1505 i16::try_from(i).map_err(|_| format!("{} does not fit in int2", i))?
1506 .to_be_bytes().to_vec()
1507 }
1508 (OID_INT4, Value::Number(n)) => {
1509 let i = n.as_i64().ok_or("not an integer")?;
1510 i32::try_from(i).map_err(|_| format!("{} does not fit in int4", i))?
1511 .to_be_bytes().to_vec()
1512 }
1513 (OID_INT8, Value::Number(n)) => {
1514 n.as_i64().ok_or("not an integer")?.to_be_bytes().to_vec()
1515 }
1516 (OID_FLOAT4, Value::Number(n)) => (as_f64(n)? as f32).to_be_bytes().to_vec(),
1517 (OID_FLOAT8, Value::Number(n)) => as_f64(n)?.to_be_bytes().to_vec(),
1518 (OID_TEXT | OID_VARCHAR | OID_NAME | OID_UNKNOWN | OID_JSON, _) => {
1520 cell(Some(v)).unwrap_or_default().into_bytes()
1521 }
1522 (OID_JSONB, _) => {
1524 let mut b = vec![1u8];
1525 b.extend_from_slice(cell(Some(v)).unwrap_or_default().as_bytes());
1526 b
1527 }
1528 (oid, val) => {
1529 let kind = match val {
1530 Value::Bool(_) => "a boolean",
1531 Value::Number(_) => "a number",
1532 Value::String(_) => "a string",
1533 Value::Array(_) => "an array",
1534 _ => "an object",
1535 };
1536 return Err(format!(
1537 "cannot send {} in binary format as type OID {} — the field holds \
1538 more than one type across documents, so it cannot be described \
1539 by a single Postgres type. Select it with a text cast, or use a \
1540 text-format client",
1541 kind, oid
1542 ));
1543 }
1544 }))
1545}
1546
1547fn row_description_fmt(cols: &[Col], oids: &[i32], fmts: &[i16]) -> Vec<u8> {
1549 let mut m = Out::msg(b'T');
1550 m.i16(cols.len() as i16);
1551 for (i, c) in cols.iter().enumerate() {
1552 m.cstr(&c.out);
1553 m.i32(0); m.i16((i + 1) as i16); m.i32(oids.get(i).copied().unwrap_or(OID_TEXT));
1556 m.i16(-1); m.i32(-1); m.i16(fmts.get(i).copied().unwrap_or(0));
1559 }
1560 m.finish()
1561}
1562
1563fn row_description(cols: &[Col], oids: &[i32]) -> Vec<u8> {
1564 row_description_fmt(cols, oids, &[])
1565}
1566
1567fn data_row_bytes(vals: &[Option<Vec<u8>>]) -> Vec<u8> {
1568 let mut m = Out::msg(b'D');
1569 m.i16(vals.len() as i16);
1570 for v in vals {
1571 match v {
1572 None => m.i32(-1),
1573 Some(b) => {
1574 m.i32(b.len() as i32);
1575 m.bytes(b);
1576 }
1577 }
1578 }
1579 m.finish()
1580}
1581
1582fn data_row(vals: &[Option<String>]) -> Vec<u8> {
1583 let owned: Vec<Option<Vec<u8>>> =
1584 vals.iter().map(|v| v.as_ref().map(|s| s.as_bytes().to_vec())).collect();
1585 data_row_bytes(&owned)
1586}
1587
1588pub fn encode_rows(rows: &[Value], project: &[Col]) -> Vec<u8> {
1598 let cols = columns_for(rows, project);
1599 let oids: Vec<i32> = cols.iter().map(|c| oid_for(rows, &c.src)).collect();
1600 let mut out = row_description(&cols, &oids);
1601 for r in rows {
1602 let vals: Vec<Option<String>> = cols.iter().map(|c| cell(r.get(&c.src))).collect();
1603 out.extend_from_slice(&data_row(&vals));
1604 }
1605 out
1606}
1607
1608pub fn encode_result(rows: &[Value], project: &[Col]) -> Vec<u8> {
1610 let mut out = encode_rows(rows, project);
1611 out.extend_from_slice(&command_complete(&format!("SELECT {}", rows.len())));
1612 out
1613}
1614
1615const OID_INT2: i32 = 21;
1642const OID_INT4: i32 = 23;
1643const OID_OID: i32 = 26;
1644const OID_FLOAT4: i32 = 700;
1645const OID_VARCHAR: i32 = 1043;
1646const OID_NAME: i32 = 19;
1647const OID_UNKNOWN: i32 = 705;
1648const OID_JSON: i32 = 114;
1649const OID_JSONB: i32 = 3802;
1650
1651fn param_count(sql: &str) -> usize {
1657 let b = sql.as_bytes();
1658 let mut i = 0usize;
1659 let mut in_s = false;
1660 let mut max = 0usize;
1661 while i < b.len() {
1662 let c = b[i];
1663 if in_s {
1664 if c == b'\'' {
1665 in_s = false;
1666 }
1667 i += 1;
1668 continue;
1669 }
1670 if c == b'\'' {
1671 in_s = true;
1672 i += 1;
1673 continue;
1674 }
1675 if c == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1676 let mut j = i + 1;
1677 let mut n = 0usize;
1678 while j < b.len() && b[j].is_ascii_digit() {
1679 n = n * 10 + (b[j] - b'0') as usize;
1680 j += 1;
1681 }
1682 max = max.max(n);
1683 i = j;
1684 continue;
1685 }
1686 i += 1;
1687 }
1688 max
1689}
1690
1691fn infer_field_oid(db: Option<&Arc<Db>>, coll: &str, field: &str) -> i32 {
1700 match field {
1704 "_seq" => return OID_INT8,
1705 "_id" | "_hash" | "_prev" | "_collection" | "_valid_from" | "_valid_to" => return OID_TEXT,
1706 _ => {}
1707 }
1708 if !field.is_empty() && crate::pgcatalog::is_catalog(coll) {
1714 if let Some(rows) = crate::pgcatalog::rows(coll, db) {
1715 return oid_for(&rows, field);
1716 }
1717 }
1718 let db = match db {
1719 Some(db) => db,
1720 None => return OID_TEXT,
1721 };
1722 if coll.is_empty() || field.is_empty() {
1723 return OID_TEXT;
1724 }
1725 let rows = match crate::nql::query(db, &format!("FROM {} LIMIT {}", coll, TYPE_SAMPLE)) {
1726 Ok((rows, _)) => rows,
1727 Err(_) => return OID_TEXT,
1728 };
1729 oid_for(&rows, field)
1733}
1734
1735fn aggregate_oid(src: &str, db: Option<&Arc<Db>>, coll: &str) -> Option<i32> {
1747 if src == "count" {
1748 return Some(OID_INT8);
1749 }
1750 for (prefix, fixed) in [
1751 ("count_", Some(OID_INT8)),
1752 ("avg_", Some(OID_FLOAT8)),
1753 ("sum_", None),
1754 ("min_", None),
1755 ("max_", None),
1756 ] {
1757 if let Some(field) = src.strip_prefix(prefix) {
1758 return Some(match fixed {
1759 Some(oid) => oid,
1760 None => match infer_field_oid(db, coll, field) {
1763 OID_INT8 => OID_INT8,
1764 OID_FLOAT8 => OID_FLOAT8,
1765 other => other,
1768 },
1769 });
1770 }
1771 }
1772 None
1773}
1774
1775const TYPE_SAMPLE: usize = 200;
1781
1782fn stmt_collection(sql: &str) -> String {
1784 let s = normalise(sql);
1785 let up = s.to_uppercase();
1786 let after = if let Some(at) = find_kw(&up, "FROM") {
1787 &s[at + 4..]
1788 } else if let Some(rest) = strip_prefix_ci(&s, "UPDATE") {
1789 return rest
1790 .split_whitespace()
1791 .next()
1792 .unwrap_or("")
1793 .rsplit('.')
1794 .next()
1795 .unwrap_or("")
1796 .trim_matches('"')
1797 .to_string();
1798 } else if let Some(rest) = strip_prefix_ci(&s, "INSERT INTO") {
1799 return rest
1800 .split(|c: char| c.is_whitespace() || c == '(')
1801 .find(|t| !t.is_empty())
1802 .unwrap_or("")
1803 .rsplit('.')
1804 .next()
1805 .unwrap_or("")
1806 .trim_matches('"')
1807 .to_string();
1808 } else {
1809 return String::new();
1810 };
1811 after
1812 .trim()
1813 .split(|c: char| c.is_whitespace())
1814 .find(|t| !t.is_empty())
1815 .unwrap_or("")
1816 .rsplit('.')
1817 .next()
1818 .unwrap_or("")
1819 .trim_matches('"')
1820 .to_string()
1821}
1822
1823fn param_fields(sql: &str, n_params: usize) -> Vec<Option<String>> {
1834 let s = normalise(sql);
1835 let mut out = vec![None; n_params];
1836
1837 let up = s.to_uppercase();
1840 if up.starts_with("INSERT") {
1841 if let (Some(open), Some(vals_at)) = (s.find('('), find_kw(&up, "VALUES")) {
1842 if open < vals_at {
1843 if let Some(close) = s[open..vals_at].rfind(')') {
1844 let cols: Vec<String> = split_top(&s[open + 1..open + close], ',')
1845 .into_iter()
1846 .map(|c| c.trim().trim_matches('"').to_string())
1847 .collect();
1848 let tail = &s[vals_at..];
1850 let mut seen = 0usize;
1851 let b = tail.as_bytes();
1852 let mut i = 0usize;
1853 let mut in_s = false;
1854 while i < b.len() {
1855 if in_s {
1856 if b[i] == b'\'' { in_s = false; }
1857 i += 1;
1858 continue;
1859 }
1860 if b[i] == b'\'' { in_s = true; i += 1; continue; }
1861 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1862 let mut j = i + 1;
1863 let mut num = 0usize;
1864 while j < b.len() && b[j].is_ascii_digit() {
1865 num = num * 10 + (b[j] - b'0') as usize;
1866 j += 1;
1867 }
1868 if num >= 1 && num <= n_params {
1869 if let Some(c) = cols.get(seen % cols.len().max(1)) {
1870 out[num - 1] = Some(c.clone());
1871 }
1872 }
1873 seen += 1;
1874 i = j;
1875 continue;
1876 }
1877 i += 1;
1878 }
1879 return out;
1880 }
1881 }
1882 }
1883 }
1884
1885 let b = s.as_bytes();
1887 let mut i = 0usize;
1888 let mut in_s = false;
1889 while i < b.len() {
1890 if in_s {
1891 if b[i] == b'\'' { in_s = false; }
1892 i += 1;
1893 continue;
1894 }
1895 if b[i] == b'\'' { in_s = true; i += 1; continue; }
1896 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1897 let mut j = i + 1;
1898 let mut num = 0usize;
1899 while j < b.len() && b[j].is_ascii_digit() {
1900 num = num * 10 + (b[j] - b'0') as usize;
1901 j += 1;
1902 }
1903 if num >= 1 && num <= n_params {
1904 let left = &s[..i];
1905 let trimmed = left.trim_end_matches(|c: char| {
1908 c.is_whitespace() || "=<>!+-*/%(,".contains(c)
1909 });
1910 let mut tok = trimmed
1913 .rsplit(|c: char| c.is_whitespace() || c == '(' || c == ',')
1914 .find(|t| !t.is_empty())
1915 .unwrap_or("")
1916 .trim_matches('"');
1917 let mut before = trimmed;
1918 for _ in 0..4 {
1919 let upper_tok = tok.to_uppercase();
1920 if upper_tok.starts_with('$')
1926 || matches!(upper_tok.as_str(),
1927 "LIKE" | "ILIKE" | "IN" | "BETWEEN" | "AND" | "OR" | "NOT" | "IS") {
1928 before = before[..before.len() - tok.len()].trim_end_matches(|c: char| {
1929 c.is_whitespace() || "=<>!(,".contains(c)
1930 });
1931 tok = before
1932 .rsplit(|c: char| c.is_whitespace() || c == '(' || c == ',')
1933 .find(|t| !t.is_empty())
1934 .unwrap_or("")
1935 .trim_matches('"');
1936 } else {
1937 break;
1938 }
1939 }
1940 if !tok.is_empty()
1941 && tok.chars().all(|c| c.is_alphanumeric() || c == '_' || c == '.')
1942 && !tok.chars().next().map(|c| c.is_ascii_digit()).unwrap_or(true)
1943 {
1944 out[num - 1] = Some(tok.rsplit('.').next().unwrap_or(tok).to_string());
1945 }
1946 }
1947 i = j;
1948 continue;
1949 }
1950 i += 1;
1951 }
1952 out
1953}
1954
1955fn clause_param_oids(sql: &str, n_params: usize) -> Vec<Option<i32>> {
1965 let s = normalise(sql);
1966 let mut out = vec![None; n_params];
1967 let b = s.as_bytes();
1968 let mut i = 0usize;
1969 let mut in_s = false;
1970 while i < b.len() {
1971 if in_s {
1972 if b[i] == b'\'' { in_s = false; }
1973 i += 1;
1974 continue;
1975 }
1976 if b[i] == b'\'' { in_s = true; i += 1; continue; }
1977 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1978 let mut j = i + 1;
1979 let mut num = 0usize;
1980 while j < b.len() && b[j].is_ascii_digit() {
1981 num = num * 10 + (b[j] - b'0') as usize;
1982 j += 1;
1983 }
1984 if num >= 1 && num <= n_params {
1985 let left = s[..i].trim_end().to_uppercase();
1986 out[num - 1] = if left.ends_with("VALID AS OF") {
1989 Some(OID_TEXT)
1990 } else if left.ends_with("AS OF SYSTEM TIME")
1991 || left.ends_with("FOR SYSTEM_TIME AS OF")
1992 || left.ends_with("AS OF")
1993 || left.ends_with("LIMIT")
1994 || left.ends_with("OFFSET")
1995 {
1996 Some(OID_INT8)
1997 } else {
1998 None
1999 };
2000 }
2001 i = j;
2002 continue;
2003 }
2004 i += 1;
2005 }
2006 out
2007}
2008
2009fn infer_param_oids(sql: &str, declared: &[i32], db: Option<&Arc<Db>>) -> Vec<i32> {
2015 let n = param_count(sql).max(declared.len());
2016 if n == 0 {
2017 return vec![];
2018 }
2019 let coll = stmt_collection(sql);
2020 let fields = param_fields(sql, n);
2021 let clauses = clause_param_oids(sql, n);
2022 (0..n)
2023 .map(|i| match declared.get(i) {
2024 Some(&oid) if oid != 0 => oid,
2025 _ => match clauses[i] {
2028 Some(oid) => oid,
2029 None => match &fields[i] {
2030 Some(f) => infer_field_oid(db, &coll, f),
2031 None => OID_TEXT,
2032 },
2033 },
2034 })
2035 .collect()
2036}
2037
2038fn decode_param(raw: Option<&[u8]>, oid: i32, format: i16) -> Result<Option<String>, String> {
2044 let bytes = match raw {
2045 None => return Ok(None),
2046 Some(b) => b,
2047 };
2048 let quote = |s: &str| format!("'{}'", s.replace('\'', "''"));
2049
2050 if format == 0 {
2051 let s = String::from_utf8_lossy(bytes).to_string();
2052 return Ok(Some(match oid {
2053 OID_BOOL => {
2054 let t = matches!(s.as_str(), "t" | "true" | "TRUE" | "1" | "yes" | "on");
2055 if t { "TRUE".into() } else { "FALSE".into() }
2056 }
2057 OID_INT2 | OID_INT4 | OID_INT8 | OID_OID | OID_FLOAT4 | OID_FLOAT8 => {
2058 if s.parse::<f64>().is_ok() { s } else { quote(&s) }
2062 }
2063 _ => quote(&s),
2068 }));
2069 }
2070 if format != 1 {
2071 return Err(format!("unsupported parameter format code {}", format));
2072 }
2073
2074 let need = |n: usize| -> Result<(), String> {
2076 if bytes.len() == n {
2077 Ok(())
2078 } else {
2079 Err(format!(
2080 "binary parameter of type OID {} should be {} bytes, got {}",
2081 oid, n, bytes.len()
2082 ))
2083 }
2084 };
2085 Ok(Some(match oid {
2086 OID_BOOL => {
2087 need(1)?;
2088 if bytes[0] != 0 { "TRUE".into() } else { "FALSE".into() }
2089 }
2090 OID_INT2 => {
2091 need(2)?;
2092 i16::from_be_bytes([bytes[0], bytes[1]]).to_string()
2093 }
2094 OID_INT4 => {
2095 need(4)?;
2096 i32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]).to_string()
2097 }
2098 OID_OID => {
2099 need(4)?;
2100 u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]).to_string()
2101 }
2102 OID_INT8 => {
2103 need(8)?;
2104 i64::from_be_bytes(bytes[..8].try_into().unwrap()).to_string()
2105 }
2106 OID_FLOAT4 => {
2107 need(4)?;
2108 let f = f32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]);
2109 fmt_float(f as f64)
2110 }
2111 OID_FLOAT8 => {
2112 need(8)?;
2113 fmt_float(f64::from_be_bytes(bytes[..8].try_into().unwrap()))
2114 }
2115 OID_TEXT | OID_VARCHAR | OID_NAME | OID_UNKNOWN | OID_JSON | 0 => {
2116 quote(&String::from_utf8_lossy(bytes))
2117 }
2118 OID_JSONB => {
2119 let body = if bytes.first() == Some(&1) { &bytes[1..] } else { bytes };
2121 quote(&String::from_utf8_lossy(body))
2122 }
2123 other => {
2124 return Err(format!(
2125 "parameter type OID {} is not supported in binary format — \
2126 the supported set is bool, int2/int4/int8, float4/float8, \
2127 text/varchar/json/jsonb. Send it as text, or cast it in the \
2128 statement",
2129 other
2130 ))
2131 }
2132 }))
2133}
2134
2135fn fmt_float(f: f64) -> String {
2137 if f.is_nan() {
2138 "'NaN'".into()
2139 } else if f.is_infinite() {
2140 if f > 0.0 { "'Infinity'".into() } else { "'-Infinity'".into() }
2141 } else if f.fract() == 0.0 && f.abs() < 1e15 {
2142 format!("{:.0}", f)
2143 } else {
2144 f.to_string()
2145 }
2146}
2147
2148fn substitute_params(sql: &str, params: &[Option<String>]) -> Result<String, String> {
2156 let b = sql.as_bytes();
2157 let mut out = String::with_capacity(sql.len() + 16);
2158 let mut i = 0usize;
2159 let mut in_s = false;
2160 while i < b.len() {
2161 let c = b[i];
2162 if in_s {
2163 out.push(c as char);
2164 if c == b'\'' { in_s = false; }
2165 i += 1;
2166 continue;
2167 }
2168 if c == b'\'' {
2169 in_s = true;
2170 out.push('\'');
2171 i += 1;
2172 continue;
2173 }
2174 if c == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
2175 let mut j = i + 1;
2176 let mut n = 0usize;
2177 while j < b.len() && b[j].is_ascii_digit() {
2178 n = n * 10 + (b[j] - b'0') as usize;
2179 j += 1;
2180 }
2181 match params.get(n.wrapping_sub(1)) {
2182 Some(Some(lit)) => out.push_str(lit),
2183 Some(None) => out.push_str("NULL"),
2184 None => {
2185 return Err(format!(
2186 "bind message supplies {} parameter(s) but the statement uses ${}",
2187 params.len(), n
2188 ))
2189 }
2190 }
2191 i = j;
2192 continue;
2193 }
2194 out.push(c as char);
2195 i += 1;
2196 }
2197 Ok(out)
2198}
2199
2200struct Prepared {
2202 sql: String,
2203 param_oids: Vec<i32>,
2206 out_shape: Option<Option<(Vec<Col>, Vec<i32>)>>,
2214}
2215
2216fn prepared_shape<'a>(
2218 p: &'a mut Prepared,
2219 db: Option<&Arc<Db>>,
2220) -> &'a Option<(Vec<Col>, Vec<i32>)> {
2221 if p.out_shape.is_none() {
2222 p.out_shape = Some(describe_shape(&p.sql, db, p.param_oids.len()));
2223 }
2224 p.out_shape.as_ref().expect("just filled")
2225}
2226
2227struct Portal {
2229 sql: String,
2230 result: Option<PortalResult>,
2236 frozen: Option<Vec<Col>>,
2246 formats: Vec<i16>,
2248 declared: Option<(Vec<Col>, Vec<i32>)>,
2255}
2256
2257impl Portal {
2258 fn format_of(&self, i: usize) -> i16 {
2261 match self.formats.len() {
2262 0 => 0,
2263 1 => self.formats[0],
2264 _ => self.formats.get(i).copied().unwrap_or(0),
2265 }
2266 }
2267 fn shape(&self, r: &PortalResult) -> (Vec<Col>, Vec<i32>) {
2269 match &self.declared {
2270 Some((cols, oids)) if self.formats.iter().any(|f| *f == 1) => {
2271 (cols.clone(), oids.clone())
2272 }
2273 _ => {
2274 let cols = columns_for(&r.rows, &r.project);
2275 let oids = cols.iter().map(|c| oid_for(&r.rows, &c.src)).collect();
2276 (cols, oids)
2277 }
2278 }
2279 }
2280}
2281
2282struct PortalResult {
2283 rows: Vec<Value>,
2284 project: Vec<Col>,
2285 has_rows: bool,
2286 tag: String,
2287 tag_counts_rows: bool,
2288 sent: usize,
2290}
2291
2292fn parse_complete() -> Vec<u8> { Out::msg(b'1').finish() }
2293fn bind_complete() -> Vec<u8> { Out::msg(b'2').finish() }
2294fn close_complete() -> Vec<u8> { Out::msg(b'3').finish() }
2295fn no_data() -> Vec<u8> { Out::msg(b'n').finish() }
2296fn portal_suspended() -> Vec<u8> { Out::msg(b's').finish() }
2297
2298fn parameter_description(oids: &[i32]) -> Vec<u8> {
2299 let mut m = Out::msg(b't');
2300 m.i16(oids.len() as i16);
2301 for o in oids {
2302 m.i32(*o);
2303 }
2304 m.finish()
2305}
2306
2307fn take_cstr(body: &[u8], at: &mut usize) -> String {
2309 let start = *at;
2310 while *at < body.len() && body[*at] != 0 {
2311 *at += 1;
2312 }
2313 let s = String::from_utf8_lossy(&body[start..*at]).to_string();
2314 if *at < body.len() {
2315 *at += 1; }
2317 s
2318}
2319
2320fn take_i16(body: &[u8], at: &mut usize) -> Result<i16, String> {
2321 if *at + 2 > body.len() {
2322 return Err("truncated message".into());
2323 }
2324 let v = i16::from_be_bytes([body[*at], body[*at + 1]]);
2325 *at += 2;
2326 Ok(v)
2327}
2328
2329fn take_i32(body: &[u8], at: &mut usize) -> Result<i32, String> {
2330 if *at + 4 > body.len() {
2331 return Err("truncated message".into());
2332 }
2333 let v = i32::from_be_bytes([body[*at], body[*at + 1], body[*at + 2], body[*at + 3]]);
2334 *at += 4;
2335 Ok(v)
2336}
2337
2338fn sample_columns(db: Option<&Arc<Db>>, coll: &str) -> Vec<Col> {
2344 let db = match db {
2345 Some(db) => db,
2346 None => return vec![],
2347 };
2348 let rows = match crate::nql::query(db, &format!("FROM {} LIMIT 25", coll)) {
2349 Ok((rows, _)) => rows,
2350 Err(_) => return vec![],
2351 };
2352 let mut names: Vec<String> = vec![];
2353 for r in &rows {
2354 if let Value::Object(m) = r {
2355 for k in m.keys() {
2356 if !names.iter().any(|n| n == k) {
2357 names.push(k.clone());
2358 }
2359 }
2360 }
2361 }
2362 names.sort();
2363 names.iter().map(|n| Col::same(n)).collect()
2364}
2365
2366fn describe_shape(
2374 sql: &str,
2375 db: Option<&Arc<Db>>,
2376 n_params: usize,
2377) -> Option<(Vec<Col>, Vec<i32>)> {
2378 let probe = probe_sql(sql, n_params);
2379
2380 if sql_engine_owns(&probe) {
2391 if let Ok(Some((done, _))) = try_catalog_select(&probe, db) {
2392 if done.project.is_empty() {
2393 return None;
2394 }
2395 let oids = done
2396 .project
2397 .iter()
2398 .map(|c| oid_for(&done.rows, &c.src))
2399 .collect();
2400 return Some((done.project, oids));
2401 }
2402 }
2403
2404 let stmt = translate(&probe).ok()?;
2405 let coll = stmt_collection(sql);
2406
2407 let cols = match stmt {
2408 Stmt::Ok(_) => return None,
2409 Stmt::Canned { cols, .. } => cols.iter().map(|c| Col::same(c)).collect(),
2410 Stmt::Query { project, .. } => {
2411 if project.is_empty() { sample_columns(db, &coll) } else { project }
2412 }
2413 Stmt::Insert { returning, .. } | Stmt::Update { returning, .. } | Stmt::Delete { returning, .. } => {
2414 if !wants_returning(sql) {
2415 return None;
2416 }
2417 if returning.is_empty() { sample_columns(db, &coll) } else { returning }
2418 }
2419 };
2420 if cols.is_empty() {
2421 return None;
2425 }
2426 let oids = cols
2427 .iter()
2428 .map(|c| {
2429 aggregate_oid(&c.src, db, &coll)
2430 .unwrap_or_else(|| infer_field_oid(db, &coll, &c.src))
2431 })
2432 .collect();
2433 Some((cols, oids))
2434}
2435
2436fn probe_sql(sql: &str, n_params: usize) -> String {
2444 let stub: Vec<Option<String>> = vec![Some("0".to_string()); n_params];
2445 substitute_params(sql, &stub).unwrap_or_else(|_| sql.to_string())
2446}
2447
2448fn ensure_executed(
2450 portal: &mut Portal,
2451 db_name: &str,
2452 db: Option<&Arc<Db>>,
2453 read_only: bool,
2454) -> Result<(), Vec<u8>> {
2455 if portal.result.is_some() {
2456 return Ok(());
2457 }
2458 let ex = execute_stmt(&portal.sql, db_name, db, read_only)?;
2459 let project = if let Some(f) = &portal.frozen {
2462 f.clone()
2463 } else {
2464 let p = if ex.project.is_empty() {
2465 columns_for(&ex.rows, &[])
2466 } else {
2467 ex.project.clone()
2468 };
2469 portal.frozen = Some(p.clone());
2470 p
2471 };
2472 portal.result = Some(PortalResult {
2473 rows: ex.rows,
2474 project,
2475 has_rows: ex.has_rows,
2476 tag: ex.tag,
2477 tag_counts_rows: ex.tag_counts_rows,
2478 sent: 0,
2479 });
2480 Ok(())
2481}
2482
2483async fn read_exact(sock: &mut TcpStream, n: usize) -> std::io::Result<Vec<u8>> {
2486 let mut buf = vec![0u8; n];
2487 sock.read_exact(&mut buf).await?;
2488 Ok(buf)
2489}
2490
2491async fn read_i32(sock: &mut TcpStream) -> std::io::Result<i32> {
2492 let b = read_exact(sock, 4).await?;
2493 Ok(i32::from_be_bytes([b[0], b[1], b[2], b[3]]))
2494}
2495
2496fn parse_startup_params(body: &[u8]) -> HashMap<String, String> {
2497 let mut out = HashMap::new();
2498 let mut parts = body.split(|b| *b == 0).map(|s| String::from_utf8_lossy(s).to_string());
2499 while let (Some(k), Some(v)) = (parts.next(), parts.next()) {
2500 if k.is_empty() {
2501 break;
2502 }
2503 out.insert(k, v);
2504 }
2505 out
2506}
2507
2508async fn handle(mut sock: TcpStream, resolver: Arc<dyn DbResolver>, read_only: bool) -> std::io::Result<()> {
2510 let params = loop {
2512 let len = read_i32(&mut sock).await?;
2513 if len < 8 || len > 1 << 20 {
2514 return Ok(()); }
2516 let code = read_i32(&mut sock).await?;
2517 let body = read_exact(&mut sock, (len - 8) as usize).await?;
2518 match code {
2519 SSL_REQUEST | GSS_REQUEST => {
2520 sock.write_all(b"N").await?;
2522 continue;
2523 }
2524 CANCEL_REQUEST => return Ok(()), PROTO_V3 => break parse_startup_params(&body),
2526 other => {
2527 let major = other >> 16;
2528 sock.write_all(&err_msg(
2529 "0A000",
2530 &format!("unsupported frontend protocol {}.{} — this endpoint speaks 3.0",
2531 major, other & 0xffff),
2532 )).await?;
2533 return Ok(());
2534 }
2535 }
2536 };
2537
2538 let db_name = params.get("database").cloned().unwrap_or_default();
2539
2540 let resolved: Option<Arc<Db>> = {
2546 let r = Arc::clone(&resolver);
2547 let name = db_name.clone();
2548 tokio::task::spawn_blocking(move || r.resolve(&name))
2549 .await
2550 .unwrap_or(None)
2551 };
2552
2553 if let Some(expected) = resolver.token() {
2555 let mut m = Out::msg(b'R');
2557 m.i32(3);
2558 sock.write_all(&m.finish()).await?;
2559
2560 let tag = read_exact(&mut sock, 1).await?;
2561 if tag[0] != b'p' {
2562 sock.write_all(&err_msg("28000", "expected a password message")).await?;
2563 return Ok(());
2564 }
2565 let len = read_i32(&mut sock).await?;
2566 if len < 4 || len > 1 << 16 {
2567 return Ok(());
2568 }
2569 let body = read_exact(&mut sock, (len - 4) as usize).await?;
2570 let supplied = String::from_utf8_lossy(&body).trim_end_matches('\0').to_string();
2571 let ok = supplied.len() == expected.len()
2573 && supplied.bytes().zip(expected.bytes()).fold(0u8, |a, (x, y)| a | (x ^ y)) == 0;
2574 if !ok {
2575 sock.write_all(&err_msg("28P01", "password authentication failed")).await?;
2576 return Ok(());
2577 }
2578 }
2579
2580 let mut m = Out::msg(b'R');
2581 m.i32(0); sock.write_all(&m.finish()).await?;
2583
2584 for (k, v) in [
2585 ("server_version", SERVER_VERSION),
2586 ("server_encoding", "UTF8"),
2587 ("client_encoding", "UTF8"),
2588 ("DateStyle", "ISO, MDY"),
2589 ("integer_datetimes", "on"),
2590 ("standard_conforming_strings", "on"),
2591 ("application_name", "nedbd"),
2592 ] {
2593 let mut p = Out::msg(b'S');
2594 p.cstr(k);
2595 p.cstr(v);
2596 sock.write_all(&p.finish()).await?;
2597 }
2598 let mut k = Out::msg(b'K');
2599 k.i32(std::process::id() as i32);
2600 k.i32(0);
2601 sock.write_all(&k.finish()).await?;
2602 sock.write_all(&ready()).await?;
2603
2604 let mut prepared: HashMap<String, Prepared> = HashMap::new();
2610 let mut portals: HashMap<String, Portal> = HashMap::new();
2611 let mut failed = false;
2615
2616 loop {
2617 let mut tag = [0u8; 1];
2618 if sock.read_exact(&mut tag).await.is_err() {
2619 return Ok(()); }
2621 let len = read_i32(&mut sock).await?;
2622 if len < 4 || len > 64 << 20 {
2623 return Ok(());
2624 }
2625 let body = read_exact(&mut sock, (len - 4) as usize).await?;
2626
2627 if failed && tag[0] != b'S' && tag[0] != b'X' {
2629 continue;
2630 }
2631
2632 match tag[0] {
2633 b'X' => return Ok(()), b'Q' => {
2636 let sql = String::from_utf8_lossy(&body).trim_end_matches('\0').to_string();
2637 let out = run_simple_query(&sql, &db_name, resolved.as_ref(), read_only);
2638 sock.write_all(&out).await?;
2639 sock.write_all(&ready()).await?;
2640 portals.remove("");
2642 }
2643
2644 b'P' => {
2646 let mut at = 0usize;
2647 let name = take_cstr(&body, &mut at);
2648 let sql = take_cstr(&body, &mut at);
2649 let n = take_i16(&body, &mut at).unwrap_or(0).max(0) as usize;
2650 let mut declared = Vec::with_capacity(n);
2651 let mut bad = false;
2652 for _ in 0..n {
2653 match take_i32(&body, &mut at) {
2654 Ok(o) => declared.push(o),
2655 Err(_) => { bad = true; break; }
2656 }
2657 }
2658 if bad {
2659 sock.write_all(&err_msg("08P01", "malformed Parse message")).await?;
2660 failed = true;
2661 continue;
2662 }
2663 let probe = probe_sql(&sql, param_count(&sql));
2673 if !sql_engine_owns(&probe) {
2674 if let Err(why) = translate(&probe) {
2675 sock.write_all(&err_msg("0A000", &why)).await?;
2676 failed = true;
2677 continue;
2678 }
2679 }
2680 let param_oids = infer_param_oids(&sql, &declared, resolved.as_ref());
2681 prepared.insert(name, Prepared { sql, param_oids, out_shape: None });
2682 sock.write_all(&parse_complete()).await?;
2683 }
2684
2685 b'B' => {
2687 let mut at = 0usize;
2688 let portal_name = take_cstr(&body, &mut at);
2689 let stmt_name = take_cstr(&body, &mut at);
2690 if !prepared.contains_key(&stmt_name) {
2691 sock.write_all(&err_msg("26000", &format!(
2692 "prepared statement {:?} does not exist", stmt_name))).await?;
2693 failed = true;
2694 continue;
2695 }
2696 let p = &prepared[&stmt_name];
2697 let mut want_formats: Vec<i16> = vec![];
2698 let res: Result<String, String> = (|| {
2699 let nfmt = take_i16(&body, &mut at)? .max(0) as usize;
2700 let mut fmts = Vec::with_capacity(nfmt);
2701 for _ in 0..nfmt {
2702 fmts.push(take_i16(&body, &mut at)?);
2703 }
2704 let nparam = take_i16(&body, &mut at)?.max(0) as usize;
2705 let mut vals: Vec<Option<String>> = Vec::with_capacity(nparam);
2706 for i in 0..nparam {
2707 let l = take_i32(&body, &mut at)?;
2708 let raw: Option<Vec<u8>> = if l < 0 {
2709 None
2710 } else {
2711 let l = l as usize;
2712 if at + l > body.len() {
2713 return Err("truncated Bind parameter".into());
2714 }
2715 let v = body[at..at + l].to_vec();
2716 at += l;
2717 Some(v)
2718 };
2719 let f = match fmts.len() {
2722 0 => 0,
2723 1 => fmts[0],
2724 _ => *fmts.get(i).unwrap_or(&0),
2725 };
2726 let oid = *p.param_oids.get(i).unwrap_or(&OID_TEXT);
2727 vals.push(decode_param(raw.as_deref(), oid, f)?);
2728 }
2729 let nres = take_i16(&body, &mut at)?.max(0) as usize;
2734 for _ in 0..nres {
2735 let f = take_i16(&body, &mut at)?;
2736 if f != 0 && f != 1 {
2737 return Err(format!("unknown result format code {}", f));
2738 }
2739 want_formats.push(f);
2740 }
2741 substitute_params(&p.sql, &vals)
2742 })();
2743 match res {
2744 Ok(sql) => {
2745 let declared = if want_formats.iter().any(|f| *f == 1) {
2748 let p = prepared.get_mut(&stmt_name).expect("checked above");
2749 prepared_shape(p, resolved.as_ref()).clone()
2750 } else {
2751 None
2752 };
2753 portals.insert(portal_name, Portal {
2754 sql, result: None, frozen: None,
2755 formats: want_formats, declared,
2756 });
2757 sock.write_all(&bind_complete()).await?;
2758 }
2759 Err(why) => {
2760 sock.write_all(&err_msg("08P01", &why)).await?;
2761 failed = true;
2762 }
2763 }
2764 }
2765
2766 b'D' => {
2768 let kind = body.first().copied().unwrap_or(b'S');
2769 let mut at = 1usize;
2770 let name = take_cstr(&body, &mut at);
2771 if kind == b'S' {
2772 if !prepared.contains_key(&name) {
2773 sock.write_all(&err_msg("26000", &format!(
2774 "prepared statement {:?} does not exist", name))).await?;
2775 failed = true;
2776 continue;
2777 }
2778 let p = prepared.get_mut(&name).expect("checked above");
2779 let oids = p.param_oids.clone();
2780 sock.write_all(¶meter_description(&oids)).await?;
2783 let out = match prepared_shape(p, resolved.as_ref()) {
2787 Some((cols, col_oids)) => row_description(cols, col_oids),
2788 None => no_data(),
2789 };
2790 sock.write_all(&out).await?;
2791 } else {
2792 let portal = match portals.get_mut(&name) {
2793 Some(p) => p,
2794 None => {
2795 sock.write_all(&err_msg("34000", &format!(
2796 "portal {:?} does not exist", name))).await?;
2797 failed = true;
2798 continue;
2799 }
2800 };
2801 match ensure_executed(portal, &db_name, resolved.as_ref(), read_only) {
2806 Err(encoded) => {
2807 sock.write_all(&encoded).await?;
2808 failed = true;
2809 }
2810 Ok(()) => {
2811 let r = portal.result.as_ref().expect("just executed");
2812 if !r.has_rows {
2813 sock.write_all(&no_data()).await?;
2814 } else {
2815 let (cols, oids) = portal.shape(r);
2816 let fmts: Vec<i16> =
2817 (0..cols.len()).map(|i| portal.format_of(i)).collect();
2818 sock.write_all(&row_description_fmt(&cols, &oids, &fmts)).await?;
2819 }
2820 }
2821 }
2822 }
2823 }
2824
2825 b'E' => {
2827 let mut at = 0usize;
2828 let name = take_cstr(&body, &mut at);
2829 let max_rows = take_i32(&body, &mut at).unwrap_or(0);
2830 let portal = match portals.get_mut(&name) {
2831 Some(p) => p,
2832 None => {
2833 sock.write_all(&err_msg("34000", &format!(
2834 "portal {:?} does not exist", name))).await?;
2835 failed = true;
2836 continue;
2837 }
2838 };
2839 if let Err(encoded) = ensure_executed(portal, &db_name, resolved.as_ref(), read_only) {
2840 sock.write_all(&encoded).await?;
2841 failed = true;
2842 continue;
2843 }
2844 let r = portal.result.as_ref().expect("just executed");
2845 if !r.has_rows {
2846 let tag = r.tag.clone();
2847 sock.write_all(&command_complete(&tag)).await?;
2848 continue;
2849 }
2850 let (cols, oids) = portal.shape(r);
2851 let limit = if max_rows > 0 {
2852 (r.sent + max_rows as usize).min(r.rows.len())
2853 } else {
2854 r.rows.len()
2855 };
2856 let mut encoded: Vec<Vec<u8>> = Vec::with_capacity(limit - r.sent);
2861 let mut fail: Option<String> = None;
2862 for row in &r.rows[r.sent..limit] {
2863 let mut vals: Vec<Option<Vec<u8>>> = Vec::with_capacity(cols.len());
2864 for (i, c) in cols.iter().enumerate() {
2865 let v = row.get(&c.src);
2866 let got = if portal.format_of(i) == 1 {
2867 cell_binary(v, oids.get(i).copied().unwrap_or(OID_TEXT))
2868 .map_err(|e| format!("column {:?}: {}", c.out, e))
2869 } else {
2870 Ok(cell(v).map(|s| s.into_bytes()))
2871 };
2872 match got {
2873 Ok(b) => vals.push(b),
2874 Err(e) => { fail = Some(e); break; }
2875 }
2876 }
2877 if fail.is_some() {
2878 break;
2879 }
2880 encoded.push(data_row_bytes(&vals));
2881 }
2882 if let Some(why) = fail {
2883 sock.write_all(&err_msg("22P03", &why)).await?;
2884 failed = true;
2885 continue;
2886 }
2887 let mut out = vec![];
2888 for e in &encoded {
2889 out.extend_from_slice(e);
2890 }
2891 let r = portal.result.as_mut().expect("just executed");
2892 r.sent = limit;
2893 if max_rows > 0 && r.sent < r.rows.len() {
2897 out.extend_from_slice(&portal_suspended());
2898 } else {
2899 let tag = if r.tag_counts_rows {
2900 format!("{} {}", r.tag, r.sent)
2901 } else {
2902 r.tag.clone()
2903 };
2904 out.extend_from_slice(&command_complete(&tag));
2905 }
2906 sock.write_all(&out).await?;
2907 }
2908
2909 b'C' => {
2911 let kind = body.first().copied().unwrap_or(b'S');
2912 let mut at = 1usize;
2913 let name = take_cstr(&body, &mut at);
2914 if kind == b'S' {
2915 prepared.remove(&name);
2916 } else {
2917 portals.remove(&name);
2918 }
2919 sock.write_all(&close_complete()).await?;
2922 }
2923
2924 b'H' => {}
2928
2929 b'S' => {
2930 failed = false;
2931 sock.write_all(&ready()).await?;
2932 }
2933
2934 other => {
2935 sock.write_all(&err_msg(
2936 "08P01",
2937 &format!("unexpected frontend message {:?}", other as char),
2938 )).await?;
2939 failed = true;
2940 }
2941 }
2942 }
2943}
2944
2945const READ_ONLY_MSG: &str =
2946 "this endpoint is running read-only (NEDBD_PG_READ_ONLY=1). Writes are \
2947 implemented but disabled on this server — unset the flag to allow them.";
2948
2949fn no_db(db_name: &str) -> Vec<u8> {
2950 err_msg("3D000", &format!(
2951 "database {:?} is not open on this server — create it first \
2952 (POST /v1/databases), or connect with -d <name>", db_name))
2953}
2954
2955fn catalog_name(n: &str) -> String {
2959 let joined: Vec<&str> = n.split('.').collect();
2960 if joined.len() >= 2 && joined[joined.len() - 2] == "information_schema" {
2961 format!("information_schema.{}", joined[joined.len() - 1])
2962 } else {
2963 joined[joined.len() - 1].to_string()
2964 }
2965}
2966
2967fn sql_engine_owns(sql: &str) -> bool {
2990 let Ok(sel) = crate::sqlselect::parse(sql) else { return false };
2991 let touched = sel.base_relations();
2992 if touched.is_empty() {
2993 return translate(sql).is_err();
2994 }
2995 touched.iter().any(|t| crate::pgcatalog::is_catalog(&catalog_name(t)))
2996}
2997
2998fn try_catalog_select(
3011 sql: &str,
3012 db: Option<&Arc<Db>>,
3013) -> Result<Option<(Executed, crate::sqlplan::Plan)>, Vec<u8>> {
3014 let sel = match crate::sqlselect::parse(sql) {
3015 Ok(sel) => sel,
3016 Err(why) => {
3017 if mentions_catalog(sql) {
3026 return Err(err_msg("0A000", &format!(
3027 "this catalogue query uses SQL this endpoint does not \
3028 implement: {}", why)));
3029 }
3030 return Ok(None);
3031 }
3032 };
3033
3034 if !sql_engine_owns(sql) {
3040 return Ok(None);
3041 }
3042
3043 let resolve = |name: &str| -> anyhow::Result<Option<Box<dyn crate::sqlselect::Relation>>> {
3044 let cname = catalog_name(name);
3045 if let Some(rows) = crate::pgcatalog::rows(&cname, db) {
3046 return Ok(Some(crate::sqlselect::from_vec(rows)));
3050 }
3051 match db {
3060 Some(db) => match crate::nql::query(db, &format!("FROM {}", cname)) {
3061 Ok((rows, _)) => Ok(Some(crate::sqlselect::from_vec(rows))),
3062 Err(_) => Ok(None),
3063 },
3064 None => Ok(None),
3065 }
3066 };
3067
3068 let (cols, rows, plan) = crate::sqlselect::execute_explain(
3069 &sel,
3070 &resolve,
3071 crate::sqljoin::JoinExec::Auto,
3072 )
3073 .map_err(|e| err_msg("42601", &e.to_string()))?;
3074
3075 Ok(Some((
3076 Executed {
3077 rows,
3078 project: cols
3082 .iter()
3083 .map(|c| Col::renamed(&c.key, &c.name))
3084 .collect(),
3085 has_rows: true,
3086 tag: "SELECT".into(),
3087 tag_counts_rows: true,
3088 },
3089 plan,
3090 )))
3091}
3092
3093fn strip_explain(sql: &str) -> Option<&str> {
3100 let t = sql.trim().trim_end_matches(';').trim();
3101 let mut rest = t.strip_prefix("EXPLAIN").or_else(|| t.strip_prefix("explain"))?;
3102 if !rest.starts_with(char::is_whitespace) {
3104 return None;
3105 }
3106 rest = rest.trim_start();
3107 loop {
3108 let low = rest.to_lowercase();
3109 if let Some(r) = low.strip_prefix("analyze").or_else(|| low.strip_prefix("analyse")) {
3110 if r.starts_with(char::is_whitespace) || r.is_empty() {
3111 rest = rest[rest.len() - r.len()..].trim_start();
3112 continue;
3113 }
3114 }
3115 if let Some(r) = low.strip_prefix("verbose") {
3116 if r.starts_with(char::is_whitespace) || r.is_empty() {
3117 rest = rest[rest.len() - r.len()..].trim_start();
3118 continue;
3119 }
3120 }
3121 break;
3122 }
3123 Some(rest)
3124}
3125
3126fn plan_result(lines: Vec<String>) -> Executed {
3129 Executed {
3130 rows: lines
3131 .into_iter()
3132 .map(|l| serde_json::json!({ "QUERY PLAN": l }))
3133 .collect(),
3134 project: vec![Col::same("QUERY PLAN")],
3135 has_rows: true,
3136 tag: "EXPLAIN".into(),
3137 tag_counts_rows: false,
3138 }
3139}
3140
3141fn mentions_catalog(sql: &str) -> bool {
3148 let low = sql.to_lowercase();
3149 low.contains("pg_catalog.")
3150 || low.contains("information_schema.")
3151 || low.contains("from pg_")
3152 || low.contains("join pg_")
3153}
3154
3155fn catalog_target(nql: &str) -> Option<String> {
3160 let coll = crate::nql::parse(nql).ok()?.coll;
3161 if crate::pgcatalog::is_catalog(&coll) {
3162 Some(coll)
3163 } else {
3164 None
3165 }
3166}
3167
3168fn wants_returning(sql: &str) -> bool {
3172 find_kw(&sql.to_uppercase(), "RETURNING").is_some()
3173}
3174
3175fn next_row_id() -> String {
3177 use std::sync::atomic::{AtomicU64, Ordering};
3178 static N: AtomicU64 = AtomicU64::new(0);
3179 let n = N.fetch_add(1, Ordering::Relaxed);
3180 let ts = std::time::SystemTime::now()
3181 .duration_since(std::time::UNIX_EPOCH)
3182 .map(|d| d.as_micros())
3183 .unwrap_or(0);
3184 format!("r{}{}", ts, n)
3185}
3186
3187pub struct Executed {
3195 pub rows: Vec<Value>,
3197 pub project: Vec<Col>,
3199 pub has_rows: bool,
3203 pub tag: String,
3206 pub tag_counts_rows: bool,
3208}
3209
3210impl Executed {
3211 fn nothing(tag: &str) -> Self {
3212 Executed { rows: vec![], project: vec![], has_rows: false, tag: tag.to_string(), tag_counts_rows: false }
3213 }
3214 fn tag_for(&self, sent: usize) -> String {
3216 if self.tag_counts_rows { format!("{} {}", self.tag, sent) } else { self.tag.clone() }
3217 }
3218}
3219
3220fn execute_stmt(
3226 stmt_sql: &str,
3227 db_name: &str,
3228 db: Option<&Arc<Db>>,
3229 read_only: bool,
3230) -> Result<Executed, Vec<u8>> {
3231 if let Some(inner) = strip_explain(stmt_sql) {
3240 if let Some((_, plan)) = try_catalog_select(inner, db)? {
3241 return Ok(plan_result(plan.render()));
3242 }
3243 let mut lines = vec![];
3244 match translate(inner) {
3245 Ok(_) => {
3246 lines.push(
3247 "NQL path — this statement is translated to NQL and \
3248 executed by the storage engine, not by the SQL evaluator."
3249 .to_string(),
3250 );
3251 lines.push(
3252 "No plan is reported, because the SQL evaluator is not \
3253 what runs it. Reporting one would describe a pipeline \
3254 that never executed."
3255 .to_string(),
3256 );
3257 lines.push(
3258 "The SQL evaluator (joins, CASE, scalar functions, a \
3259 hash-join planner) currently serves catalogue queries."
3260 .to_string(),
3261 );
3262 }
3263 Err(why) => lines.push(format!("cannot be executed: {why}")),
3264 }
3265 return Ok(plan_result(lines));
3266 }
3267
3268 if let Some((done, _plan)) = try_catalog_select(stmt_sql, db)? {
3269 return Ok(done);
3270 }
3271
3272 let stmt = translate(stmt_sql).map_err(|why| err_msg("0A000", &why))?;
3273
3274 macro_rules! need_db {
3277 () => {
3278 match db {
3279 Some(db) => db,
3280 None => return Err(no_db(db_name)),
3281 }
3282 };
3283 }
3284 macro_rules! need_write {
3285 () => {
3286 if read_only {
3287 return Err(err_msg("25006", READ_ONLY_MSG));
3288 }
3289 };
3290 }
3291
3292 match stmt {
3293 Stmt::Ok(tag) => Ok(Executed::nothing(if tag.is_empty() { "SELECT 0" } else { tag })),
3294
3295 Stmt::Canned { cols, row } => {
3296 let mut obj = serde_json::Map::new();
3299 for (c, v) in cols.iter().zip(row.iter()) {
3300 obj.insert(c.clone(), Value::String(v.clone()));
3301 }
3302 Ok(Executed {
3303 rows: vec![Value::Object(obj)],
3304 project: cols.iter().map(|c| Col::same(c)).collect(),
3305 has_rows: true,
3306 tag: "SELECT".into(),
3307 tag_counts_rows: true,
3308 })
3309 }
3310
3311 Stmt::Query { nql, project } => {
3312 if let Some(coll) = catalog_target(&nql) {
3322 let rows = crate::pgcatalog::rows(&coll, db)
3323 .expect("catalog_target only returns names pgcatalog serves");
3324 let rows = crate::nql::query_rows(rows, &nql)
3325 .map_err(|e| err_msg("42601", &e.to_string()))?;
3326 return Ok(Executed {
3327 rows, project, has_rows: true,
3328 tag: "SELECT".into(), tag_counts_rows: true,
3329 });
3330 }
3331 let db = need_db!();
3332 let (rows, _) = crate::nql::query(db, &nql).map_err(|e| {
3333 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3334 })?;
3335 Ok(Executed { rows, project, has_rows: true, tag: "SELECT".into(), tag_counts_rows: true })
3336 }
3337
3338 Stmt::Insert { coll, rows, returning } => {
3339 let db = need_db!();
3340 need_write!();
3341 let mut written: Vec<Value> = vec![];
3342 for (i, r) in rows.iter().enumerate() {
3343 let id = match &r.id {
3347 Some(id) => id.clone(),
3348 None => format!("{}-{}", next_row_id(), i),
3349 };
3350 let node = db
3351 .put(&coll, &id, Value::Object(r.doc.clone()),
3352 r.caused_by.clone(), r.valid_from.clone(), r.valid_to.clone())
3353 .map_err(|e| err_msg("XX000", &format!("INSERT failed: {}", e)))?;
3354 written.push(crate::nql::node_to_json(&node));
3355 }
3356 let n = written.len();
3357 let has_rows = wants_returning(stmt_sql);
3358 Ok(Executed {
3359 rows: if has_rows { written } else { vec![] },
3360 project: returning,
3361 has_rows,
3362 tag: format!("INSERT 0 {}", n),
3364 tag_counts_rows: false,
3365 })
3366 }
3367
3368 Stmt::Update { coll, set, nql, returning } => {
3369 let db = need_db!();
3370 need_write!();
3371 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
3374 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3375 })?;
3376 let mut written: Vec<Value> = vec![];
3377 for row in &matched {
3378 let id = match row.get("_id").and_then(|v| v.as_str()) {
3379 Some(id) => id.to_string(),
3380 None => continue,
3381 };
3382 let mut doc = match db.get(&coll, &id) {
3386 Some(n) => match n.data {
3387 Value::Object(m) => m,
3388 _ => serde_json::Map::new(),
3389 },
3390 None => continue,
3391 };
3392 for (k, v) in &set {
3393 doc.insert(k.clone(), v.clone());
3394 }
3395 let node = db
3398 .put(&coll, &id, Value::Object(doc), vec![], None, None)
3399 .map_err(|e| err_msg("XX000", &format!("UPDATE failed: {}", e)))?;
3400 written.push(crate::nql::node_to_json(&node));
3401 }
3402 let n = written.len();
3403 let has_rows = wants_returning(stmt_sql);
3404 Ok(Executed {
3405 rows: if has_rows { written } else { vec![] },
3406 project: returning,
3407 has_rows,
3408 tag: format!("UPDATE {}", n),
3409 tag_counts_rows: false,
3410 })
3411 }
3412
3413 Stmt::Delete { coll, nql, returning } => {
3414 let db = need_db!();
3415 need_write!();
3416 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
3417 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3418 })?;
3419 let returned = matched.clone();
3422 let mut n = 0usize;
3423 for row in &matched {
3424 if let Some(id) = row.get("_id").and_then(|v| v.as_str()) {
3425 match db.delete(&coll, id) {
3426 Ok(true) => n += 1,
3427 Ok(false) => {}
3428 Err(e) => return Err(err_msg("XX000", &format!("DELETE failed: {}", e))),
3429 }
3430 }
3431 }
3432 let has_rows = wants_returning(stmt_sql);
3433 Ok(Executed {
3434 rows: if has_rows { returned } else { vec![] },
3435 project: returning,
3436 has_rows,
3437 tag: format!("DELETE {}", n),
3438 tag_counts_rows: false,
3439 })
3440 }
3441 }
3442}
3443
3444fn run_simple_query(sql: &str, db_name: &str, db: Option<&Arc<Db>>, read_only: bool) -> Vec<u8> {
3446 let mut out = vec![];
3447 let statements = split_statements(sql);
3448 if statements.is_empty() {
3449 return Out::msg(b'I').finish();
3451 }
3452 for stmt_sql in statements {
3453 match execute_stmt(&stmt_sql, db_name, db, read_only) {
3454 Err(encoded) => {
3456 out.extend_from_slice(&encoded);
3457 return out;
3458 }
3459 Ok(ex) => {
3460 if ex.has_rows {
3461 out.extend_from_slice(&encode_rows(&ex.rows, &ex.project));
3462 }
3463 out.extend_from_slice(&command_complete(&ex.tag_for(ex.rows.len())));
3464 }
3465 }
3466 }
3467 out
3468}
3469
3470fn split_statements(sql: &str) -> Vec<String> {
3472 let mut out = vec![];
3473 let mut cur = String::new();
3474 let mut in_s = false;
3475 for c in sql.chars() {
3476 match c {
3477 '\'' => { in_s = !in_s; cur.push(c); }
3478 ';' if !in_s => {
3479 if !cur.trim().is_empty() { out.push(cur.clone()); }
3480 cur.clear();
3481 }
3482 _ => cur.push(c),
3483 }
3484 }
3485 if !cur.trim().is_empty() {
3486 out.push(cur);
3487 }
3488 out
3489}
3490
3491pub async fn run(host: &str, port: u16, resolver: Arc<dyn DbResolver>) -> anyhow::Result<()> {
3493 let read_only = std::env::var("NEDBD_PG_READ_ONLY")
3497 .map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
3498 .unwrap_or(false);
3499 let listener = TcpListener::bind((host, port)).await?;
3500 println!(" pgwire postgres endpoint on {}:{} — psql / DBeaver / psycopg ({})",
3501 host, port,
3502 if read_only { "SELECT only — read-only mode" } else { "SELECT + INSERT/UPDATE/DELETE" });
3503 loop {
3504 let (sock, _peer) = match listener.accept().await {
3505 Ok(v) => v,
3506 Err(e) => {
3507 eprintln!(" [pgwire] accept failed: {}", e);
3508 continue;
3509 }
3510 };
3511 let r = Arc::clone(&resolver);
3512 tokio::spawn(async move {
3513 let _ = sock.set_nodelay(true);
3514 if let Err(e) = handle(sock, r, read_only).await {
3515 if e.kind() != std::io::ErrorKind::UnexpectedEof
3517 && e.kind() != std::io::ErrorKind::ConnectionReset
3518 {
3519 eprintln!(" [pgwire] connection error: {}", e);
3520 }
3521 }
3522 });
3523 }
3524}
3525
3526#[cfg(test)]
3529mod explain_tests {
3530 use super::*;
3531
3532 #[test]
3533 fn a_bare_explain_is_stripped() {
3534 assert_eq!(strip_explain("EXPLAIN SELECT 1"), Some("SELECT 1"));
3535 assert_eq!(strip_explain("explain select 1"), Some("select 1"));
3536 assert_eq!(strip_explain(" EXPLAIN SELECT 1 ; "), Some("SELECT 1"));
3537 }
3538
3539 #[test]
3540 fn analyze_and_verbose_are_accepted_and_ignored() {
3541 assert_eq!(strip_explain("EXPLAIN ANALYZE SELECT 1"), Some("SELECT 1"));
3545 assert_eq!(strip_explain("EXPLAIN ANALYSE SELECT 1"), Some("SELECT 1"));
3546 assert_eq!(strip_explain("EXPLAIN VERBOSE SELECT 1"), Some("SELECT 1"));
3547 assert_eq!(strip_explain("EXPLAIN ANALYZE VERBOSE SELECT 1"), Some("SELECT 1"));
3548 assert_eq!(strip_explain("explain analyze verbose select 1"), Some("select 1"));
3549 }
3550
3551 #[test]
3552 fn a_word_merely_starting_with_explain_is_not_a_keyword() {
3553 assert_eq!(strip_explain("EXPLAINED SELECT 1"), None);
3554 assert_eq!(strip_explain("SELECT 1"), None);
3555 assert_eq!(strip_explain("SELECT explain FROM t"), None);
3556 }
3557
3558 #[test]
3559 fn a_column_named_analyze_is_not_eaten() {
3560 assert_eq!(strip_explain("EXPLAIN analyzed_view"), Some("analyzed_view"));
3563 }
3564
3565 #[test]
3566 fn the_plan_result_has_postgres_shape() {
3567 let e = plan_result(vec!["Seq Scan on t".into(), "note".into()]);
3568 assert_eq!(e.project.len(), 1);
3569 assert_eq!(e.project[0].out, "QUERY PLAN");
3570 assert_eq!(e.rows.len(), 2);
3571 assert_eq!(e.rows[0]["QUERY PLAN"], "Seq Scan on t");
3572 assert_eq!(e.tag, "EXPLAIN");
3573 assert!(!e.tag_counts_rows);
3575 }
3576}
3577
3578#[cfg(test)]
3579mod tests {
3580 use super::*;
3581 use serde_json::json;
3582
3583 fn q(sql: &str) -> String {
3584 match translate(sql) {
3585 Ok(Stmt::Query { nql, .. }) => nql,
3586 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3587 }
3588 }
3589 fn proj(sql: &str) -> Vec<String> {
3591 match translate(sql) {
3592 Ok(Stmt::Query { project, .. }) => project.iter().map(|c| c.out.clone()).collect(),
3593 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3594 }
3595 }
3596 fn proj_pairs(sql: &str) -> Vec<(String, String)> {
3598 match translate(sql) {
3599 Ok(Stmt::Query { project, .. }) =>
3600 project.iter().map(|c| (c.src.clone(), c.out.clone())).collect(),
3601 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3602 }
3603 }
3604 fn names(cols: &[Col]) -> Vec<String> { cols.iter().map(|c| c.out.clone()).collect() }
3605
3606 fn cols_of(sql: &str) -> Vec<Col> {
3610 match translate(sql).unwrap() {
3611 Stmt::Query { project, .. } => project,
3612 other => panic!("{:?}", other),
3613 }
3614 }
3615
3616 #[test]
3617 fn select_star_becomes_bare_from() {
3618 assert_eq!(q("SELECT * FROM orders"), "FROM orders");
3619 assert_eq!(q("select * from orders;"), "FROM orders");
3620 assert_eq!(proj("SELECT * FROM orders"), Vec::<String>::new());
3621 }
3622
3623 #[test]
3624 fn a_column_list_becomes_a_projection_not_a_clause() {
3625 assert_eq!(q("SELECT status, total FROM orders"), "FROM orders");
3628 assert_eq!(proj("SELECT status, total FROM orders"), vec!["status", "total"]);
3629 }
3630
3631 #[test]
3632 fn a_qualifier_reduces_to_the_field_while_an_ALIAS_is_the_name_the_client_sees() {
3633 let cols = cols_of("SELECT o.status AS s, o.total total, o.region FROM orders o");
3641 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
3642 vec!["status", "total", "region"]);
3643 assert_eq!(cols.iter().map(|c| c.out.clone()).collect::<Vec<_>>(),
3644 vec!["s", "total", "region"]);
3645 assert_eq!(q("SELECT * FROM public.orders"), "FROM orders");
3646 assert_eq!(q("SELECT * FROM \"orders\""), "FROM orders");
3647 }
3648
3649 #[test]
3650 fn a_select_list_may_MIX_columns_with_an_aggregate() {
3651 assert_eq!(q("SELECT status, count(*) AS count_1 FROM orders GROUP BY status"),
3659 "FROM orders GROUP BY status COUNT");
3660 assert_eq!(q("SELECT status, count(*) FROM orders WHERE total > 1 GROUP BY status ORDER BY status LIMIT 5"),
3663 "FROM orders WHERE total > 1 GROUP BY status COUNT ORDER BY status LIMIT 5");
3664 assert_eq!(q("SELECT count(*) FROM orders"), "FROM orders COUNT");
3666 assert_eq!(q("SELECT sum(total) FROM orders"), "FROM orders SUM total");
3667 let e = translate("SELECT status, count(*) FROM orders GROUP BY status, region").unwrap_err();
3671 assert!(e.contains("GROUP BY takes one key"), "{}", e);
3672 let cols = cols_of("SELECT status, count(*) AS count_1 FROM orders GROUP BY status");
3673 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
3674 vec!["status", "count"]);
3675 assert_eq!(cols.iter().map(|c| c.out.clone()).collect::<Vec<_>>(),
3676 vec!["status", "count_1"]);
3677
3678 let cols = cols_of("SELECT status, count(*), sum(total) FROM orders GROUP BY status");
3681 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
3682 vec!["status", "count", "sum_total"]);
3683 assert_eq!(q("SELECT status, count(*), sum(total) FROM orders GROUP BY status"),
3684 "FROM orders GROUP BY status SUM total");
3685
3686 assert_eq!(q("SELECT o.status, sum(o.total) FROM orders o GROUP BY o.status"),
3688 "FROM orders GROUP BY status SUM total");
3689
3690 let e = translate("SELECT status, sum(total), avg(total) FROM orders GROUP BY status")
3693 .unwrap_err();
3694 assert!(e.contains("only one of SUM/AVG/MIN/MAX"), "{}", e);
3695
3696 let e = translate("SELECT status, total, count(*) FROM orders GROUP BY status")
3698 .unwrap_err();
3699 assert!(e.contains("must appear in the GROUP BY clause"), "{}", e);
3700 }
3701
3702 #[test]
3703 fn ORDER_BY_an_ordinal_resolves_to_that_select_list_column() {
3704 assert_eq!(q("SELECT status, total FROM orders ORDER BY 1"),
3709 "FROM orders ORDER BY status");
3710 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 DESC"),
3711 "FROM orders ORDER BY total DESC");
3712 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 DESC, 1"),
3714 "FROM orders ORDER BY total DESC, status");
3715 assert_eq!(q("SELECT status, total FROM orders ORDER BY 1, total DESC"),
3716 "FROM orders ORDER BY status, total DESC");
3717 assert_eq!(q("SELECT status, count(*) AS n FROM orders GROUP BY status ORDER BY 1"),
3721 "FROM orders GROUP BY status COUNT ORDER BY status");
3722 assert_eq!(q("SELECT status, count(*) AS n FROM orders GROUP BY status ORDER BY 2 DESC"),
3724 "FROM orders GROUP BY status COUNT ORDER BY count DESC");
3725 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 LIMIT 1"),
3728 "FROM orders ORDER BY total LIMIT 1");
3729 assert_eq!(q("SELECT status FROM orders WHERE total > 1 ORDER BY 1"),
3731 "FROM orders WHERE total > 1 ORDER BY status");
3732
3733 let e = translate("SELECT status FROM orders ORDER BY 4").unwrap_err();
3737 assert!(e.contains("out of range") && e.contains("1 column"), "{}", e);
3738 let e = translate("SELECT * FROM orders ORDER BY 1").unwrap_err();
3739 assert!(e.contains("no list to index"), "{}", e);
3740 }
3741
3742 #[test]
3743 fn count_of_a_subquery_flattens_only_when_the_two_counts_MUST_agree() {
3744 assert_eq!(
3748 q("SELECT count(*) AS count_1 FROM (SELECT orders._id AS a, orders.status AS b \
3749 FROM orders WHERE orders.status = 'paid') AS anon_1"),
3750 r#"FROM orders COUNT WHERE status = "paid""#);
3753 assert_eq!(q("SELECT count(*) FROM (SELECT orders._id FROM orders) AS anon_1"),
3755 "FROM orders COUNT");
3756 assert_eq!(q("SELECT count(*) FROM (SELECT _id FROM orders ORDER BY total DESC) AS a"),
3758 "FROM orders COUNT");
3759 let cols = cols_of("SELECT count(*) AS count_1 FROM (SELECT _id FROM orders) AS a");
3761 assert_eq!(cols[0].src, "count");
3762 assert_eq!(cols[0].out, "count_1");
3763
3764 for sql in [
3767 "SELECT count(*) FROM (SELECT _id FROM orders LIMIT 1) AS a",
3769 "SELECT count(*) FROM (SELECT _id FROM orders OFFSET 1) AS a",
3770 "SELECT count(*) FROM (SELECT status FROM orders GROUP BY status) AS a",
3772 "SELECT count(*) FROM (SELECT count(*) FROM orders) AS a",
3774 "SELECT count(*) FROM (SELECT sum(total) FROM orders) AS a",
3775 "SELECT count(*), status FROM (SELECT status FROM orders) AS a",
3777 "SELECT status FROM (SELECT status FROM orders) AS a",
3778 "SELECT count(*) FROM (SELECT x FROM (SELECT _id AS x FROM orders) AS b) AS a",
3780 ] {
3781 let e = translate(sql).unwrap_err();
3782 assert!(e.contains("subqueries in FROM"), "{} -> {}", sql, e);
3783 }
3784
3785 for (sql, needle) in [
3790 ("SELECT count(*) FROM (SELECT DISTINCT status FROM orders) AS a", "DISTINCT"),
3791 ("SELECT count(*) FROM (SELECT a FROM t UNION SELECT b FROM u) AS x", "UNION"),
3792 ] {
3793 let e = translate(sql).unwrap_err();
3794 assert!(e.contains(needle), "{} -> {}", sql, e);
3795 }
3796 }
3797
3798 #[test]
3799 fn a_QUALIFIED_column_in_WHERE_finds_its_field_instead_of_ZERO_ROWS() {
3800 assert_eq!(q("SELECT _id FROM orders WHERE orders.status = 'paid'"),
3808 r#"FROM orders WHERE status = "paid""#);
3809 assert_eq!(q("SELECT _id FROM orders WHERE orders.total > 50"),
3810 "FROM orders WHERE total > 50");
3811 assert_eq!(q("SELECT _id FROM orders ORDER BY orders.total DESC LIMIT 2"),
3813 "FROM orders ORDER BY total DESC LIMIT 2");
3814 assert_eq!(q("SELECT status, count(*) FROM orders GROUP BY orders.status"),
3815 "FROM orders GROUP BY status COUNT");
3816
3817 assert_eq!(q("SELECT o.status FROM orders o WHERE o.status = 'paid'"),
3821 r#"FROM orders WHERE status = "paid""#);
3822 assert_eq!(q("SELECT o.status FROM orders AS o WHERE o.total > 1"),
3823 "FROM orders WHERE total > 1");
3824
3825 let e = translate("SELECT _id FROM orders WHERE nosuch.status = 'paid'").unwrap_err();
3830 assert!(e.contains("no table or alias named \"nosuch\""), "{}", e);
3831 let e = translate("SELECT _id FROM orders o WHERE p.status = 'paid'").unwrap_err();
3832 assert!(e.contains("aliased \"o\""), "the message names the alias in scope: {}", e);
3833
3834 assert_eq!(q("SELECT _id FROM orders WHERE status = 'pa.id'"),
3836 r#"FROM orders WHERE status = "pa.id""#);
3837 assert_eq!(q("SELECT _id FROM orders WHERE total > 1.5"),
3839 "FROM orders WHERE total > 1.5");
3840
3841 match translate("UPDATE orders o SET status = 'x' WHERE o.total > 5").unwrap() {
3843 Stmt::Update { coll, nql, .. } => {
3844 assert_eq!(coll, "orders", "the alias is not part of the collection name");
3845 assert_eq!(nql, "FROM orders WHERE total > 5");
3846 }
3847 other => panic!("{:?}", other),
3848 }
3849 match translate("DELETE FROM orders o WHERE o.status = 'paid'").unwrap() {
3850 Stmt::Delete { coll, nql, .. } => {
3851 assert_eq!(coll, "orders");
3852 assert_eq!(nql, r#"FROM orders WHERE status = "paid""#);
3853 }
3854 other => panic!("{:?}", other),
3855 }
3856
3857 assert_eq!(q("SELECT _id FROM orders AS OF SYSTEM TIME 3 WHERE orders.total > 1"),
3859 "FROM orders AS OF 3 WHERE total > 1");
3860 }
3861
3862 #[test]
3863 fn where_clauses_pass_through_with_sql_literals_rewritten() {
3864 assert_eq!(q("SELECT * FROM orders WHERE status = 'paid'"),
3865 r#"FROM orders WHERE status = "paid""#);
3866 assert_eq!(q("SELECT * FROM orders WHERE status <> 'paid'"),
3867 r#"FROM orders WHERE status != "paid""#);
3868 assert_eq!(q("SELECT * FROM orders WHERE status IN ('paid','open')"),
3869 r#"FROM orders WHERE status IN ("paid","open")"#);
3870 }
3871
3872 #[test]
3875 fn a_doubled_sql_quote_is_one_literal_character() {
3876 assert_eq!(q("SELECT * FROM t WHERE name = 'it''s'"),
3877 r#"FROM t WHERE name = "it's""#);
3878 }
3879
3880 #[test]
3883 fn a_double_quote_inside_a_sql_literal_is_escaped_for_nql() {
3884 assert_eq!(q(r#"SELECT * FROM t WHERE name = 'say "hi"'"#),
3885 r#"FROM t WHERE name = "say \"hi\"""#);
3886 }
3887
3888 #[test]
3889 fn the_shared_clauses_are_handed_to_nql_unchanged() {
3890 assert_eq!(q("SELECT * FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5"),
3891 "FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5");
3892 assert_eq!(q("SELECT * FROM orders GROUP BY region"), "FROM orders GROUP BY region");
3893 assert_eq!(q("SELECT * FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC"),
3894 "FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC");
3895 }
3896
3897 #[test]
3904 fn an_aggregate_is_one_column_named_as_sql_names_it() {
3905 assert_eq!(proj_pairs("SELECT COUNT(*) FROM orders"),
3906 vec![("count".to_string(), "count".to_string())]);
3907 assert_eq!(proj_pairs("SELECT SUM(total) FROM orders"),
3908 vec![("sum_total".to_string(), "sum".to_string())]);
3909 assert_eq!(proj_pairs("SELECT avg(total) FROM orders"),
3910 vec![("avg_total".to_string(), "avg".to_string())]);
3911 assert_eq!(proj_pairs("SELECT MIN(total) FROM orders"),
3912 vec![("min_total".to_string(), "min".to_string())]);
3913 let rows = vec![json!({"count": 4, "sum_total": 420, "value": 420})];
3915 let p = vec![Col::renamed("sum_total", "sum")];
3916 let cols = columns_for(&rows, &p);
3917 assert_eq!(names(&cols), vec!["sum"], "one column, SQL's name");
3918 assert_eq!(cell(rows[0].get(&cols[0].src)), Some("420".to_string()));
3919 }
3920
3921 #[test]
3926 fn a_bare_column_with_group_by_is_refused_not_nulled() {
3927 let e = translate("SELECT region, total FROM orders GROUP BY region").unwrap_err();
3928 assert!(e.contains("must appear in the GROUP BY clause"), "{}", e);
3929 assert!(e.contains("total"), "the message names the offending column: {}", e);
3930
3931 assert!(translate("SELECT region FROM orders GROUP BY region").is_ok());
3933 assert!(translate("SELECT region, count FROM orders GROUP BY region").is_ok());
3934 assert!(translate("SELECT SUM(total) FROM orders GROUP BY region").is_ok());
3936 assert!(translate("SELECT * FROM orders GROUP BY region").is_ok());
3938 }
3939
3940 #[test]
3941 fn count_star_becomes_nql_count() {
3942 assert_eq!(q("SELECT COUNT(*) FROM orders"), "FROM orders COUNT");
3943 assert_eq!(q("SELECT count(*) FROM orders WHERE total > 5"),
3944 "FROM orders COUNT WHERE total > 5");
3945 }
3946
3947 #[test]
3948 fn aggregates_carry_their_target_column() {
3949 assert_eq!(q("SELECT SUM(total) FROM orders"), "FROM orders SUM total");
3950 assert_eq!(q("SELECT avg(total) FROM orders WHERE region = 'eu'"),
3951 r#"FROM orders AVG total WHERE region = "eu""#);
3952 assert!(translate("SELECT SUM(*) FROM orders").is_err());
3953 }
3954
3955 #[test]
3958 fn as_of_system_time_bridges_to_nql_as_of() {
3959 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42"),
3960 "FROM orders AS OF 42");
3961 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42 WHERE total > 1"),
3962 "FROM orders AS OF 42 WHERE total > 1");
3963 let e = translate("SELECT * FROM orders AS OF SYSTEM TIME '2026-01-01'").unwrap_err();
3965 assert!(e.contains("sequence number"), "{}", e);
3966 }
3967
3968 #[test]
3969 fn handshake_queries_are_answered_so_clients_can_connect() {
3970 assert!(matches!(translate("SELECT version()"), Ok(Stmt::Canned { .. })));
3971 assert!(matches!(translate("SHOW transaction_isolation"), Ok(Stmt::Canned { .. })));
3972 assert!(matches!(translate("SELECT current_schema()"), Ok(Stmt::Canned { .. })));
3973 assert!(matches!(translate("SET extra_float_digits = 3"), Ok(Stmt::Ok(_))));
3974 assert!(matches!(translate("BEGIN"), Ok(Stmt::Ok(_))));
3975 assert!(matches!(translate(""), Ok(Stmt::Ok(_))));
3976 }
3977
3978 #[test]
3981 fn unsupported_sql_is_refused_with_a_reason() {
3982 for (sql, expect) in [
3983 ("INSERT INTO t VALUES (1)", "explicit column list"),
3984 ("CREATE TABLE t (a int)", "DDL"),
3985 ("TRUNCATE t", "append-only"),
3986 ("GRANT ALL ON t TO x", "privilege system"),
3987 ("SELECT * FROM a JOIN b ON a.x = b.x", "JOIN is not supported"),
3988 ("SELECT * FROM a UNION SELECT * FROM b", "UNION"),
3989 ("SELECT DISTINCT region FROM orders", "GROUP BY"),
3990 ("SELECT * FROM (SELECT 1) x", "subqueries in FROM"),
3991 ("SELECT * FROM a, b", "more than one collection"),
3992 ("SELECT lower(status) FROM orders", "expressions in the select list"),
3993 ("VACUUM", "only SELECT"),
3994 ] {
3995 let e = translate(sql).unwrap_err();
3996 assert!(e.contains(expect), "for {:?} expected {:?} in {:?}", sql, expect, e);
3997 }
3998 }
3999
4000 fn ins(sql: &str) -> (String, Vec<InsertRow>, Vec<Col>) {
4008 match translate(sql) {
4009 Ok(Stmt::Insert { coll, rows, returning }) => (coll, rows, returning),
4010 other => panic!("expected INSERT for {:?}, got {:?}", sql, other),
4011 }
4012 }
4013
4014 #[test]
4015 fn insert_becomes_a_put_per_row() {
4016 let (coll, rows, ret) = ins("INSERT INTO orders (_id, status, total) VALUES ('o1', 'paid', 120)");
4017 assert_eq!(coll, "orders");
4018 assert_eq!(rows.len(), 1);
4019 assert_eq!(rows[0].id.as_deref(), Some("o1"));
4020 assert_eq!(rows[0].doc.get("status"), Some(&json!("paid")));
4021 assert_eq!(rows[0].doc.get("total"), Some(&json!(120)));
4022 assert!(!rows[0].doc.contains_key("_id"));
4024 assert!(ret.is_empty());
4025 }
4026
4027 #[test]
4028 fn a_multi_row_insert_yields_one_row_each() {
4029 let (_, rows, _) = ins(
4030 "INSERT INTO t (id, n) VALUES ('a', 1), ('b', 2), ('c', 3)");
4031 assert_eq!(rows.len(), 3);
4032 assert_eq!(rows[1].id.as_deref(), Some("b"));
4033 assert_eq!(rows[2].doc.get("n"), Some(&json!(3)));
4034 }
4035
4036 #[test]
4037 fn an_insert_without_an_id_column_lets_the_server_assign_one() {
4038 let (_, rows, _) = ins("INSERT INTO t (n) VALUES (1)");
4039 assert_eq!(rows[0].id, None, "the executor mints a unique key");
4040 assert_eq!(rows[0].doc.get("n"), Some(&json!(1)));
4041 }
4042
4043 #[test]
4046 fn insert_lifts_provenance_out_of_reserved_columns() {
4047 let (_, rows, _) = ins(
4048 "INSERT INTO audit (_id, _caused_by, _valid_from, kind) \
4049 VALUES ('e1', 'abc123', '2026-01-01', 'reprice')");
4050 assert_eq!(rows[0].caused_by, vec!["abc123".to_string()]);
4051 assert_eq!(rows[0].valid_from.as_deref(), Some("2026-01-01"));
4052 assert_eq!(rows[0].doc.get("kind"), Some(&json!("reprice")));
4053 for k in ["_id", "_caused_by", "_valid_from"] {
4055 assert!(!rows[0].doc.contains_key(k), "{} leaked into the doc", k);
4056 }
4057 }
4058
4059 #[test]
4060 fn insert_values_cover_the_scalar_types() {
4061 let (_, rows, _) = ins(
4062 "INSERT INTO t (s, i, f, b, n) VALUES ('x', 42, 1.5, TRUE, NULL)");
4063 assert_eq!(rows[0].doc.get("s"), Some(&json!("x")));
4064 assert_eq!(rows[0].doc.get("i"), Some(&json!(42)));
4065 assert_eq!(rows[0].doc.get("f"), Some(&json!(1.5)));
4066 assert_eq!(rows[0].doc.get("b"), Some(&json!(true)));
4067 assert_eq!(rows[0].doc.get("n"), Some(&Value::Null));
4068 }
4069
4070 #[test]
4073 fn insert_literals_survive_quotes_and_commas() {
4074 let (_, rows, _) = ins("INSERT INTO t (a, b) VALUES ('it''s', 'x,y')");
4075 assert_eq!(rows[0].doc.get("a"), Some(&json!("it's")));
4076 assert_eq!(rows[0].doc.get("b"), Some(&json!("x,y")));
4077 }
4078
4079 #[test]
4080 fn insert_refuses_what_it_cannot_store_faithfully() {
4081 assert!(translate("INSERT INTO t (a) VALUES (1 + 1)").is_err());
4083 assert!(translate("INSERT INTO t (a) VALUES (now())").is_err());
4084 let e = translate("INSERT INTO t (a, b) VALUES (1)").unwrap_err();
4086 assert!(e.contains("values for"), "{}", e);
4087 let e2 = translate("INSERT INTO t VALUES (1)").unwrap_err();
4089 assert!(e2.contains("explicit column list"), "{}", e2);
4090 }
4091
4092 #[test]
4093 fn update_finds_rows_with_the_full_predicate_surface() {
4094 match translate("UPDATE orders SET status = 'void' WHERE total < 50 AND region IN ('eu')") {
4095 Ok(Stmt::Update { coll, set, nql, .. }) => {
4096 assert_eq!(coll, "orders");
4097 assert_eq!(set, vec![("status".to_string(), json!("void"))]);
4098 assert_eq!(nql, r#"FROM orders WHERE total < 50 AND region IN ("eu")"#);
4100 }
4101 other => panic!("expected UPDATE, got {:?}", other),
4102 }
4103 }
4104
4105 #[test]
4106 fn update_without_where_targets_the_whole_collection() {
4107 match translate("UPDATE t SET a = 1") {
4109 Ok(Stmt::Update { nql, .. }) => assert_eq!(nql, "FROM t"),
4110 other => panic!("expected UPDATE, got {:?}", other),
4111 }
4112 }
4113
4114 #[test]
4115 fn update_handles_several_assignments() {
4116 match translate("UPDATE t SET a = 1, b = 'x,y', c = NULL WHERE id = 'k'") {
4117 Ok(Stmt::Update { set, .. }) => {
4118 assert_eq!(set.len(), 3);
4119 assert_eq!(set[1], ("b".to_string(), json!("x,y")));
4120 assert_eq!(set[2], ("c".to_string(), Value::Null));
4121 }
4122 other => panic!("expected UPDATE, got {:?}", other),
4123 }
4124 assert!(translate("UPDATE t SET").is_err());
4125 assert!(translate("UPDATE t SET a").is_err());
4126 }
4127
4128 #[test]
4129 fn delete_becomes_a_predicate_over_the_collection() {
4130 match translate("DELETE FROM orders WHERE status = 'void'") {
4131 Ok(Stmt::Delete { coll, nql, .. }) => {
4132 assert_eq!(coll, "orders");
4133 assert_eq!(nql, r#"FROM orders WHERE status = "void""#);
4134 }
4135 other => panic!("expected DELETE, got {:?}", other),
4136 }
4137 match translate("DELETE FROM t") {
4138 Ok(Stmt::Delete { nql, .. }) => assert_eq!(nql, "FROM t"),
4139 other => panic!("expected DELETE, got {:?}", other),
4140 }
4141 }
4142
4143 #[test]
4144 fn returning_is_parsed_off_every_write() {
4145 let (_, _, ret) = ins("INSERT INTO t (a) VALUES (1) RETURNING a, _id");
4146 assert_eq!(ret.iter().map(|c| c.out.clone()).collect::<Vec<_>>(), vec!["a", "_id"]);
4147 let (_, _, star) = ins("INSERT INTO t (a) VALUES (1) RETURNING *");
4150 assert!(star.is_empty());
4151 assert!(wants_returning("INSERT INTO t (a) VALUES (1) RETURNING *"));
4152 assert!(!wants_returning("INSERT INTO t (a) VALUES (1)"));
4153
4154 match translate("UPDATE t SET a = 1 WHERE id = 'k' RETURNING a") {
4155 Ok(Stmt::Update { nql, returning, .. }) => {
4156 assert_eq!(returning.len(), 1);
4157 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql);
4159 }
4160 other => panic!("expected UPDATE, got {:?}", other),
4161 }
4162 match translate("DELETE FROM t WHERE id = 'k' RETURNING *") {
4163 Ok(Stmt::Delete { nql, .. }) =>
4164 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql),
4165 other => panic!("expected DELETE, got {:?}", other),
4166 }
4167 }
4168
4169 #[test]
4170 fn a_keyword_inside_a_value_is_not_a_clause() {
4171 match translate("UPDATE t SET note = 'where returning from' WHERE id = 'k'") {
4172 Ok(Stmt::Update { set, nql, .. }) => {
4173 assert_eq!(set[0].1, json!("where returning from"));
4174 assert_eq!(nql, r#"FROM t WHERE id = "k""#);
4175 }
4176 other => panic!("expected UPDATE, got {:?}", other),
4177 }
4178 }
4179
4180 #[test]
4181 fn split_top_respects_quotes_and_nesting() {
4182 assert_eq!(split_top("a, b, c", ',').len(), 3);
4183 assert_eq!(split_top("(1, 2), (3, 4)", ',').len(), 2);
4184 assert_eq!(split_top("'a,b', c", ',').len(), 2);
4185 assert_eq!(split_top("'it''s, fine', c", ',').len(), 2);
4186 }
4187
4188 #[test]
4189 fn comments_and_whitespace_do_not_confuse_the_translator() {
4190 assert_eq!(q("SELECT *\n FROM orders -- trailing note\n"), "FROM orders");
4191 assert_eq!(q("SELECT /* inline */ * FROM orders"), "FROM orders");
4192 assert_eq!(q("SELECT * FROM t WHERE note = 'from here to JOIN'"),
4194 r#"FROM t WHERE note = "from here to JOIN""#);
4195 }
4196
4197 #[test]
4198 fn find_kw_ignores_quotes_parens_and_substrings() {
4199 assert_eq!(find_kw("SELECT A FROM B", "FROM"), Some(9));
4200 assert_eq!(find_kw("SELECT 'FROM' FROM B", "FROM"), Some(14));
4201 assert_eq!(find_kw("SELECT F(x FROM y) FROM B", "FROM"), Some(19));
4202 assert_eq!(find_kw("SELECT FROMAGE", "FROM"), None);
4203 assert_eq!(find_kw("SELECT X_FROM", "FROM"), None);
4204 }
4205
4206 #[test]
4209 fn provenance_columns_sort_after_the_users_own_fields() {
4210 let rows = vec![json!({"_id":"1","_hash":"ab","status":"paid","total":9})];
4211 assert_eq!(names(&columns_for(&rows, &[])),
4212 vec!["status", "total", "_hash", "_id"]);
4213 }
4214
4215 #[test]
4216 fn an_explicit_projection_sets_the_column_order() {
4217 let rows = vec![json!({"a":1,"b":2})];
4218 let p = vec![Col::same("b"), Col::same("a")];
4219 assert_eq!(names(&columns_for(&rows, &p)), vec!["b", "a"]);
4220 }
4221
4222 #[test]
4223 fn columns_are_the_union_across_sparse_rows() {
4224 let rows = vec![json!({"a":1}), json!({"b":2})];
4226 assert_eq!(names(&columns_for(&rows, &[])), vec!["a", "b"]);
4227 }
4228
4229 #[test]
4230 fn type_oids_follow_the_first_non_null_value() {
4231 let rows = vec![json!({"i":1,"f":1.5,"b":true,"s":"x","n":null})];
4232 assert_eq!(oid_for(&rows, "i"), OID_INT8);
4233 assert_eq!(oid_for(&rows, "f"), OID_FLOAT8);
4234 assert_eq!(oid_for(&rows, "b"), OID_BOOL);
4235 assert_eq!(oid_for(&rows, "s"), OID_TEXT);
4236 assert_eq!(oid_for(&rows, "n"), OID_TEXT);
4238 assert_eq!(oid_for(&rows, "absent"), OID_TEXT);
4239 }
4240
4241 #[test]
4242 fn a_column_that_is_null_in_the_first_row_still_gets_its_type() {
4243 let rows = vec![json!({"v": null}), json!({"v": 7})];
4244 assert_eq!(oid_for(&rows, "v"), OID_INT8);
4245 }
4246
4247 #[test]
4248 fn cells_render_in_postgres_text_format() {
4249 assert_eq!(cell(Some(&json!("x"))), Some("x".to_string()));
4250 assert_eq!(cell(Some(&json!(true))), Some("t".to_string()));
4251 assert_eq!(cell(Some(&json!(false))), Some("f".to_string()));
4252 assert_eq!(cell(Some(&json!(42))), Some("42".to_string()));
4253 assert_eq!(cell(Some(&json!(null))), None);
4254 assert_eq!(cell(None), None);
4255 assert_eq!(cell(Some(&json!({"a":1}))), Some("{\"a\":1}".to_string()));
4257 }
4258
4259 #[test]
4262 fn message_framing_length_excludes_the_tag() {
4263 let mut m = Out::msg(b'Z');
4264 m.bytes(b"I");
4265 let bytes = m.finish();
4266 assert_eq!(bytes[0], b'Z');
4267 assert_eq!(i32::from_be_bytes([bytes[1], bytes[2], bytes[3], bytes[4]]), 5);
4268 assert_eq!(bytes.len(), 6);
4269 }
4270
4271 #[test]
4272 fn a_result_set_encodes_as_description_then_rows_then_complete() {
4273 let rows = vec![json!({"a": 1}), json!({"a": 2})];
4274 let out = encode_result(&rows, &[]);
4275 assert_eq!(out[0], b'T');
4276 let tags: Vec<u8> = {
4277 let mut t = vec![];
4279 let mut i = 0usize;
4280 while i < out.len() {
4281 t.push(out[i]);
4282 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
4283 i += 1 + len;
4284 }
4285 t
4286 };
4287 assert_eq!(tags, vec![b'T', b'D', b'D', b'C'],
4288 "one description, one row each, one completion");
4289 }
4290
4291 #[test]
4296 fn a_write_with_returning_emits_exactly_one_command_complete() {
4297 let rows = vec![json!({"_id": "o1", "total": 9})];
4298 let mut out = encode_rows(&rows, &[Col::same("_id")]);
4299 out.extend_from_slice(&command_complete("INSERT 0 1"));
4300 let mut tags = vec![];
4301 let mut i = 0usize;
4302 while i < out.len() {
4303 tags.push(out[i]);
4304 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
4305 i += 1 + len;
4306 }
4307 assert_eq!(tags, vec![b'T', b'D', b'C'], "one description, one row, ONE tag");
4308 assert_eq!(tags.iter().filter(|t| **t == b'C').count(), 1);
4309 assert!(!encode_rows(&rows, &[]).contains(&b'C')
4311 || encode_rows(&rows, &[]).iter().filter(|b| **b == b'C').count() > 0);
4312 let bare = encode_rows(&rows, &[Col::same("_id")]);
4313 let mut bare_tags = vec![];
4314 let mut j = 0usize;
4315 while j < bare.len() {
4316 bare_tags.push(bare[j]);
4317 let len = i32::from_be_bytes([bare[j+1], bare[j+2], bare[j+3], bare[j+4]]) as usize;
4318 j += 1 + len;
4319 }
4320 assert_eq!(bare_tags, vec![b'T', b'D'], "encode_rows never appends a tag");
4321 }
4322
4323 #[test]
4324 fn an_empty_result_still_sends_a_description() {
4325 let out = encode_result(&[], &[Col::same("a")]);
4326 assert_eq!(out[0], b'T', "clients need the shape even with no rows");
4327 }
4328
4329 #[test]
4330 fn statements_split_on_top_level_semicolons_only() {
4331 assert_eq!(split_statements("SELECT 1; SELECT 2").len(), 2);
4332 assert_eq!(split_statements("SELECT ';'").len(), 1);
4333 assert_eq!(split_statements("SELECT 1;").len(), 1);
4334 assert_eq!(split_statements(" ").len(), 0);
4335 }
4336
4337 #[test]
4338 fn an_error_names_its_sqlstate() {
4339 let e = String::from_utf8_lossy(&err_msg("0A000", "x")).to_string();
4340 assert!(e.contains("ERROR"));
4341 assert!(e.contains("0A000"));
4342 }
4343
4344 #[test]
4347 fn placeholders_are_counted_outside_string_literals() {
4348 assert_eq!(param_count("SELECT a FROM t WHERE b = $1 AND c = $2"), 2);
4349 assert_eq!(param_count("SELECT a FROM t"), 0);
4350 assert_eq!(param_count("WHERE a = $2 OR b = $2 OR c = $1"), 2);
4352 assert_eq!(param_count("SELECT a FROM t WHERE b = '$1'"), 0,
4353 "a placeholder inside a literal is data, not a parameter");
4354 assert_eq!(param_count("WHERE a = $10 AND b = $1"), 10,
4355 "two-digit indexes must not be read as $1 followed by 0");
4356 }
4357
4358 #[test]
4359 fn parameters_are_spliced_as_literals() {
4360 let out = substitute_params("WHERE a = $1 AND b = $2 AND c = $3",
4361 &[Some("'x'".into()), Some("42".into()), None]).unwrap();
4362 assert_eq!(out, "WHERE a = 'x' AND b = 42 AND c = NULL");
4363 }
4364
4365 #[test]
4366 fn substitution_leaves_string_literals_alone() {
4367 let out = substitute_params("WHERE a = '$1' AND b = $1", &[Some("9".into())]).unwrap();
4368 assert_eq!(out, "WHERE a = '$1' AND b = 9");
4369 }
4370
4371 #[test]
4372 fn too_few_parameters_is_an_error_not_a_silent_null() {
4373 let e = substitute_params("WHERE a = $2", &[Some("1".into())]).unwrap_err();
4376 assert!(e.contains("$2"), "{}", e);
4377 }
4378
4379 #[test]
4380 fn a_quote_in_a_parameter_cannot_escape_its_literal() {
4381 let lit = decode_param(Some(b"it's"), OID_TEXT, 0).unwrap().unwrap();
4382 assert_eq!(lit, "'it''s'");
4383 let out = substitute_params("WHERE a = $1", &[Some(lit)]).unwrap();
4385 assert_eq!(out, "WHERE a = 'it''s'");
4386 }
4387
4388 #[test]
4389 fn binary_parameters_decode_in_every_width_psycopg_sends() {
4390 assert_eq!(decode_param(Some(&[0x00, 0x2a]), OID_INT2, 1).unwrap().unwrap(), "42");
4393 assert_eq!(decode_param(Some(&[0, 0, 0, 7]), OID_INT4, 1).unwrap().unwrap(), "7");
4394 assert_eq!(
4395 decode_param(Some(&[0, 0, 0, 0, 0, 0, 0, 9]), OID_INT8, 1).unwrap().unwrap(), "9");
4396 assert_eq!(
4397 decode_param(Some(&0x400c_0000_0000_0000u64.to_be_bytes()), OID_FLOAT8, 1)
4398 .unwrap().unwrap(), "3.5");
4399 assert_eq!(decode_param(Some(&[1]), OID_BOOL, 1).unwrap().unwrap(), "TRUE");
4400 assert_eq!(decode_param(Some(&[0]), OID_BOOL, 1).unwrap().unwrap(), "FALSE");
4401 }
4402
4403 #[test]
4404 fn a_negative_binary_integer_keeps_its_sign() {
4405 assert_eq!(decode_param(Some(&(-5i32).to_be_bytes()), OID_INT4, 1).unwrap().unwrap(), "-5");
4406 assert_eq!(decode_param(Some(&(-5i16).to_be_bytes()), OID_INT2, 1).unwrap().unwrap(), "-5");
4407 }
4408
4409 #[test]
4410 fn a_binary_parameter_of_the_wrong_width_is_refused() {
4411 let e = decode_param(Some(&[0x2a]), OID_INT4, 1).unwrap_err();
4414 assert!(e.contains("4 bytes"), "{}", e);
4415 }
4416
4417 #[test]
4418 fn an_unspecified_text_parameter_is_treated_as_a_string() {
4419 assert_eq!(decode_param(Some(b"hello"), 0, 0).unwrap().unwrap(), "'hello'");
4422 }
4423
4424 #[test]
4425 fn a_null_parameter_decodes_to_none_in_every_format() {
4426 assert_eq!(decode_param(None, OID_TEXT, 0).unwrap(), None);
4427 assert_eq!(decode_param(None, OID_INT8, 1).unwrap(), None);
4428 }
4429
4430 #[test]
4431 fn an_unsupported_binary_type_says_so_by_name() {
4432 let e = decode_param(Some(&[0u8; 8]), 1114, 1).unwrap_err();
4433 assert!(e.contains("1114"), "{}", e);
4434 assert!(e.contains("text"), "the error should point at the way out: {}", e);
4435 }
4436
4437 #[test]
4438 fn a_text_number_that_is_not_a_number_gets_quoted() {
4439 assert_eq!(decode_param(Some(b"oops"), OID_INT8, 0).unwrap().unwrap(), "'oops'");
4442 }
4443
4444 #[test]
4445 fn a_client_declared_type_is_believed_over_inference() {
4446 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[OID_INT4, 0], None);
4449 assert_eq!(oids, vec![OID_INT4, OID_TEXT]);
4450 }
4451
4452 #[test]
4453 fn parameter_arity_is_taken_from_the_sql_when_the_client_declares_none() {
4454 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[], None);
4457 assert_eq!(oids.len(), 2);
4458 }
4459
4460 #[test]
4461 fn the_field_behind_each_placeholder_is_identified() {
4462 assert_eq!(
4463 param_fields("SELECT a FROM t WHERE qty > $1 AND status = $2", 2),
4464 vec![Some("qty".to_string()), Some("status".to_string())]);
4465 }
4466
4467 #[test]
4468 fn word_operators_do_not_hide_the_field() {
4469 assert_eq!(param_fields("SELECT a FROM t WHERE name LIKE $1", 1),
4470 vec![Some("name".to_string())]);
4471 assert_eq!(param_fields("SELECT a FROM t WHERE qty BETWEEN $1 AND $2", 2),
4472 vec![Some("qty".to_string()), Some("qty".to_string())]);
4473 assert_eq!(param_fields("SELECT a FROM t WHERE region IN ($1, $2)", 2),
4474 vec![Some("region".to_string()), Some("region".to_string())]);
4475 }
4476
4477 #[test]
4478 fn a_clause_position_types_from_the_grammar_not_from_a_column() {
4479 assert_eq!(
4483 infer_param_oids("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", &[], None),
4484 vec![OID_INT8, OID_TEXT]);
4485 assert_eq!(infer_param_oids("SELECT a FROM t AS OF $1", &[], None), vec![OID_INT8]);
4486 assert_eq!(
4489 infer_param_oids("SELECT a FROM t VALID AS OF $1", &[], None), vec![OID_TEXT]);
4490 assert_eq!(
4491 infer_param_oids("SELECT a FROM t LIMIT $1 OFFSET $2", &[], None),
4492 vec![OID_INT8, OID_INT8]);
4493 }
4494
4495 #[test]
4496 fn an_aggregate_column_types_from_what_the_aggregate_means() {
4497 assert_eq!(aggregate_oid("count", None, "t"), Some(OID_INT8));
4501 assert_eq!(aggregate_oid("avg_fee", None, "t"), Some(OID_FLOAT8),
4502 "an average is fractional even over integers");
4503 assert_eq!(aggregate_oid("max__seq", None, "t"), Some(OID_INT8));
4506 assert_eq!(aggregate_oid("total", None, "t"), None, "not an aggregate");
4507 }
4508
4509 #[test]
4510 fn the_parse_probe_uses_a_literal_that_every_clause_accepts() {
4511 let probe = probe_sql("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", 2);
4515 assert!(!probe.contains("NULL"), "{}", probe);
4516 assert!(translate(&probe).is_ok(), "the probe must parse: {}", probe);
4517 }
4518
4519 #[test]
4520 fn a_column_with_mixed_types_across_documents_is_advertised_as_text() {
4521 let rows = vec![json!({"x": 3}), json!({"x": "n/a"})];
4525 assert_eq!(oid_for(&rows, "x"), OID_TEXT);
4526 let rows = vec![json!({"x": 3}), json!({"x": 1.5})];
4528 assert_eq!(oid_for(&rows, "x"), OID_FLOAT8);
4529 let rows = vec![json!({"x": Value::Null}), json!({"x": 7})];
4531 assert_eq!(oid_for(&rows, "x"), OID_INT8);
4532 }
4533
4534 #[test]
4535 fn binary_output_encodes_each_advertised_type() {
4536 assert_eq!(cell_binary(Some(&json!(true)), OID_BOOL).unwrap().unwrap(), vec![1]);
4537 assert_eq!(cell_binary(Some(&json!(42)), OID_INT8).unwrap().unwrap(),
4538 42i64.to_be_bytes().to_vec());
4539 assert_eq!(cell_binary(Some(&json!(3.5)), OID_FLOAT8).unwrap().unwrap(),
4540 3.5f64.to_be_bytes().to_vec());
4541 assert_eq!(cell_binary(Some(&json!("hi")), OID_TEXT).unwrap().unwrap(), b"hi".to_vec());
4543 assert_eq!(cell_binary(Some(&Value::Null), OID_INT8).unwrap(), None);
4544 assert_eq!(cell(Some(&json!(true))).unwrap(), "t");
4546 }
4547
4548 #[test]
4549 fn a_value_that_does_not_fit_its_advertised_binary_type_is_refused() {
4550 let e = cell_binary(Some(&json!("nope")), OID_INT8).unwrap_err();
4555 assert!(e.contains("a string"), "{}", e);
4556 assert!(e.contains("more than one type"), "the error should explain WHY: {}", e);
4557 }
4558
4559 #[test]
4560 fn a_row_description_carries_the_requested_format_per_column() {
4561 let cols = [Col::same("a"), Col::same("b")];
4562 let m = row_description_fmt(&cols, &[OID_INT8, OID_TEXT], &[1, 0]);
4563 assert_eq!(m[0], b'T');
4564 assert_eq!(m[m.len() - 1], 0, "the last column was requested as text");
4566 }
4567
4568 #[test]
4569 fn a_qualified_column_resolves_to_its_bare_name() {
4570 assert_eq!(param_fields("SELECT a FROM t WHERE t.qty = $1", 1),
4571 vec![Some("qty".to_string())]);
4572 }
4573
4574 #[test]
4575 fn insert_placeholders_map_positionally_to_the_column_list() {
4576 assert_eq!(
4577 param_fields("INSERT INTO t (_id, qty, status) VALUES ($1, $2, $3)", 3),
4578 vec![Some("_id".to_string()), Some("qty".to_string()), Some("status".to_string())]);
4579 }
4580
4581 #[test]
4582 fn a_set_clause_placeholder_finds_its_column() {
4583 assert_eq!(param_fields("UPDATE t SET status = $1 WHERE _id = $2", 2),
4584 vec![Some("status".to_string()), Some("_id".to_string())]);
4585 }
4586
4587 #[test]
4588 fn the_target_collection_is_found_for_every_statement_kind() {
4589 assert_eq!(stmt_collection("SELECT a FROM inv WHERE b = $1"), "inv");
4590 assert_eq!(stmt_collection("UPDATE inv SET a = $1"), "inv");
4591 assert_eq!(stmt_collection("DELETE FROM inv WHERE a = $1"), "inv");
4592 assert_eq!(stmt_collection("INSERT INTO inv (a) VALUES ($1)"), "inv");
4593 assert_eq!(stmt_collection("SELECT a FROM public.inv"), "inv");
4595 assert_eq!(stmt_collection("INSERT INTO inv(a) VALUES ($1)"), "inv");
4596 }
4597
4598 #[test]
4599 fn engine_metadata_fields_type_without_touching_storage() {
4600 assert_eq!(infer_field_oid(None, "t", "_seq"), OID_INT8);
4601 assert_eq!(infer_field_oid(None, "t", "_id"), OID_TEXT);
4602 }
4603
4604 #[test]
4605 fn the_protocol_acknowledgements_are_single_empty_messages() {
4606 for (m, tag) in [
4608 (parse_complete(), b'1'), (bind_complete(), b'2'),
4609 (close_complete(), b'3'), (no_data(), b'n'), (portal_suspended(), b's'),
4610 ] {
4611 assert_eq!(m.len(), 5, "{:?}", tag as char);
4612 assert_eq!(m[0], tag);
4613 assert_eq!(i32::from_be_bytes([m[1], m[2], m[3], m[4]]), 4);
4614 }
4615 }
4616
4617 #[test]
4618 fn parameter_description_reports_its_arity_and_types() {
4619 let m = parameter_description(&[OID_TEXT, OID_INT8]);
4620 assert_eq!(m[0], b't');
4621 assert_eq!(i16::from_be_bytes([m[5], m[6]]), 2);
4622 assert_eq!(i32::from_be_bytes([m[7], m[8], m[9], m[10]]), OID_TEXT);
4623 assert_eq!(i32::from_be_bytes([m[11], m[12], m[13], m[14]]), OID_INT8);
4624 }
4625
4626 #[test]
4627 fn a_cstring_is_taken_without_its_terminator() {
4628 let body = b"one\0two\0".to_vec();
4629 let mut at = 0usize;
4630 assert_eq!(take_cstr(&body, &mut at), "one");
4631 assert_eq!(take_cstr(&body, &mut at), "two");
4632 assert_eq!(at, body.len());
4633 }
4634
4635 #[test]
4636 fn truncated_integers_are_reported_rather_than_read_past_the_end() {
4637 let body = vec![0u8, 1];
4638 let mut at = 0usize;
4639 assert!(take_i32(&body, &mut at).is_err());
4640 let mut at = 0usize;
4641 assert!(take_i16(&body, &mut at).is_ok());
4642 }
4643
4644 #[test]
4645 fn a_binary_result_format_request_is_refused_rather_than_faked() {
4646 let out = encode_rows(&[], &[Col::same("a")]);
4649 let desc_format = &out[out.len() - 2..];
4650 assert_eq!(i16::from_be_bytes([desc_format[0], desc_format[1]]), 0,
4651 "every column is advertised as text format");
4652 }
4653
4654 #[test]
4655 fn a_float_parameter_does_not_render_as_rust_infinity() {
4656 assert_eq!(fmt_float(f64::INFINITY), "'Infinity'");
4657 assert_eq!(fmt_float(f64::NEG_INFINITY), "'-Infinity'");
4658 assert_eq!(fmt_float(f64::NAN), "'NaN'");
4659 assert_eq!(fmt_float(3.0), "3", "a whole float should not gain a .0 tail");
4660 assert_eq!(fmt_float(3.5), "3.5");
4661 }
4662}