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 let bare = expr.rsplit('.').next().unwrap_or(expr).trim_matches('"');
1213 let is_column = !bare.is_empty()
1214 && !bare.starts_with(|c: char| c.is_ascii_digit())
1215 && bare.chars().all(|c| c.is_alphanumeric() || c == '_' || c == '$');
1216 if !is_column {
1217 return Err(format!(
1218 "expressions in the select list are not supported ({:?}) — \
1219 supported: *, a column list, COUNT(*), or SUM/AVG/MIN/MAX(col). \
1220 Compute it in your client, or read the column and map it there",
1221 p));
1222 }
1223 let name = bare;
1224 project.push(Col::renamed(name, alias.unwrap_or(name)));
1225 }
1226 }
1227
1228 let (alias, tail) = split_table_alias(tail);
1237 let mut tail = strip_column_qualifiers(tail, coll, alias.as_deref())?;
1238 let tu = tail.to_uppercase();
1239 if let Some(at) = find_kw(&tu, "AS OF SYSTEM TIME") {
1240 let before = tail[..at].to_string();
1241 let after = tail[at + "AS OF SYSTEM TIME".len()..].trim_start().to_string();
1242 let end = after.find(' ').unwrap_or(after.len());
1244 let seq = after[..end].trim().trim_matches('\'').trim_matches('"').to_string();
1245 if seq.parse::<u64>().is_err() {
1246 return Err(format!(
1247 "AS OF SYSTEM TIME takes a NEDB sequence number here, not a timestamp (got {:?}). \
1248 NEDB's history is sequence-addressed and never garbage-collected, so a seq is \
1249 exact where a wall-clock time would be approximate", seq));
1250 }
1251 tail = format!("{} AS OF {} {}", before.trim(), seq, after[end..].trim())
1252 .trim()
1253 .to_string();
1254 }
1255
1256 let tu_ord = tail.to_uppercase();
1269 if let Some(ob_at) = find_kw(&tu_ord, "ORDER BY") {
1270 let start = ob_at + "ORDER BY".len();
1271 let end = ["LIMIT", "OFFSET", "GROUP BY", "TRACE", "TRAVERSE", "SEARCH"]
1273 .iter()
1274 .filter_map(|k| find_kw(&tu_ord[start..], k).map(|at| start + at))
1275 .min()
1276 .unwrap_or(tail.len());
1277 let mut keys = vec![];
1278 for item in split_top_level(&tail[start..end], ',') {
1279 let item = item.trim();
1280 if item.is_empty() {
1281 continue;
1282 }
1283 let mut parts = item.split_whitespace();
1284 let first = parts.next().unwrap_or("");
1285 let rest: Vec<&str> = parts.collect();
1286 match first.parse::<usize>() {
1287 Ok(n) if n >= 1 => {
1288 let col = project.get(n - 1).ok_or_else(|| {
1289 if project.is_empty() {
1290 format!(
1291 "ORDER BY {} is a select-list POSITION, and `SELECT *` \
1292 has no list to index — name the column instead", n)
1293 } else {
1294 format!(
1295 "ORDER BY {} is out of range: the select list has {} \
1296 column(s)", n, project.len())
1297 }
1298 })?;
1299 keys.push(
1300 std::iter::once(col.src.as_str())
1301 .chain(rest.iter().copied())
1302 .collect::<Vec<_>>()
1303 .join(" "),
1304 );
1305 }
1306 _ => keys.push(item.to_string()),
1309 }
1310 }
1311 tail = format!("{} ORDER BY {} {}", &tail[..ob_at], keys.join(", "), &tail[end..])
1312 .split_whitespace()
1313 .collect::<Vec<_>>()
1314 .join(" ");
1315 }
1316
1317 let mut gkey: Option<String> = None;
1325 let tu_all = tail.to_uppercase();
1326 if let Some(gb_at) = find_kw(&tu_all, "GROUP BY") {
1327 let head = tail[..gb_at].trim_end().to_string();
1328 let after = tail[gb_at + "GROUP BY".len()..].trim_start();
1329 let key_end = after.find(|c: char| c == ' ' || c == ',').unwrap_or(after.len());
1330 let group_key = after[..key_end].trim().trim_matches('"').to_string();
1331 let after_key = after[key_end..].trim_start();
1332 gkey = Some(group_key.clone());
1333
1334 if after_key.starts_with(',') {
1338 return Err(format!(
1339 "GROUP BY takes one key here (got {:?} and more) — NQL groups by a \
1340 single field, and grouping by only the first would aggregate over \
1341 rows the query meant to keep apart",
1342 group_key));
1343 }
1344
1345 for c in &project {
1346 let ok = c.src == group_key
1347 || c.src == "count"
1348 || agg_srcs.contains(&c.src);
1349 if !ok {
1350 return Err(format!(
1351 "column {:?} must appear in the GROUP BY clause or be used in an \
1352 aggregate function — a grouped row carries the group key, `count`, \
1353 and the aggregate, nothing else",
1354 c.src));
1355 }
1356 }
1357
1358 tail = format!("{} GROUP BY {}{} {}", head, group_key, agg_clause, after_key)
1369 .split_whitespace()
1370 .collect::<Vec<_>>()
1371 .join(" ");
1372 agg_clause.clear();
1373 }
1374
1375 let tu_hav = tail.to_uppercase();
1390 if let Some(h_at) = find_kw(&tu_hav, "HAVING") {
1391 let start = h_at + "HAVING".len();
1392 let end = ["ORDER BY", "LIMIT", "OFFSET"]
1393 .iter()
1394 .filter_map(|k| find_kw(&tu_hav[start..], k).map(|at| start + at))
1395 .min()
1396 .unwrap_or(tail.len());
1397 let clause = tail[start..end].to_string();
1398 let lhs_end = clause
1400 .find(|c: char| "<>=!".contains(c))
1401 .unwrap_or(clause.len());
1402 let lhs = clause[..lhs_end].trim();
1403 if !lhs.is_empty() {
1404 let lu = lhs.to_uppercase();
1405 let is_count = lu == "COUNT" || lu.replace(' ', "") == "COUNT(*)"
1412 || project.iter().any(|c| c.out.eq_ignore_ascii_case(lhs) && c.src == "count");
1413 let mapped = if is_count {
1414 Some("count".to_string())
1415 } else {
1416 agg_srcs.iter().find(|s| s.eq_ignore_ascii_case(lhs)).cloned().or_else(|| {
1418 project.iter()
1419 .find(|c| c.out.eq_ignore_ascii_case(lhs) && agg_srcs.contains(&c.src))
1420 .map(|c| c.src.clone())
1421 })
1422 };
1423 match mapped {
1424 Some(m) => {
1425 let rewritten = format!("{} {}", m, clause[lhs_end..].trim());
1429 tail = format!("{} HAVING {} {}",
1430 tail[..h_at].trim(), rewritten.trim(), tail[end..].trim())
1431 .trim().to_string();
1432 }
1433 None if gkey.as_deref().map(|g| g.eq_ignore_ascii_case(lhs)) == Some(true) => {}
1434 None => {
1435 return Err(format!(
1436 "HAVING names {:?}, which this grouped row does not carry. \
1437 It has the group key{}{}. Filtering on anything else would \
1438 answer zero rows rather than report a mistake",
1439 lhs,
1440 gkey.as_deref().map(|g| format!(" ({:?})", g)).unwrap_or_default(),
1441 if agg_srcs.is_empty() { String::new() }
1442 else { format!(", plus {}", agg_srcs.join(", ")) }));
1443 }
1444 }
1445 }
1446 }
1447
1448
1449 let tail = sql_literals_to_nql(&tail);
1450 let nql = format!("FROM {}{}{}", coll,
1451 if agg_clause.is_empty() { String::new() } else { agg_clause },
1452 if tail.is_empty() { String::new() } else { format!(" {}", tail) });
1453
1454 Ok(Stmt::Query { nql: nql.trim().to_string(), project })
1455}
1456
1457const SERVER_VERSION: &str = "15.0";
1458
1459pub fn version_string() -> String {
1461 full_version_string()
1462}
1463
1464fn full_version_string() -> String {
1465 format!(
1466 "PostgreSQL {} (NEDB {}) — tamper-evident, append-only, permanent \
1467 history. SELECT + INSERT/UPDATE/DELETE; an UPDATE is a new version, \
1468 so prior values stay readable with AS OF SYSTEM TIME.",
1469 SERVER_VERSION,
1470 env!("CARGO_PKG_VERSION")
1471 )
1472}
1473
1474fn columns_for(rows: &[Value], project: &[Col]) -> Vec<Col> {
1482 if !project.is_empty() {
1483 return project.to_vec();
1484 }
1485 let mut plain: Vec<String> = vec![];
1486 let mut meta: Vec<String> = vec![];
1487 for r in rows {
1488 if let Value::Object(m) = r {
1489 for k in m.keys() {
1490 let target = if k.starts_with('_') { &mut meta } else { &mut plain };
1491 if !target.contains(k) {
1492 target.push(k.clone());
1493 }
1494 }
1495 }
1496 }
1497 plain.sort();
1498 meta.sort();
1499 plain.extend(meta);
1500 plain.into_iter().map(|k| Col::same(&k)).collect()
1501}
1502
1503fn oid_of_value(v: &Value) -> Option<i32> {
1505 match v {
1506 Value::Null => None,
1507 Value::Bool(_) => Some(OID_BOOL),
1508 Value::Number(n) => Some(if n.is_i64() || n.is_u64() { OID_INT8 } else { OID_FLOAT8 }),
1509 Value::String(_) => Some(OID_TEXT),
1510 _ => Some(OID_TEXT),
1512 }
1513}
1514
1515fn unify_oid(a: i32, b: i32) -> i32 {
1522 if a == b {
1523 return a;
1524 }
1525 match (a, b) {
1526 (OID_INT8, OID_FLOAT8) | (OID_FLOAT8, OID_INT8) => OID_FLOAT8,
1527 _ => OID_TEXT,
1528 }
1529}
1530
1531pub fn oid_for_column(rows: &[Value], col: &str) -> i32 {
1543 oid_for(rows, col)
1544}
1545
1546fn oid_for(rows: &[Value], col: &str) -> i32 {
1547 let mut acc: Option<i32> = None;
1548 for r in rows {
1549 if let Some(o) = r.get(col).and_then(oid_of_value) {
1550 acc = Some(match acc {
1551 None => o,
1552 Some(prev) => unify_oid(prev, o),
1553 });
1554 if acc == Some(OID_TEXT) {
1555 break; }
1557 }
1558 }
1559 acc.unwrap_or(OID_TEXT)
1560}
1561
1562fn cell(v: Option<&Value>) -> Option<String> {
1564 match v {
1565 None | Some(Value::Null) => None, Some(Value::String(s)) => Some(s.clone()),
1567 Some(Value::Bool(b)) => Some(if *b { "t".into() } else { "f".into() }),
1568 Some(other) => Some(other.to_string()),
1569 }
1570}
1571
1572fn cell_binary(v: Option<&Value>, oid: i32) -> Result<Option<Vec<u8>>, String> {
1584 let v = match v {
1585 None | Some(Value::Null) => return Ok(None),
1586 Some(v) => v,
1587 };
1588 let as_f64 = |n: &serde_json::Number| n.as_f64()
1589 .ok_or_else(|| "a number too large to send as float8".to_string());
1590 Ok(Some(match (oid, v) {
1591 (OID_BOOL, Value::Bool(b)) => vec![u8::from(*b)],
1592 (OID_INT2, Value::Number(n)) => {
1593 let i = n.as_i64().ok_or("not an integer")?;
1594 i16::try_from(i).map_err(|_| format!("{} does not fit in int2", i))?
1595 .to_be_bytes().to_vec()
1596 }
1597 (OID_INT4, Value::Number(n)) => {
1598 let i = n.as_i64().ok_or("not an integer")?;
1599 i32::try_from(i).map_err(|_| format!("{} does not fit in int4", i))?
1600 .to_be_bytes().to_vec()
1601 }
1602 (OID_INT8, Value::Number(n)) => {
1603 n.as_i64().ok_or("not an integer")?.to_be_bytes().to_vec()
1604 }
1605 (OID_FLOAT4, Value::Number(n)) => (as_f64(n)? as f32).to_be_bytes().to_vec(),
1606 (OID_FLOAT8, Value::Number(n)) => as_f64(n)?.to_be_bytes().to_vec(),
1607 (OID_TEXT | OID_VARCHAR | OID_NAME | OID_UNKNOWN | OID_JSON, _) => {
1609 cell(Some(v)).unwrap_or_default().into_bytes()
1610 }
1611 (OID_JSONB, _) => {
1613 let mut b = vec![1u8];
1614 b.extend_from_slice(cell(Some(v)).unwrap_or_default().as_bytes());
1615 b
1616 }
1617 (oid, val) => {
1618 let kind = match val {
1619 Value::Bool(_) => "a boolean",
1620 Value::Number(_) => "a number",
1621 Value::String(_) => "a string",
1622 Value::Array(_) => "an array",
1623 _ => "an object",
1624 };
1625 return Err(format!(
1626 "cannot send {} in binary format as type OID {} — the field holds \
1627 more than one type across documents, so it cannot be described \
1628 by a single Postgres type. Select it with a text cast, or use a \
1629 text-format client",
1630 kind, oid
1631 ));
1632 }
1633 }))
1634}
1635
1636fn row_description_fmt(cols: &[Col], oids: &[i32], fmts: &[i16]) -> Vec<u8> {
1638 let mut m = Out::msg(b'T');
1639 m.i16(cols.len() as i16);
1640 for (i, c) in cols.iter().enumerate() {
1641 m.cstr(&c.out);
1642 m.i32(0); m.i16((i + 1) as i16); m.i32(oids.get(i).copied().unwrap_or(OID_TEXT));
1645 m.i16(-1); m.i32(-1); m.i16(fmts.get(i).copied().unwrap_or(0));
1648 }
1649 m.finish()
1650}
1651
1652fn row_description(cols: &[Col], oids: &[i32]) -> Vec<u8> {
1653 row_description_fmt(cols, oids, &[])
1654}
1655
1656fn data_row_bytes(vals: &[Option<Vec<u8>>]) -> Vec<u8> {
1657 let mut m = Out::msg(b'D');
1658 m.i16(vals.len() as i16);
1659 for v in vals {
1660 match v {
1661 None => m.i32(-1),
1662 Some(b) => {
1663 m.i32(b.len() as i32);
1664 m.bytes(b);
1665 }
1666 }
1667 }
1668 m.finish()
1669}
1670
1671fn data_row(vals: &[Option<String>]) -> Vec<u8> {
1672 let owned: Vec<Option<Vec<u8>>> =
1673 vals.iter().map(|v| v.as_ref().map(|s| s.as_bytes().to_vec())).collect();
1674 data_row_bytes(&owned)
1675}
1676
1677pub fn encode_rows(rows: &[Value], project: &[Col]) -> Vec<u8> {
1687 let cols = columns_for(rows, project);
1688 let oids: Vec<i32> = cols.iter().map(|c| oid_for(rows, &c.src)).collect();
1689 let mut out = row_description(&cols, &oids);
1690 for r in rows {
1691 let vals: Vec<Option<String>> = cols.iter().map(|c| cell(r.get(&c.src))).collect();
1692 out.extend_from_slice(&data_row(&vals));
1693 }
1694 out
1695}
1696
1697pub fn encode_result(rows: &[Value], project: &[Col]) -> Vec<u8> {
1699 let mut out = encode_rows(rows, project);
1700 out.extend_from_slice(&command_complete(&format!("SELECT {}", rows.len())));
1701 out
1702}
1703
1704const OID_INT2: i32 = 21;
1731const OID_INT4: i32 = 23;
1732const OID_OID: i32 = 26;
1733const OID_FLOAT4: i32 = 700;
1734const OID_VARCHAR: i32 = 1043;
1735const OID_NAME: i32 = 19;
1736const OID_UNKNOWN: i32 = 705;
1737const OID_JSON: i32 = 114;
1738const OID_JSONB: i32 = 3802;
1739
1740fn param_count(sql: &str) -> usize {
1746 let b = sql.as_bytes();
1747 let mut i = 0usize;
1748 let mut in_s = false;
1749 let mut max = 0usize;
1750 while i < b.len() {
1751 let c = b[i];
1752 if in_s {
1753 if c == b'\'' {
1754 in_s = false;
1755 }
1756 i += 1;
1757 continue;
1758 }
1759 if c == b'\'' {
1760 in_s = true;
1761 i += 1;
1762 continue;
1763 }
1764 if c == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1765 let mut j = i + 1;
1766 let mut n = 0usize;
1767 while j < b.len() && b[j].is_ascii_digit() {
1768 n = n * 10 + (b[j] - b'0') as usize;
1769 j += 1;
1770 }
1771 max = max.max(n);
1772 i = j;
1773 continue;
1774 }
1775 i += 1;
1776 }
1777 max
1778}
1779
1780fn infer_field_oid(db: Option<&Arc<Db>>, coll: &str, field: &str) -> i32 {
1789 match field {
1793 "_seq" => return OID_INT8,
1794 "_id" | "_hash" | "_prev" | "_collection" | "_valid_from" | "_valid_to" => return OID_TEXT,
1795 _ => {}
1796 }
1797 if !field.is_empty() && crate::pgcatalog::is_catalog(coll) {
1803 if let Some(rows) = crate::pgcatalog::rows(coll, db) {
1804 return oid_for(&rows, field);
1805 }
1806 }
1807 let db = match db {
1808 Some(db) => db,
1809 None => return OID_TEXT,
1810 };
1811 if coll.is_empty() || field.is_empty() {
1812 return OID_TEXT;
1813 }
1814 let rows = match crate::nql::query(db, &format!("FROM {} LIMIT {}", coll, TYPE_SAMPLE)) {
1815 Ok((rows, _)) => rows,
1816 Err(_) => return OID_TEXT,
1817 };
1818 oid_for(&rows, field)
1822}
1823
1824fn aggregate_oid(src: &str, db: Option<&Arc<Db>>, coll: &str) -> Option<i32> {
1836 if src == "count" {
1837 return Some(OID_INT8);
1838 }
1839 for (prefix, fixed) in [
1840 ("count_", Some(OID_INT8)),
1841 ("avg_", Some(OID_FLOAT8)),
1842 ("sum_", None),
1843 ("min_", None),
1844 ("max_", None),
1845 ] {
1846 if let Some(field) = src.strip_prefix(prefix) {
1847 return Some(match fixed {
1848 Some(oid) => oid,
1849 None => match infer_field_oid(db, coll, field) {
1852 OID_INT8 => OID_INT8,
1853 OID_FLOAT8 => OID_FLOAT8,
1854 other => other,
1857 },
1858 });
1859 }
1860 }
1861 None
1862}
1863
1864const TYPE_SAMPLE: usize = 200;
1870
1871fn stmt_collection(sql: &str) -> String {
1873 let s = normalise(sql);
1874 let up = s.to_uppercase();
1875 let after = if let Some(at) = find_kw(&up, "FROM") {
1876 &s[at + 4..]
1877 } else if let Some(rest) = strip_prefix_ci(&s, "UPDATE") {
1878 return rest
1879 .split_whitespace()
1880 .next()
1881 .unwrap_or("")
1882 .rsplit('.')
1883 .next()
1884 .unwrap_or("")
1885 .trim_matches('"')
1886 .to_string();
1887 } else if let Some(rest) = strip_prefix_ci(&s, "INSERT INTO") {
1888 return rest
1889 .split(|c: char| c.is_whitespace() || c == '(')
1890 .find(|t| !t.is_empty())
1891 .unwrap_or("")
1892 .rsplit('.')
1893 .next()
1894 .unwrap_or("")
1895 .trim_matches('"')
1896 .to_string();
1897 } else {
1898 return String::new();
1899 };
1900 after
1901 .trim()
1902 .split(|c: char| c.is_whitespace())
1903 .find(|t| !t.is_empty())
1904 .unwrap_or("")
1905 .rsplit('.')
1906 .next()
1907 .unwrap_or("")
1908 .trim_matches('"')
1909 .to_string()
1910}
1911
1912fn param_fields(sql: &str, n_params: usize) -> Vec<Option<String>> {
1923 let s = normalise(sql);
1924 let mut out = vec![None; n_params];
1925
1926 let up = s.to_uppercase();
1929 if up.starts_with("INSERT") {
1930 if let (Some(open), Some(vals_at)) = (s.find('('), find_kw(&up, "VALUES")) {
1931 if open < vals_at {
1932 if let Some(close) = s[open..vals_at].rfind(')') {
1933 let cols: Vec<String> = split_top(&s[open + 1..open + close], ',')
1934 .into_iter()
1935 .map(|c| c.trim().trim_matches('"').to_string())
1936 .collect();
1937 let tail = &s[vals_at..];
1939 let mut seen = 0usize;
1940 let b = tail.as_bytes();
1941 let mut i = 0usize;
1942 let mut in_s = false;
1943 while i < b.len() {
1944 if in_s {
1945 if b[i] == b'\'' { in_s = false; }
1946 i += 1;
1947 continue;
1948 }
1949 if b[i] == b'\'' { in_s = true; i += 1; continue; }
1950 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1951 let mut j = i + 1;
1952 let mut num = 0usize;
1953 while j < b.len() && b[j].is_ascii_digit() {
1954 num = num * 10 + (b[j] - b'0') as usize;
1955 j += 1;
1956 }
1957 if num >= 1 && num <= n_params {
1958 if let Some(c) = cols.get(seen % cols.len().max(1)) {
1959 out[num - 1] = Some(c.clone());
1960 }
1961 }
1962 seen += 1;
1963 i = j;
1964 continue;
1965 }
1966 i += 1;
1967 }
1968 return out;
1969 }
1970 }
1971 }
1972 }
1973
1974 let b = s.as_bytes();
1976 let mut i = 0usize;
1977 let mut in_s = false;
1978 while i < b.len() {
1979 if in_s {
1980 if b[i] == b'\'' { in_s = false; }
1981 i += 1;
1982 continue;
1983 }
1984 if b[i] == b'\'' { in_s = true; i += 1; continue; }
1985 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1986 let mut j = i + 1;
1987 let mut num = 0usize;
1988 while j < b.len() && b[j].is_ascii_digit() {
1989 num = num * 10 + (b[j] - b'0') as usize;
1990 j += 1;
1991 }
1992 if num >= 1 && num <= n_params {
1993 let left = &s[..i];
1994 let trimmed = left.trim_end_matches(|c: char| {
1997 c.is_whitespace() || "=<>!+-*/%(,".contains(c)
1998 });
1999 let mut tok = trimmed
2002 .rsplit(|c: char| c.is_whitespace() || c == '(' || c == ',')
2003 .find(|t| !t.is_empty())
2004 .unwrap_or("")
2005 .trim_matches('"');
2006 let mut before = trimmed;
2007 for _ in 0..4 {
2008 let upper_tok = tok.to_uppercase();
2009 if upper_tok.starts_with('$')
2015 || matches!(upper_tok.as_str(),
2016 "LIKE" | "ILIKE" | "IN" | "BETWEEN" | "AND" | "OR" | "NOT" | "IS") {
2017 before = before[..before.len() - tok.len()].trim_end_matches(|c: char| {
2018 c.is_whitespace() || "=<>!(,".contains(c)
2019 });
2020 tok = before
2021 .rsplit(|c: char| c.is_whitespace() || c == '(' || c == ',')
2022 .find(|t| !t.is_empty())
2023 .unwrap_or("")
2024 .trim_matches('"');
2025 } else {
2026 break;
2027 }
2028 }
2029 if !tok.is_empty()
2030 && tok.chars().all(|c| c.is_alphanumeric() || c == '_' || c == '.')
2031 && !tok.chars().next().map(|c| c.is_ascii_digit()).unwrap_or(true)
2032 {
2033 out[num - 1] = Some(tok.rsplit('.').next().unwrap_or(tok).to_string());
2034 }
2035 }
2036 i = j;
2037 continue;
2038 }
2039 i += 1;
2040 }
2041 out
2042}
2043
2044fn clause_param_oids(sql: &str, n_params: usize) -> Vec<Option<i32>> {
2054 let s = normalise(sql);
2055 let mut out = vec![None; n_params];
2056 let b = s.as_bytes();
2057 let mut i = 0usize;
2058 let mut in_s = false;
2059 while i < b.len() {
2060 if in_s {
2061 if b[i] == b'\'' { in_s = false; }
2062 i += 1;
2063 continue;
2064 }
2065 if b[i] == b'\'' { in_s = true; i += 1; continue; }
2066 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
2067 let mut j = i + 1;
2068 let mut num = 0usize;
2069 while j < b.len() && b[j].is_ascii_digit() {
2070 num = num * 10 + (b[j] - b'0') as usize;
2071 j += 1;
2072 }
2073 if num >= 1 && num <= n_params {
2074 let left = s[..i].trim_end().to_uppercase();
2075 out[num - 1] = if left.ends_with("VALID AS OF") {
2078 Some(OID_TEXT)
2079 } else if left.ends_with("AS OF SYSTEM TIME")
2080 || left.ends_with("FOR SYSTEM_TIME AS OF")
2081 || left.ends_with("AS OF")
2082 || left.ends_with("LIMIT")
2083 || left.ends_with("OFFSET")
2084 {
2085 Some(OID_INT8)
2086 } else {
2087 None
2088 };
2089 }
2090 i = j;
2091 continue;
2092 }
2093 i += 1;
2094 }
2095 out
2096}
2097
2098fn infer_param_oids(sql: &str, declared: &[i32], db: Option<&Arc<Db>>) -> Vec<i32> {
2104 let n = param_count(sql).max(declared.len());
2105 if n == 0 {
2106 return vec![];
2107 }
2108 let coll = stmt_collection(sql);
2109 let fields = param_fields(sql, n);
2110 let clauses = clause_param_oids(sql, n);
2111 (0..n)
2112 .map(|i| match declared.get(i) {
2113 Some(&oid) if oid != 0 => oid,
2114 _ => match clauses[i] {
2117 Some(oid) => oid,
2118 None => match &fields[i] {
2119 Some(f) => infer_field_oid(db, &coll, f),
2120 None => OID_TEXT,
2121 },
2122 },
2123 })
2124 .collect()
2125}
2126
2127fn decode_param(raw: Option<&[u8]>, oid: i32, format: i16) -> Result<Option<String>, String> {
2133 let bytes = match raw {
2134 None => return Ok(None),
2135 Some(b) => b,
2136 };
2137 let quote = |s: &str| format!("'{}'", s.replace('\'', "''"));
2138
2139 if format == 0 {
2140 let s = String::from_utf8_lossy(bytes).to_string();
2141 return Ok(Some(match oid {
2142 OID_BOOL => {
2143 let t = matches!(s.as_str(), "t" | "true" | "TRUE" | "1" | "yes" | "on");
2144 if t { "TRUE".into() } else { "FALSE".into() }
2145 }
2146 OID_INT2 | OID_INT4 | OID_INT8 | OID_OID | OID_FLOAT4 | OID_FLOAT8 => {
2147 if s.parse::<f64>().is_ok() { s } else { quote(&s) }
2151 }
2152 _ => quote(&s),
2157 }));
2158 }
2159 if format != 1 {
2160 return Err(format!("unsupported parameter format code {}", format));
2161 }
2162
2163 let need = |n: usize| -> Result<(), String> {
2165 if bytes.len() == n {
2166 Ok(())
2167 } else {
2168 Err(format!(
2169 "binary parameter of type OID {} should be {} bytes, got {}",
2170 oid, n, bytes.len()
2171 ))
2172 }
2173 };
2174 Ok(Some(match oid {
2175 OID_BOOL => {
2176 need(1)?;
2177 if bytes[0] != 0 { "TRUE".into() } else { "FALSE".into() }
2178 }
2179 OID_INT2 => {
2180 need(2)?;
2181 i16::from_be_bytes([bytes[0], bytes[1]]).to_string()
2182 }
2183 OID_INT4 => {
2184 need(4)?;
2185 i32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]).to_string()
2186 }
2187 OID_OID => {
2188 need(4)?;
2189 u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]).to_string()
2190 }
2191 OID_INT8 => {
2192 need(8)?;
2193 i64::from_be_bytes(bytes[..8].try_into().unwrap()).to_string()
2194 }
2195 OID_FLOAT4 => {
2196 need(4)?;
2197 let f = f32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]);
2198 fmt_float(f as f64)
2199 }
2200 OID_FLOAT8 => {
2201 need(8)?;
2202 fmt_float(f64::from_be_bytes(bytes[..8].try_into().unwrap()))
2203 }
2204 OID_TEXT | OID_VARCHAR | OID_NAME | OID_UNKNOWN | OID_JSON | 0 => {
2205 quote(&String::from_utf8_lossy(bytes))
2206 }
2207 OID_JSONB => {
2208 let body = if bytes.first() == Some(&1) { &bytes[1..] } else { bytes };
2210 quote(&String::from_utf8_lossy(body))
2211 }
2212 other => {
2213 return Err(format!(
2214 "parameter type OID {} is not supported in binary format — \
2215 the supported set is bool, int2/int4/int8, float4/float8, \
2216 text/varchar/json/jsonb. Send it as text, or cast it in the \
2217 statement",
2218 other
2219 ))
2220 }
2221 }))
2222}
2223
2224fn fmt_float(f: f64) -> String {
2226 if f.is_nan() {
2227 "'NaN'".into()
2228 } else if f.is_infinite() {
2229 if f > 0.0 { "'Infinity'".into() } else { "'-Infinity'".into() }
2230 } else if f.fract() == 0.0 && f.abs() < 1e15 {
2231 format!("{:.0}", f)
2232 } else {
2233 f.to_string()
2234 }
2235}
2236
2237fn substitute_params(sql: &str, params: &[Option<String>]) -> Result<String, String> {
2245 let b = sql.as_bytes();
2246 let mut out = String::with_capacity(sql.len() + 16);
2247 let mut i = 0usize;
2248 let mut in_s = false;
2249 while i < b.len() {
2250 let c = b[i];
2251 if in_s {
2252 out.push(c as char);
2253 if c == b'\'' { in_s = false; }
2254 i += 1;
2255 continue;
2256 }
2257 if c == b'\'' {
2258 in_s = true;
2259 out.push('\'');
2260 i += 1;
2261 continue;
2262 }
2263 if c == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
2264 let mut j = i + 1;
2265 let mut n = 0usize;
2266 while j < b.len() && b[j].is_ascii_digit() {
2267 n = n * 10 + (b[j] - b'0') as usize;
2268 j += 1;
2269 }
2270 match params.get(n.wrapping_sub(1)) {
2271 Some(Some(lit)) => out.push_str(lit),
2272 Some(None) => out.push_str("NULL"),
2273 None => {
2274 return Err(format!(
2275 "bind message supplies {} parameter(s) but the statement uses ${}",
2276 params.len(), n
2277 ))
2278 }
2279 }
2280 i = j;
2281 continue;
2282 }
2283 out.push(c as char);
2284 i += 1;
2285 }
2286 Ok(out)
2287}
2288
2289struct Prepared {
2291 sql: String,
2292 param_oids: Vec<i32>,
2295 out_shape: Option<Option<(Vec<Col>, Vec<i32>)>>,
2303}
2304
2305fn prepared_shape<'a>(
2307 p: &'a mut Prepared,
2308 db: Option<&Arc<Db>>,
2309) -> &'a Option<(Vec<Col>, Vec<i32>)> {
2310 if p.out_shape.is_none() {
2311 p.out_shape = Some(describe_shape(&p.sql, db, p.param_oids.len()));
2312 }
2313 p.out_shape.as_ref().expect("just filled")
2314}
2315
2316struct Portal {
2318 sql: String,
2319 result: Option<PortalResult>,
2325 frozen: Option<Vec<Col>>,
2335 formats: Vec<i16>,
2337 declared: Option<(Vec<Col>, Vec<i32>)>,
2344}
2345
2346impl Portal {
2347 fn format_of(&self, i: usize) -> i16 {
2350 match self.formats.len() {
2351 0 => 0,
2352 1 => self.formats[0],
2353 _ => self.formats.get(i).copied().unwrap_or(0),
2354 }
2355 }
2356 fn shape(&self, r: &PortalResult) -> (Vec<Col>, Vec<i32>) {
2358 match &self.declared {
2359 Some((cols, oids)) if self.formats.iter().any(|f| *f == 1) => {
2360 (cols.clone(), oids.clone())
2361 }
2362 _ => {
2363 let cols = columns_for(&r.rows, &r.project);
2364 let oids = cols.iter().map(|c| oid_for(&r.rows, &c.src)).collect();
2365 (cols, oids)
2366 }
2367 }
2368 }
2369}
2370
2371struct PortalResult {
2372 rows: Vec<Value>,
2373 project: Vec<Col>,
2374 has_rows: bool,
2375 tag: String,
2376 tag_counts_rows: bool,
2377 sent: usize,
2379}
2380
2381fn parse_complete() -> Vec<u8> { Out::msg(b'1').finish() }
2382fn bind_complete() -> Vec<u8> { Out::msg(b'2').finish() }
2383fn close_complete() -> Vec<u8> { Out::msg(b'3').finish() }
2384fn no_data() -> Vec<u8> { Out::msg(b'n').finish() }
2385fn portal_suspended() -> Vec<u8> { Out::msg(b's').finish() }
2386
2387fn parameter_description(oids: &[i32]) -> Vec<u8> {
2388 let mut m = Out::msg(b't');
2389 m.i16(oids.len() as i16);
2390 for o in oids {
2391 m.i32(*o);
2392 }
2393 m.finish()
2394}
2395
2396fn take_cstr(body: &[u8], at: &mut usize) -> String {
2398 let start = *at;
2399 while *at < body.len() && body[*at] != 0 {
2400 *at += 1;
2401 }
2402 let s = String::from_utf8_lossy(&body[start..*at]).to_string();
2403 if *at < body.len() {
2404 *at += 1; }
2406 s
2407}
2408
2409fn take_i16(body: &[u8], at: &mut usize) -> Result<i16, String> {
2410 if *at + 2 > body.len() {
2411 return Err("truncated message".into());
2412 }
2413 let v = i16::from_be_bytes([body[*at], body[*at + 1]]);
2414 *at += 2;
2415 Ok(v)
2416}
2417
2418fn take_i32(body: &[u8], at: &mut usize) -> Result<i32, String> {
2419 if *at + 4 > body.len() {
2420 return Err("truncated message".into());
2421 }
2422 let v = i32::from_be_bytes([body[*at], body[*at + 1], body[*at + 2], body[*at + 3]]);
2423 *at += 4;
2424 Ok(v)
2425}
2426
2427fn sample_columns(db: Option<&Arc<Db>>, coll: &str) -> Vec<Col> {
2433 let db = match db {
2434 Some(db) => db,
2435 None => return vec![],
2436 };
2437 let rows = match crate::nql::query(db, &format!("FROM {} LIMIT 25", coll)) {
2438 Ok((rows, _)) => rows,
2439 Err(_) => return vec![],
2440 };
2441 let mut names: Vec<String> = vec![];
2442 for r in &rows {
2443 if let Value::Object(m) = r {
2444 for k in m.keys() {
2445 if !names.iter().any(|n| n == k) {
2446 names.push(k.clone());
2447 }
2448 }
2449 }
2450 }
2451 names.sort();
2452 names.iter().map(|n| Col::same(n)).collect()
2453}
2454
2455fn describe_shape(
2463 sql: &str,
2464 db: Option<&Arc<Db>>,
2465 n_params: usize,
2466) -> Option<(Vec<Col>, Vec<i32>)> {
2467 let probe = probe_sql(sql, n_params);
2468
2469 if sql_engine_owns(&probe) {
2480 if let Ok(Some((done, _))) = try_catalog_select(&probe, db) {
2481 if done.project.is_empty() {
2482 return None;
2483 }
2484 let oids = done
2485 .project
2486 .iter()
2487 .map(|c| oid_for(&done.rows, &c.src))
2488 .collect();
2489 return Some((done.project, oids));
2490 }
2491 }
2492
2493 let stmt = translate(&probe).ok()?;
2494 let coll = stmt_collection(sql);
2495
2496 let cols = match stmt {
2497 Stmt::Ok(_) => return None,
2498 Stmt::Canned { cols, .. } => cols.iter().map(|c| Col::same(c)).collect(),
2499 Stmt::Query { project, .. } => {
2500 if project.is_empty() { sample_columns(db, &coll) } else { project }
2501 }
2502 Stmt::Insert { returning, .. } | Stmt::Update { returning, .. } | Stmt::Delete { returning, .. } => {
2503 if !wants_returning(sql) {
2504 return None;
2505 }
2506 if returning.is_empty() { sample_columns(db, &coll) } else { returning }
2507 }
2508 };
2509 if cols.is_empty() {
2510 return None;
2514 }
2515 let oids = cols
2516 .iter()
2517 .map(|c| {
2518 aggregate_oid(&c.src, db, &coll)
2519 .unwrap_or_else(|| infer_field_oid(db, &coll, &c.src))
2520 })
2521 .collect();
2522 Some((cols, oids))
2523}
2524
2525fn probe_sql(sql: &str, n_params: usize) -> String {
2533 let stub: Vec<Option<String>> = vec![Some("0".to_string()); n_params];
2534 substitute_params(sql, &stub).unwrap_or_else(|_| sql.to_string())
2535}
2536
2537fn ensure_executed(
2539 portal: &mut Portal,
2540 db_name: &str,
2541 db: Option<&Arc<Db>>,
2542 read_only: bool,
2543) -> Result<(), Vec<u8>> {
2544 if portal.result.is_some() {
2545 return Ok(());
2546 }
2547 let ex = execute_stmt(&portal.sql, db_name, db, read_only)?;
2548 let project = if let Some(f) = &portal.frozen {
2551 f.clone()
2552 } else {
2553 let p = if ex.project.is_empty() {
2554 columns_for(&ex.rows, &[])
2555 } else {
2556 ex.project.clone()
2557 };
2558 portal.frozen = Some(p.clone());
2559 p
2560 };
2561 portal.result = Some(PortalResult {
2562 rows: ex.rows,
2563 project,
2564 has_rows: ex.has_rows,
2565 tag: ex.tag,
2566 tag_counts_rows: ex.tag_counts_rows,
2567 sent: 0,
2568 });
2569 Ok(())
2570}
2571
2572async fn read_exact(sock: &mut TcpStream, n: usize) -> std::io::Result<Vec<u8>> {
2575 let mut buf = vec![0u8; n];
2576 sock.read_exact(&mut buf).await?;
2577 Ok(buf)
2578}
2579
2580async fn read_i32(sock: &mut TcpStream) -> std::io::Result<i32> {
2581 let b = read_exact(sock, 4).await?;
2582 Ok(i32::from_be_bytes([b[0], b[1], b[2], b[3]]))
2583}
2584
2585fn parse_startup_params(body: &[u8]) -> HashMap<String, String> {
2586 let mut out = HashMap::new();
2587 let mut parts = body.split(|b| *b == 0).map(|s| String::from_utf8_lossy(s).to_string());
2588 while let (Some(k), Some(v)) = (parts.next(), parts.next()) {
2589 if k.is_empty() {
2590 break;
2591 }
2592 out.insert(k, v);
2593 }
2594 out
2595}
2596
2597async fn handle(mut sock: TcpStream, resolver: Arc<dyn DbResolver>, read_only: bool) -> std::io::Result<()> {
2599 let params = loop {
2601 let len = read_i32(&mut sock).await?;
2602 if len < 8 || len > 1 << 20 {
2603 return Ok(()); }
2605 let code = read_i32(&mut sock).await?;
2606 let body = read_exact(&mut sock, (len - 8) as usize).await?;
2607 match code {
2608 SSL_REQUEST | GSS_REQUEST => {
2609 sock.write_all(b"N").await?;
2611 continue;
2612 }
2613 CANCEL_REQUEST => return Ok(()), PROTO_V3 => break parse_startup_params(&body),
2615 other => {
2616 let major = other >> 16;
2617 sock.write_all(&err_msg(
2618 "0A000",
2619 &format!("unsupported frontend protocol {}.{} — this endpoint speaks 3.0",
2620 major, other & 0xffff),
2621 )).await?;
2622 return Ok(());
2623 }
2624 }
2625 };
2626
2627 let db_name = params.get("database").cloned().unwrap_or_default();
2628
2629 let resolved: Option<Arc<Db>> = {
2635 let r = Arc::clone(&resolver);
2636 let name = db_name.clone();
2637 tokio::task::spawn_blocking(move || r.resolve(&name))
2638 .await
2639 .unwrap_or(None)
2640 };
2641
2642 if let Some(expected) = resolver.token() {
2644 let mut m = Out::msg(b'R');
2646 m.i32(3);
2647 sock.write_all(&m.finish()).await?;
2648
2649 let tag = read_exact(&mut sock, 1).await?;
2650 if tag[0] != b'p' {
2651 sock.write_all(&err_msg("28000", "expected a password message")).await?;
2652 return Ok(());
2653 }
2654 let len = read_i32(&mut sock).await?;
2655 if len < 4 || len > 1 << 16 {
2656 return Ok(());
2657 }
2658 let body = read_exact(&mut sock, (len - 4) as usize).await?;
2659 let supplied = String::from_utf8_lossy(&body).trim_end_matches('\0').to_string();
2660 let ok = supplied.len() == expected.len()
2662 && supplied.bytes().zip(expected.bytes()).fold(0u8, |a, (x, y)| a | (x ^ y)) == 0;
2663 if !ok {
2664 sock.write_all(&err_msg("28P01", "password authentication failed")).await?;
2665 return Ok(());
2666 }
2667 }
2668
2669 let mut m = Out::msg(b'R');
2670 m.i32(0); sock.write_all(&m.finish()).await?;
2672
2673 for (k, v) in [
2674 ("server_version", SERVER_VERSION),
2675 ("server_encoding", "UTF8"),
2676 ("client_encoding", "UTF8"),
2677 ("DateStyle", "ISO, MDY"),
2678 ("integer_datetimes", "on"),
2679 ("standard_conforming_strings", "on"),
2680 ("application_name", "nedbd"),
2681 ] {
2682 let mut p = Out::msg(b'S');
2683 p.cstr(k);
2684 p.cstr(v);
2685 sock.write_all(&p.finish()).await?;
2686 }
2687 let mut k = Out::msg(b'K');
2688 k.i32(std::process::id() as i32);
2689 k.i32(0);
2690 sock.write_all(&k.finish()).await?;
2691 sock.write_all(&ready()).await?;
2692
2693 let mut prepared: HashMap<String, Prepared> = HashMap::new();
2699 let mut portals: HashMap<String, Portal> = HashMap::new();
2700 let mut failed = false;
2704
2705 loop {
2706 let mut tag = [0u8; 1];
2707 if sock.read_exact(&mut tag).await.is_err() {
2708 return Ok(()); }
2710 let len = read_i32(&mut sock).await?;
2711 if len < 4 || len > 64 << 20 {
2712 return Ok(());
2713 }
2714 let body = read_exact(&mut sock, (len - 4) as usize).await?;
2715
2716 if failed && tag[0] != b'S' && tag[0] != b'X' {
2718 continue;
2719 }
2720
2721 match tag[0] {
2722 b'X' => return Ok(()), b'Q' => {
2725 let sql = String::from_utf8_lossy(&body).trim_end_matches('\0').to_string();
2726 let out = run_simple_query(&sql, &db_name, resolved.as_ref(), read_only);
2727 sock.write_all(&out).await?;
2728 sock.write_all(&ready()).await?;
2729 portals.remove("");
2731 }
2732
2733 b'P' => {
2735 let mut at = 0usize;
2736 let name = take_cstr(&body, &mut at);
2737 let sql = take_cstr(&body, &mut at);
2738 let n = take_i16(&body, &mut at).unwrap_or(0).max(0) as usize;
2739 let mut declared = Vec::with_capacity(n);
2740 let mut bad = false;
2741 for _ in 0..n {
2742 match take_i32(&body, &mut at) {
2743 Ok(o) => declared.push(o),
2744 Err(_) => { bad = true; break; }
2745 }
2746 }
2747 if bad {
2748 sock.write_all(&err_msg("08P01", "malformed Parse message")).await?;
2749 failed = true;
2750 continue;
2751 }
2752 let probe = probe_sql(&sql, param_count(&sql));
2762 if !sql_engine_owns(&probe) {
2763 if let Err(why) = translate(&probe) {
2764 sock.write_all(&err_msg("0A000", &why)).await?;
2765 failed = true;
2766 continue;
2767 }
2768 }
2769 let param_oids = infer_param_oids(&sql, &declared, resolved.as_ref());
2770 prepared.insert(name, Prepared { sql, param_oids, out_shape: None });
2771 sock.write_all(&parse_complete()).await?;
2772 }
2773
2774 b'B' => {
2776 let mut at = 0usize;
2777 let portal_name = take_cstr(&body, &mut at);
2778 let stmt_name = take_cstr(&body, &mut at);
2779 if !prepared.contains_key(&stmt_name) {
2780 sock.write_all(&err_msg("26000", &format!(
2781 "prepared statement {:?} does not exist", stmt_name))).await?;
2782 failed = true;
2783 continue;
2784 }
2785 let p = &prepared[&stmt_name];
2786 let mut want_formats: Vec<i16> = vec![];
2787 let res: Result<String, String> = (|| {
2788 let nfmt = take_i16(&body, &mut at)? .max(0) as usize;
2789 let mut fmts = Vec::with_capacity(nfmt);
2790 for _ in 0..nfmt {
2791 fmts.push(take_i16(&body, &mut at)?);
2792 }
2793 let nparam = take_i16(&body, &mut at)?.max(0) as usize;
2794 let mut vals: Vec<Option<String>> = Vec::with_capacity(nparam);
2795 for i in 0..nparam {
2796 let l = take_i32(&body, &mut at)?;
2797 let raw: Option<Vec<u8>> = if l < 0 {
2798 None
2799 } else {
2800 let l = l as usize;
2801 if at + l > body.len() {
2802 return Err("truncated Bind parameter".into());
2803 }
2804 let v = body[at..at + l].to_vec();
2805 at += l;
2806 Some(v)
2807 };
2808 let f = match fmts.len() {
2811 0 => 0,
2812 1 => fmts[0],
2813 _ => *fmts.get(i).unwrap_or(&0),
2814 };
2815 let oid = *p.param_oids.get(i).unwrap_or(&OID_TEXT);
2816 vals.push(decode_param(raw.as_deref(), oid, f)?);
2817 }
2818 let nres = take_i16(&body, &mut at)?.max(0) as usize;
2823 for _ in 0..nres {
2824 let f = take_i16(&body, &mut at)?;
2825 if f != 0 && f != 1 {
2826 return Err(format!("unknown result format code {}", f));
2827 }
2828 want_formats.push(f);
2829 }
2830 substitute_params(&p.sql, &vals)
2831 })();
2832 match res {
2833 Ok(sql) => {
2834 let declared = if want_formats.iter().any(|f| *f == 1) {
2837 let p = prepared.get_mut(&stmt_name).expect("checked above");
2838 prepared_shape(p, resolved.as_ref()).clone()
2839 } else {
2840 None
2841 };
2842 portals.insert(portal_name, Portal {
2843 sql, result: None, frozen: None,
2844 formats: want_formats, declared,
2845 });
2846 sock.write_all(&bind_complete()).await?;
2847 }
2848 Err(why) => {
2849 sock.write_all(&err_msg("08P01", &why)).await?;
2850 failed = true;
2851 }
2852 }
2853 }
2854
2855 b'D' => {
2857 let kind = body.first().copied().unwrap_or(b'S');
2858 let mut at = 1usize;
2859 let name = take_cstr(&body, &mut at);
2860 if kind == b'S' {
2861 if !prepared.contains_key(&name) {
2862 sock.write_all(&err_msg("26000", &format!(
2863 "prepared statement {:?} does not exist", name))).await?;
2864 failed = true;
2865 continue;
2866 }
2867 let p = prepared.get_mut(&name).expect("checked above");
2868 let oids = p.param_oids.clone();
2869 sock.write_all(¶meter_description(&oids)).await?;
2872 let out = match prepared_shape(p, resolved.as_ref()) {
2876 Some((cols, col_oids)) => row_description(cols, col_oids),
2877 None => no_data(),
2878 };
2879 sock.write_all(&out).await?;
2880 } else {
2881 let portal = match portals.get_mut(&name) {
2882 Some(p) => p,
2883 None => {
2884 sock.write_all(&err_msg("34000", &format!(
2885 "portal {:?} does not exist", name))).await?;
2886 failed = true;
2887 continue;
2888 }
2889 };
2890 match ensure_executed(portal, &db_name, resolved.as_ref(), read_only) {
2895 Err(encoded) => {
2896 sock.write_all(&encoded).await?;
2897 failed = true;
2898 }
2899 Ok(()) => {
2900 let r = portal.result.as_ref().expect("just executed");
2901 if !r.has_rows {
2902 sock.write_all(&no_data()).await?;
2903 } else {
2904 let (cols, oids) = portal.shape(r);
2905 let fmts: Vec<i16> =
2906 (0..cols.len()).map(|i| portal.format_of(i)).collect();
2907 sock.write_all(&row_description_fmt(&cols, &oids, &fmts)).await?;
2908 }
2909 }
2910 }
2911 }
2912 }
2913
2914 b'E' => {
2916 let mut at = 0usize;
2917 let name = take_cstr(&body, &mut at);
2918 let max_rows = take_i32(&body, &mut at).unwrap_or(0);
2919 let portal = match portals.get_mut(&name) {
2920 Some(p) => p,
2921 None => {
2922 sock.write_all(&err_msg("34000", &format!(
2923 "portal {:?} does not exist", name))).await?;
2924 failed = true;
2925 continue;
2926 }
2927 };
2928 if let Err(encoded) = ensure_executed(portal, &db_name, resolved.as_ref(), read_only) {
2929 sock.write_all(&encoded).await?;
2930 failed = true;
2931 continue;
2932 }
2933 let r = portal.result.as_ref().expect("just executed");
2934 if !r.has_rows {
2935 let tag = r.tag.clone();
2936 sock.write_all(&command_complete(&tag)).await?;
2937 continue;
2938 }
2939 let (cols, oids) = portal.shape(r);
2940 let limit = if max_rows > 0 {
2941 (r.sent + max_rows as usize).min(r.rows.len())
2942 } else {
2943 r.rows.len()
2944 };
2945 let mut encoded: Vec<Vec<u8>> = Vec::with_capacity(limit - r.sent);
2950 let mut fail: Option<String> = None;
2951 for row in &r.rows[r.sent..limit] {
2952 let mut vals: Vec<Option<Vec<u8>>> = Vec::with_capacity(cols.len());
2953 for (i, c) in cols.iter().enumerate() {
2954 let v = row.get(&c.src);
2955 let got = if portal.format_of(i) == 1 {
2956 cell_binary(v, oids.get(i).copied().unwrap_or(OID_TEXT))
2957 .map_err(|e| format!("column {:?}: {}", c.out, e))
2958 } else {
2959 Ok(cell(v).map(|s| s.into_bytes()))
2960 };
2961 match got {
2962 Ok(b) => vals.push(b),
2963 Err(e) => { fail = Some(e); break; }
2964 }
2965 }
2966 if fail.is_some() {
2967 break;
2968 }
2969 encoded.push(data_row_bytes(&vals));
2970 }
2971 if let Some(why) = fail {
2972 sock.write_all(&err_msg("22P03", &why)).await?;
2973 failed = true;
2974 continue;
2975 }
2976 let mut out = vec![];
2977 for e in &encoded {
2978 out.extend_from_slice(e);
2979 }
2980 let r = portal.result.as_mut().expect("just executed");
2981 r.sent = limit;
2982 if max_rows > 0 && r.sent < r.rows.len() {
2986 out.extend_from_slice(&portal_suspended());
2987 } else {
2988 let tag = if r.tag_counts_rows {
2989 format!("{} {}", r.tag, r.sent)
2990 } else {
2991 r.tag.clone()
2992 };
2993 out.extend_from_slice(&command_complete(&tag));
2994 }
2995 sock.write_all(&out).await?;
2996 }
2997
2998 b'C' => {
3000 let kind = body.first().copied().unwrap_or(b'S');
3001 let mut at = 1usize;
3002 let name = take_cstr(&body, &mut at);
3003 if kind == b'S' {
3004 prepared.remove(&name);
3005 } else {
3006 portals.remove(&name);
3007 }
3008 sock.write_all(&close_complete()).await?;
3011 }
3012
3013 b'H' => {}
3017
3018 b'S' => {
3019 failed = false;
3020 sock.write_all(&ready()).await?;
3021 }
3022
3023 other => {
3024 sock.write_all(&err_msg(
3025 "08P01",
3026 &format!("unexpected frontend message {:?}", other as char),
3027 )).await?;
3028 failed = true;
3029 }
3030 }
3031 }
3032}
3033
3034const READ_ONLY_MSG: &str =
3035 "this endpoint is running read-only (NEDBD_PG_READ_ONLY=1). Writes are \
3036 implemented but disabled on this server — unset the flag to allow them.";
3037
3038fn no_db(db_name: &str) -> Vec<u8> {
3039 err_msg("3D000", &format!(
3040 "database {:?} is not open on this server — create it first \
3041 (POST /v1/databases), or connect with -d <name>", db_name))
3042}
3043
3044fn catalog_name(n: &str) -> String {
3048 let joined: Vec<&str> = n.split('.').collect();
3049 if joined.len() >= 2 && joined[joined.len() - 2] == "information_schema" {
3050 format!("information_schema.{}", joined[joined.len() - 1])
3051 } else {
3052 joined[joined.len() - 1].to_string()
3053 }
3054}
3055
3056fn sql_engine_owns(sql: &str) -> bool {
3079 let Ok(sel) = crate::sqlselect::parse(sql) else { return false };
3080 let touched = sel.base_relations();
3081 if touched.is_empty() {
3082 return translate(sql).is_err();
3083 }
3084 touched.iter().any(|t| crate::pgcatalog::is_catalog(&catalog_name(t)))
3085}
3086
3087fn try_catalog_select(
3100 sql: &str,
3101 db: Option<&Arc<Db>>,
3102) -> Result<Option<(Executed, crate::sqlplan::Plan)>, Vec<u8>> {
3103 let sel = match crate::sqlselect::parse(sql) {
3104 Ok(sel) => sel,
3105 Err(why) => {
3106 if mentions_catalog(sql) {
3115 return Err(err_msg("0A000", &format!(
3116 "this catalogue query uses SQL this endpoint does not \
3117 implement: {}", why)));
3118 }
3119 return Ok(None);
3120 }
3121 };
3122
3123 if !sql_engine_owns(sql) {
3129 return Ok(None);
3130 }
3131
3132 let resolve = |name: &str| -> anyhow::Result<Option<Box<dyn crate::sqlselect::Relation>>> {
3133 let cname = catalog_name(name);
3134 if let Some(rows) = crate::pgcatalog::rows(&cname, db) {
3135 return Ok(Some(crate::sqlselect::from_vec(rows)));
3139 }
3140 match db {
3149 Some(db) => match crate::nql::query(db, &format!("FROM {}", cname)) {
3150 Ok((rows, _)) => Ok(Some(crate::sqlselect::from_vec(rows))),
3151 Err(_) => Ok(None),
3152 },
3153 None => Ok(None),
3154 }
3155 };
3156
3157 let (cols, rows, plan) = crate::sqlselect::execute_explain(
3158 &sel,
3159 &resolve,
3160 crate::sqljoin::JoinExec::Auto,
3161 )
3162 .map_err(|e| err_msg("42601", &e.to_string()))?;
3163
3164 Ok(Some((
3165 Executed {
3166 rows,
3167 project: cols
3171 .iter()
3172 .map(|c| Col::renamed(&c.key, &c.name))
3173 .collect(),
3174 has_rows: true,
3175 tag: "SELECT".into(),
3176 tag_counts_rows: true,
3177 },
3178 plan,
3179 )))
3180}
3181
3182fn strip_explain(sql: &str) -> Option<&str> {
3189 let t = sql.trim().trim_end_matches(';').trim();
3190 let mut rest = t.strip_prefix("EXPLAIN").or_else(|| t.strip_prefix("explain"))?;
3191 if !rest.starts_with(char::is_whitespace) {
3193 return None;
3194 }
3195 rest = rest.trim_start();
3196 loop {
3197 let low = rest.to_lowercase();
3198 if let Some(r) = low.strip_prefix("analyze").or_else(|| low.strip_prefix("analyse")) {
3199 if r.starts_with(char::is_whitespace) || r.is_empty() {
3200 rest = rest[rest.len() - r.len()..].trim_start();
3201 continue;
3202 }
3203 }
3204 if let Some(r) = low.strip_prefix("verbose") {
3205 if r.starts_with(char::is_whitespace) || r.is_empty() {
3206 rest = rest[rest.len() - r.len()..].trim_start();
3207 continue;
3208 }
3209 }
3210 break;
3211 }
3212 Some(rest)
3213}
3214
3215fn plan_result(lines: Vec<String>) -> Executed {
3218 Executed {
3219 rows: lines
3220 .into_iter()
3221 .map(|l| serde_json::json!({ "QUERY PLAN": l }))
3222 .collect(),
3223 project: vec![Col::same("QUERY PLAN")],
3224 has_rows: true,
3225 tag: "EXPLAIN".into(),
3226 tag_counts_rows: false,
3227 }
3228}
3229
3230fn mentions_catalog(sql: &str) -> bool {
3237 let low = sql.to_lowercase();
3238 low.contains("pg_catalog.")
3239 || low.contains("information_schema.")
3240 || low.contains("from pg_")
3241 || low.contains("join pg_")
3242}
3243
3244fn catalog_target(nql: &str) -> Option<String> {
3249 let coll = crate::nql::parse(nql).ok()?.coll;
3250 if crate::pgcatalog::is_catalog(&coll) {
3251 Some(coll)
3252 } else {
3253 None
3254 }
3255}
3256
3257fn wants_returning(sql: &str) -> bool {
3261 find_kw(&sql.to_uppercase(), "RETURNING").is_some()
3262}
3263
3264fn next_row_id() -> String {
3266 use std::sync::atomic::{AtomicU64, Ordering};
3267 static N: AtomicU64 = AtomicU64::new(0);
3268 let n = N.fetch_add(1, Ordering::Relaxed);
3269 let ts = std::time::SystemTime::now()
3270 .duration_since(std::time::UNIX_EPOCH)
3271 .map(|d| d.as_micros())
3272 .unwrap_or(0);
3273 format!("r{}{}", ts, n)
3274}
3275
3276pub struct Executed {
3284 pub rows: Vec<Value>,
3286 pub project: Vec<Col>,
3288 pub has_rows: bool,
3292 pub tag: String,
3295 pub tag_counts_rows: bool,
3297}
3298
3299impl Executed {
3300 fn nothing(tag: &str) -> Self {
3301 Executed { rows: vec![], project: vec![], has_rows: false, tag: tag.to_string(), tag_counts_rows: false }
3302 }
3303 fn tag_for(&self, sent: usize) -> String {
3305 if self.tag_counts_rows { format!("{} {}", self.tag, sent) } else { self.tag.clone() }
3306 }
3307}
3308
3309fn execute_stmt(
3315 stmt_sql: &str,
3316 db_name: &str,
3317 db: Option<&Arc<Db>>,
3318 read_only: bool,
3319) -> Result<Executed, Vec<u8>> {
3320 if let Some(inner) = strip_explain(stmt_sql) {
3329 if let Some((_, plan)) = try_catalog_select(inner, db)? {
3330 return Ok(plan_result(plan.render()));
3331 }
3332 let mut lines = vec![];
3333 match translate(inner) {
3334 Ok(_) => {
3335 lines.push(
3336 "NQL path — this statement is translated to NQL and \
3337 executed by the storage engine, not by the SQL evaluator."
3338 .to_string(),
3339 );
3340 lines.push(
3341 "No plan is reported, because the SQL evaluator is not \
3342 what runs it. Reporting one would describe a pipeline \
3343 that never executed."
3344 .to_string(),
3345 );
3346 lines.push(
3347 "The SQL evaluator (joins, CASE, scalar functions, a \
3348 hash-join planner) currently serves catalogue queries."
3349 .to_string(),
3350 );
3351 }
3352 Err(why) => lines.push(format!("cannot be executed: {why}")),
3353 }
3354 return Ok(plan_result(lines));
3355 }
3356
3357 if let Some((done, _plan)) = try_catalog_select(stmt_sql, db)? {
3358 return Ok(done);
3359 }
3360
3361 let stmt = translate(stmt_sql).map_err(|why| err_msg("0A000", &why))?;
3362
3363 macro_rules! need_db {
3366 () => {
3367 match db {
3368 Some(db) => db,
3369 None => return Err(no_db(db_name)),
3370 }
3371 };
3372 }
3373 macro_rules! need_write {
3374 () => {
3375 if read_only {
3376 return Err(err_msg("25006", READ_ONLY_MSG));
3377 }
3378 };
3379 }
3380
3381 match stmt {
3382 Stmt::Ok(tag) => Ok(Executed::nothing(if tag.is_empty() { "SELECT 0" } else { tag })),
3383
3384 Stmt::Canned { cols, row } => {
3385 let mut obj = serde_json::Map::new();
3388 for (c, v) in cols.iter().zip(row.iter()) {
3389 obj.insert(c.clone(), Value::String(v.clone()));
3390 }
3391 Ok(Executed {
3392 rows: vec![Value::Object(obj)],
3393 project: cols.iter().map(|c| Col::same(c)).collect(),
3394 has_rows: true,
3395 tag: "SELECT".into(),
3396 tag_counts_rows: true,
3397 })
3398 }
3399
3400 Stmt::Query { nql, project } => {
3401 if let Some(coll) = catalog_target(&nql) {
3411 let rows = crate::pgcatalog::rows(&coll, db)
3412 .expect("catalog_target only returns names pgcatalog serves");
3413 let rows = crate::nql::query_rows(rows, &nql)
3414 .map_err(|e| err_msg("42601", &e.to_string()))?;
3415 return Ok(Executed {
3416 rows, project, has_rows: true,
3417 tag: "SELECT".into(), tag_counts_rows: true,
3418 });
3419 }
3420 let db = need_db!();
3421 let (rows, _) = crate::nql::query(db, &nql).map_err(|e| {
3422 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3423 })?;
3424 Ok(Executed { rows, project, has_rows: true, tag: "SELECT".into(), tag_counts_rows: true })
3425 }
3426
3427 Stmt::Insert { coll, rows, returning } => {
3428 let db = need_db!();
3429 need_write!();
3430 let mut written: Vec<Value> = vec![];
3431 for (i, r) in rows.iter().enumerate() {
3432 let id = match &r.id {
3436 Some(id) => id.clone(),
3437 None => format!("{}-{}", next_row_id(), i),
3438 };
3439 let node = db
3440 .put(&coll, &id, Value::Object(r.doc.clone()),
3441 r.caused_by.clone(), r.valid_from.clone(), r.valid_to.clone())
3442 .map_err(|e| err_msg("XX000", &format!("INSERT failed: {}", e)))?;
3443 written.push(crate::nql::node_to_json(&node));
3444 }
3445 let n = written.len();
3446 let has_rows = wants_returning(stmt_sql);
3447 Ok(Executed {
3448 rows: if has_rows { written } else { vec![] },
3449 project: returning,
3450 has_rows,
3451 tag: format!("INSERT 0 {}", n),
3453 tag_counts_rows: false,
3454 })
3455 }
3456
3457 Stmt::Update { coll, set, nql, returning } => {
3458 let db = need_db!();
3459 need_write!();
3460 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
3463 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3464 })?;
3465 let mut written: Vec<Value> = vec![];
3466 for row in &matched {
3467 let id = match row.get("_id").and_then(|v| v.as_str()) {
3468 Some(id) => id.to_string(),
3469 None => continue,
3470 };
3471 let mut doc = match db.get(&coll, &id) {
3475 Some(n) => match n.data {
3476 Value::Object(m) => m,
3477 _ => serde_json::Map::new(),
3478 },
3479 None => continue,
3480 };
3481 for (k, v) in &set {
3482 doc.insert(k.clone(), v.clone());
3483 }
3484 let node = db
3487 .put(&coll, &id, Value::Object(doc), vec![], None, None)
3488 .map_err(|e| err_msg("XX000", &format!("UPDATE failed: {}", e)))?;
3489 written.push(crate::nql::node_to_json(&node));
3490 }
3491 let n = written.len();
3492 let has_rows = wants_returning(stmt_sql);
3493 Ok(Executed {
3494 rows: if has_rows { written } else { vec![] },
3495 project: returning,
3496 has_rows,
3497 tag: format!("UPDATE {}", n),
3498 tag_counts_rows: false,
3499 })
3500 }
3501
3502 Stmt::Delete { coll, nql, returning } => {
3503 let db = need_db!();
3504 need_write!();
3505 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
3506 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3507 })?;
3508 let returned = matched.clone();
3511 let mut n = 0usize;
3512 for row in &matched {
3513 if let Some(id) = row.get("_id").and_then(|v| v.as_str()) {
3514 match db.delete(&coll, id) {
3515 Ok(true) => n += 1,
3516 Ok(false) => {}
3517 Err(e) => return Err(err_msg("XX000", &format!("DELETE failed: {}", e))),
3518 }
3519 }
3520 }
3521 let has_rows = wants_returning(stmt_sql);
3522 Ok(Executed {
3523 rows: if has_rows { returned } else { vec![] },
3524 project: returning,
3525 has_rows,
3526 tag: format!("DELETE {}", n),
3527 tag_counts_rows: false,
3528 })
3529 }
3530 }
3531}
3532
3533fn run_simple_query(sql: &str, db_name: &str, db: Option<&Arc<Db>>, read_only: bool) -> Vec<u8> {
3535 let mut out = vec![];
3536 let statements = split_statements(sql);
3537 if statements.is_empty() {
3538 return Out::msg(b'I').finish();
3540 }
3541 for stmt_sql in statements {
3542 match execute_stmt(&stmt_sql, db_name, db, read_only) {
3543 Err(encoded) => {
3545 out.extend_from_slice(&encoded);
3546 return out;
3547 }
3548 Ok(ex) => {
3549 if ex.has_rows {
3550 out.extend_from_slice(&encode_rows(&ex.rows, &ex.project));
3551 }
3552 out.extend_from_slice(&command_complete(&ex.tag_for(ex.rows.len())));
3553 }
3554 }
3555 }
3556 out
3557}
3558
3559fn split_statements(sql: &str) -> Vec<String> {
3561 let mut out = vec![];
3562 let mut cur = String::new();
3563 let mut in_s = false;
3564 for c in sql.chars() {
3565 match c {
3566 '\'' => { in_s = !in_s; cur.push(c); }
3567 ';' if !in_s => {
3568 if !cur.trim().is_empty() { out.push(cur.clone()); }
3569 cur.clear();
3570 }
3571 _ => cur.push(c),
3572 }
3573 }
3574 if !cur.trim().is_empty() {
3575 out.push(cur);
3576 }
3577 out
3578}
3579
3580pub async fn run(host: &str, port: u16, resolver: Arc<dyn DbResolver>) -> anyhow::Result<()> {
3582 let read_only = std::env::var("NEDBD_PG_READ_ONLY")
3586 .map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
3587 .unwrap_or(false);
3588 let listener = TcpListener::bind((host, port)).await?;
3589 println!(" pgwire postgres endpoint on {}:{} — psql / DBeaver / psycopg ({})",
3590 host, port,
3591 if read_only { "SELECT only — read-only mode" } else { "SELECT + INSERT/UPDATE/DELETE" });
3592 loop {
3593 let (sock, _peer) = match listener.accept().await {
3594 Ok(v) => v,
3595 Err(e) => {
3596 eprintln!(" [pgwire] accept failed: {}", e);
3597 continue;
3598 }
3599 };
3600 let r = Arc::clone(&resolver);
3601 tokio::spawn(async move {
3602 let _ = sock.set_nodelay(true);
3603 if let Err(e) = handle(sock, r, read_only).await {
3604 if e.kind() != std::io::ErrorKind::UnexpectedEof
3606 && e.kind() != std::io::ErrorKind::ConnectionReset
3607 {
3608 eprintln!(" [pgwire] connection error: {}", e);
3609 }
3610 }
3611 });
3612 }
3613}
3614
3615#[cfg(test)]
3618mod explain_tests {
3619 use super::*;
3620
3621 #[test]
3622 fn a_bare_explain_is_stripped() {
3623 assert_eq!(strip_explain("EXPLAIN SELECT 1"), Some("SELECT 1"));
3624 assert_eq!(strip_explain("explain select 1"), Some("select 1"));
3625 assert_eq!(strip_explain(" EXPLAIN SELECT 1 ; "), Some("SELECT 1"));
3626 }
3627
3628 #[test]
3629 fn analyze_and_verbose_are_accepted_and_ignored() {
3630 assert_eq!(strip_explain("EXPLAIN ANALYZE SELECT 1"), Some("SELECT 1"));
3634 assert_eq!(strip_explain("EXPLAIN ANALYSE SELECT 1"), Some("SELECT 1"));
3635 assert_eq!(strip_explain("EXPLAIN VERBOSE SELECT 1"), Some("SELECT 1"));
3636 assert_eq!(strip_explain("EXPLAIN ANALYZE VERBOSE SELECT 1"), Some("SELECT 1"));
3637 assert_eq!(strip_explain("explain analyze verbose select 1"), Some("select 1"));
3638 }
3639
3640 #[test]
3641 fn a_word_merely_starting_with_explain_is_not_a_keyword() {
3642 assert_eq!(strip_explain("EXPLAINED SELECT 1"), None);
3643 assert_eq!(strip_explain("SELECT 1"), None);
3644 assert_eq!(strip_explain("SELECT explain FROM t"), None);
3645 }
3646
3647 #[test]
3648 fn a_column_named_analyze_is_not_eaten() {
3649 assert_eq!(strip_explain("EXPLAIN analyzed_view"), Some("analyzed_view"));
3652 }
3653
3654 #[test]
3655 fn the_plan_result_has_postgres_shape() {
3656 let e = plan_result(vec!["Seq Scan on t".into(), "note".into()]);
3657 assert_eq!(e.project.len(), 1);
3658 assert_eq!(e.project[0].out, "QUERY PLAN");
3659 assert_eq!(e.rows.len(), 2);
3660 assert_eq!(e.rows[0]["QUERY PLAN"], "Seq Scan on t");
3661 assert_eq!(e.tag, "EXPLAIN");
3662 assert!(!e.tag_counts_rows);
3664 }
3665}
3666
3667#[cfg(test)]
3668mod tests {
3669 use super::*;
3670 use serde_json::json;
3671
3672 fn q(sql: &str) -> String {
3673 match translate(sql) {
3674 Ok(Stmt::Query { nql, .. }) => nql,
3675 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3676 }
3677 }
3678 fn proj(sql: &str) -> Vec<String> {
3680 match translate(sql) {
3681 Ok(Stmt::Query { project, .. }) => project.iter().map(|c| c.out.clone()).collect(),
3682 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3683 }
3684 }
3685 fn proj_pairs(sql: &str) -> Vec<(String, String)> {
3687 match translate(sql) {
3688 Ok(Stmt::Query { project, .. }) =>
3689 project.iter().map(|c| (c.src.clone(), c.out.clone())).collect(),
3690 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3691 }
3692 }
3693 fn names(cols: &[Col]) -> Vec<String> { cols.iter().map(|c| c.out.clone()).collect() }
3694
3695 fn cols_of(sql: &str) -> Vec<Col> {
3699 match translate(sql).unwrap() {
3700 Stmt::Query { project, .. } => project,
3701 other => panic!("{:?}", other),
3702 }
3703 }
3704
3705 #[test]
3706 fn select_star_becomes_bare_from() {
3707 assert_eq!(q("SELECT * FROM orders"), "FROM orders");
3708 assert_eq!(q("select * from orders;"), "FROM orders");
3709 assert_eq!(proj("SELECT * FROM orders"), Vec::<String>::new());
3710 }
3711
3712 #[test]
3713 fn a_column_list_becomes_a_projection_not_a_clause() {
3714 assert_eq!(q("SELECT status, total FROM orders"), "FROM orders");
3717 assert_eq!(proj("SELECT status, total FROM orders"), vec!["status", "total"]);
3718 }
3719
3720 #[test]
3721 fn a_qualifier_reduces_to_the_field_while_an_ALIAS_is_the_name_the_client_sees() {
3722 let cols = cols_of("SELECT o.status AS s, o.total total, o.region FROM orders o");
3730 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
3731 vec!["status", "total", "region"]);
3732 assert_eq!(cols.iter().map(|c| c.out.clone()).collect::<Vec<_>>(),
3733 vec!["s", "total", "region"]);
3734 assert_eq!(q("SELECT * FROM public.orders"), "FROM orders");
3735 assert_eq!(q("SELECT * FROM \"orders\""), "FROM orders");
3736 }
3737
3738 #[test]
3739 fn a_select_list_may_MIX_columns_with_an_aggregate() {
3740 assert_eq!(q("SELECT status, count(*) AS count_1 FROM orders GROUP BY status"),
3748 "FROM orders GROUP BY status COUNT");
3749 assert_eq!(q("SELECT status, count(*) FROM orders WHERE total > 1 GROUP BY status ORDER BY status LIMIT 5"),
3752 "FROM orders WHERE total > 1 GROUP BY status COUNT ORDER BY status LIMIT 5");
3753 assert_eq!(q("SELECT count(*) FROM orders"), "FROM orders COUNT");
3755 assert_eq!(q("SELECT sum(total) FROM orders"), "FROM orders SUM total");
3756 let e = translate("SELECT status, count(*) FROM orders GROUP BY status, region").unwrap_err();
3760 assert!(e.contains("GROUP BY takes one key"), "{}", e);
3761 let cols = cols_of("SELECT status, count(*) AS count_1 FROM orders GROUP BY status");
3762 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
3763 vec!["status", "count"]);
3764 assert_eq!(cols.iter().map(|c| c.out.clone()).collect::<Vec<_>>(),
3765 vec!["status", "count_1"]);
3766
3767 let cols = cols_of("SELECT status, count(*), sum(total) FROM orders GROUP BY status");
3770 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
3771 vec!["status", "count", "sum_total"]);
3772 assert_eq!(q("SELECT status, count(*), sum(total) FROM orders GROUP BY status"),
3773 "FROM orders GROUP BY status SUM total");
3774
3775 assert_eq!(q("SELECT o.status, sum(o.total) FROM orders o GROUP BY o.status"),
3777 "FROM orders GROUP BY status SUM total");
3778
3779 let e = translate("SELECT status, sum(total), avg(total) FROM orders GROUP BY status")
3782 .unwrap_err();
3783 assert!(e.contains("only one of SUM/AVG/MIN/MAX"), "{}", e);
3784
3785 let e = translate("SELECT status, total, count(*) FROM orders GROUP BY status")
3787 .unwrap_err();
3788 assert!(e.contains("must appear in the GROUP BY clause"), "{}", e);
3789 }
3790
3791 #[test]
3792 fn ORDER_BY_an_ordinal_resolves_to_that_select_list_column() {
3793 assert_eq!(q("SELECT status, total FROM orders ORDER BY 1"),
3798 "FROM orders ORDER BY status");
3799 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 DESC"),
3800 "FROM orders ORDER BY total DESC");
3801 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 DESC, 1"),
3803 "FROM orders ORDER BY total DESC, status");
3804 assert_eq!(q("SELECT status, total FROM orders ORDER BY 1, total DESC"),
3805 "FROM orders ORDER BY status, total DESC");
3806 assert_eq!(q("SELECT status, count(*) AS n FROM orders GROUP BY status ORDER BY 1"),
3810 "FROM orders GROUP BY status COUNT ORDER BY status");
3811 assert_eq!(q("SELECT status, count(*) AS n FROM orders GROUP BY status ORDER BY 2 DESC"),
3813 "FROM orders GROUP BY status COUNT ORDER BY count DESC");
3814 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 LIMIT 1"),
3817 "FROM orders ORDER BY total LIMIT 1");
3818 assert_eq!(q("SELECT status FROM orders WHERE total > 1 ORDER BY 1"),
3820 "FROM orders WHERE total > 1 ORDER BY status");
3821
3822 let e = translate("SELECT status FROM orders ORDER BY 4").unwrap_err();
3826 assert!(e.contains("out of range") && e.contains("1 column"), "{}", e);
3827 let e = translate("SELECT * FROM orders ORDER BY 1").unwrap_err();
3828 assert!(e.contains("no list to index"), "{}", e);
3829 }
3830
3831 #[test]
3832 fn count_of_a_subquery_flattens_only_when_the_two_counts_MUST_agree() {
3833 assert_eq!(
3837 q("SELECT count(*) AS count_1 FROM (SELECT orders._id AS a, orders.status AS b \
3838 FROM orders WHERE orders.status = 'paid') AS anon_1"),
3839 r#"FROM orders COUNT WHERE status = "paid""#);
3842 assert_eq!(q("SELECT count(*) FROM (SELECT orders._id FROM orders) AS anon_1"),
3844 "FROM orders COUNT");
3845 assert_eq!(q("SELECT count(*) FROM (SELECT _id FROM orders ORDER BY total DESC) AS a"),
3847 "FROM orders COUNT");
3848 let cols = cols_of("SELECT count(*) AS count_1 FROM (SELECT _id FROM orders) AS a");
3850 assert_eq!(cols[0].src, "count");
3851 assert_eq!(cols[0].out, "count_1");
3852
3853 for sql in [
3856 "SELECT count(*) FROM (SELECT _id FROM orders LIMIT 1) AS a",
3858 "SELECT count(*) FROM (SELECT _id FROM orders OFFSET 1) AS a",
3859 "SELECT count(*) FROM (SELECT status FROM orders GROUP BY status) AS a",
3861 "SELECT count(*) FROM (SELECT count(*) FROM orders) AS a",
3863 "SELECT count(*) FROM (SELECT sum(total) FROM orders) AS a",
3864 "SELECT count(*), status FROM (SELECT status FROM orders) AS a",
3866 "SELECT status FROM (SELECT status FROM orders) AS a",
3867 "SELECT count(*) FROM (SELECT x FROM (SELECT _id AS x FROM orders) AS b) AS a",
3869 ] {
3870 let e = translate(sql).unwrap_err();
3871 assert!(e.contains("subqueries in FROM"), "{} -> {}", sql, e);
3872 }
3873
3874 for (sql, needle) in [
3879 ("SELECT count(*) FROM (SELECT DISTINCT status FROM orders) AS a", "DISTINCT"),
3880 ("SELECT count(*) FROM (SELECT a FROM t UNION SELECT b FROM u) AS x", "UNION"),
3881 ] {
3882 let e = translate(sql).unwrap_err();
3883 assert!(e.contains(needle), "{} -> {}", sql, e);
3884 }
3885 }
3886
3887 #[test]
3888 fn a_QUALIFIED_column_in_WHERE_finds_its_field_instead_of_ZERO_ROWS() {
3889 assert_eq!(q("SELECT _id FROM orders WHERE orders.status = 'paid'"),
3897 r#"FROM orders WHERE status = "paid""#);
3898 assert_eq!(q("SELECT _id FROM orders WHERE orders.total > 50"),
3899 "FROM orders WHERE total > 50");
3900 assert_eq!(q("SELECT _id FROM orders ORDER BY orders.total DESC LIMIT 2"),
3902 "FROM orders ORDER BY total DESC LIMIT 2");
3903 assert_eq!(q("SELECT status, count(*) FROM orders GROUP BY orders.status"),
3904 "FROM orders GROUP BY status COUNT");
3905
3906 assert_eq!(q("SELECT o.status FROM orders o WHERE o.status = 'paid'"),
3910 r#"FROM orders WHERE status = "paid""#);
3911 assert_eq!(q("SELECT o.status FROM orders AS o WHERE o.total > 1"),
3912 "FROM orders WHERE total > 1");
3913
3914 let e = translate("SELECT _id FROM orders WHERE nosuch.status = 'paid'").unwrap_err();
3919 assert!(e.contains("no table or alias named \"nosuch\""), "{}", e);
3920 let e = translate("SELECT _id FROM orders o WHERE p.status = 'paid'").unwrap_err();
3921 assert!(e.contains("aliased \"o\""), "the message names the alias in scope: {}", e);
3922
3923 assert_eq!(q("SELECT _id FROM orders WHERE status = 'pa.id'"),
3925 r#"FROM orders WHERE status = "pa.id""#);
3926 assert_eq!(q("SELECT _id FROM orders WHERE total > 1.5"),
3928 "FROM orders WHERE total > 1.5");
3929
3930 match translate("UPDATE orders o SET status = 'x' WHERE o.total > 5").unwrap() {
3932 Stmt::Update { coll, nql, .. } => {
3933 assert_eq!(coll, "orders", "the alias is not part of the collection name");
3934 assert_eq!(nql, "FROM orders WHERE total > 5");
3935 }
3936 other => panic!("{:?}", other),
3937 }
3938 match translate("DELETE FROM orders o WHERE o.status = 'paid'").unwrap() {
3939 Stmt::Delete { coll, nql, .. } => {
3940 assert_eq!(coll, "orders");
3941 assert_eq!(nql, r#"FROM orders WHERE status = "paid""#);
3942 }
3943 other => panic!("{:?}", other),
3944 }
3945
3946 assert_eq!(q("SELECT _id FROM orders AS OF SYSTEM TIME 3 WHERE orders.total > 1"),
3948 "FROM orders AS OF 3 WHERE total > 1");
3949 }
3950
3951 #[test]
3952 fn where_clauses_pass_through_with_sql_literals_rewritten() {
3953 assert_eq!(q("SELECT * FROM orders WHERE status = 'paid'"),
3954 r#"FROM orders WHERE status = "paid""#);
3955 assert_eq!(q("SELECT * FROM orders WHERE status <> 'paid'"),
3956 r#"FROM orders WHERE status != "paid""#);
3957 assert_eq!(q("SELECT * FROM orders WHERE status IN ('paid','open')"),
3958 r#"FROM orders WHERE status IN ("paid","open")"#);
3959 }
3960
3961 #[test]
3964 fn a_doubled_sql_quote_is_one_literal_character() {
3965 assert_eq!(q("SELECT * FROM t WHERE name = 'it''s'"),
3966 r#"FROM t WHERE name = "it's""#);
3967 }
3968
3969 #[test]
3972 fn a_double_quote_inside_a_sql_literal_is_escaped_for_nql() {
3973 assert_eq!(q(r#"SELECT * FROM t WHERE name = 'say "hi"'"#),
3974 r#"FROM t WHERE name = "say \"hi\"""#);
3975 }
3976
3977 #[test]
3978 fn the_shared_clauses_are_handed_to_nql_unchanged() {
3979 assert_eq!(q("SELECT * FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5"),
3980 "FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5");
3981 assert_eq!(q("SELECT * FROM orders GROUP BY region"), "FROM orders GROUP BY region");
3982 assert_eq!(q("SELECT * FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC"),
3983 "FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC");
3984 }
3985
3986 #[test]
3993 fn an_aggregate_is_one_column_named_as_sql_names_it() {
3994 assert_eq!(proj_pairs("SELECT COUNT(*) FROM orders"),
3995 vec![("count".to_string(), "count".to_string())]);
3996 assert_eq!(proj_pairs("SELECT SUM(total) FROM orders"),
3997 vec![("sum_total".to_string(), "sum".to_string())]);
3998 assert_eq!(proj_pairs("SELECT avg(total) FROM orders"),
3999 vec![("avg_total".to_string(), "avg".to_string())]);
4000 assert_eq!(proj_pairs("SELECT MIN(total) FROM orders"),
4001 vec![("min_total".to_string(), "min".to_string())]);
4002 let rows = vec![json!({"count": 4, "sum_total": 420, "value": 420})];
4004 let p = vec![Col::renamed("sum_total", "sum")];
4005 let cols = columns_for(&rows, &p);
4006 assert_eq!(names(&cols), vec!["sum"], "one column, SQL's name");
4007 assert_eq!(cell(rows[0].get(&cols[0].src)), Some("420".to_string()));
4008 }
4009
4010 #[test]
4015 fn a_bare_column_with_group_by_is_refused_not_nulled() {
4016 let e = translate("SELECT region, total FROM orders GROUP BY region").unwrap_err();
4017 assert!(e.contains("must appear in the GROUP BY clause"), "{}", e);
4018 assert!(e.contains("total"), "the message names the offending column: {}", e);
4019
4020 assert!(translate("SELECT region FROM orders GROUP BY region").is_ok());
4022 assert!(translate("SELECT region, count FROM orders GROUP BY region").is_ok());
4023 assert!(translate("SELECT SUM(total) FROM orders GROUP BY region").is_ok());
4025 assert!(translate("SELECT * FROM orders GROUP BY region").is_ok());
4027 }
4028
4029 #[test]
4030 fn count_star_becomes_nql_count() {
4031 assert_eq!(q("SELECT COUNT(*) FROM orders"), "FROM orders COUNT");
4032 assert_eq!(q("SELECT count(*) FROM orders WHERE total > 5"),
4033 "FROM orders COUNT WHERE total > 5");
4034 }
4035
4036 #[test]
4037 fn aggregates_carry_their_target_column() {
4038 assert_eq!(q("SELECT SUM(total) FROM orders"), "FROM orders SUM total");
4039 assert_eq!(q("SELECT avg(total) FROM orders WHERE region = 'eu'"),
4040 r#"FROM orders AVG total WHERE region = "eu""#);
4041 assert!(translate("SELECT SUM(*) FROM orders").is_err());
4042 }
4043
4044 #[test]
4047 fn as_of_system_time_bridges_to_nql_as_of() {
4048 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42"),
4049 "FROM orders AS OF 42");
4050 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42 WHERE total > 1"),
4051 "FROM orders AS OF 42 WHERE total > 1");
4052 let e = translate("SELECT * FROM orders AS OF SYSTEM TIME '2026-01-01'").unwrap_err();
4054 assert!(e.contains("sequence number"), "{}", e);
4055 }
4056
4057 #[test]
4066 fn a_select_list_expression_is_refused_rather_than_answered_blank() {
4067 for sql in [
4068 "SELECT total * 2 FROM orders",
4069 "SELECT total, total*2 AS doubled FROM orders",
4070 "SELECT total + 1 FROM orders",
4071 "SELECT status || 'x' FROM orders",
4072 "SELECT -total FROM orders",
4073 "SELECT lower(status) FROM orders",
4074 ] {
4075 let e = translate(sql).unwrap_err();
4076 assert!(e.contains("expressions in the select list"), "{} -> {}", sql, e);
4077 }
4078 assert_eq!(q("SELECT _id, status FROM orders"), "FROM orders");
4081 assert_eq!(q("SELECT \"status\" FROM orders"), "FROM orders");
4082 assert_eq!(q("SELECT orders.status FROM orders"), "FROM orders");
4083 assert_eq!(q("SELECT o.status FROM orders o"), "FROM orders");
4084 assert_eq!(q("SELECT total AS t FROM orders"), "FROM orders");
4085 assert!(translate("SELECT count(*) FROM orders").is_ok());
4086 assert!(translate("SELECT sum(total) FROM orders").is_ok());
4087 }
4088
4089 #[test]
4097 fn having_is_translated_to_nqls_spelling_and_refuses_an_unknown_key() {
4098 for sql in [
4100 "SELECT status, count(*) AS n FROM orders GROUP BY status HAVING count(*) > 1",
4101 "SELECT status, count(*) AS n FROM orders GROUP BY status HAVING n > 1",
4102 "SELECT status, count(*) FROM orders GROUP BY status HAVING COUNT > 1",
4103 "SELECT status, count(*) FROM orders GROUP BY status HAVING count > 1",
4104 ] {
4105 let got = q(sql);
4106 assert_eq!(got, "FROM orders GROUP BY status COUNT HAVING count > 1",
4107 "{} -> {}", sql, got);
4108 }
4109 assert_eq!(q("SELECT status, sum(total) AS s FROM orders GROUP BY status HAVING s > 100"),
4111 "FROM orders GROUP BY status SUM total HAVING sum_total > 100");
4112 assert_eq!(q("SELECT status, sum(total) FROM orders GROUP BY status HAVING sum_total > 100"),
4114 "FROM orders GROUP BY status SUM total HAVING sum_total > 100");
4115 assert_eq!(q("SELECT status, count(*) FROM orders GROUP BY status HAVING status > 'a'"),
4118 "FROM orders GROUP BY status COUNT HAVING status > \"a\"");
4119 let e = translate(
4121 "SELECT status, count(*) FROM orders GROUP BY status HAVING nosuch > 1").unwrap_err();
4122 assert!(e.contains("HAVING names") && e.contains("nosuch"), "{}", e);
4123 assert!(e.contains("zero rows"), "the message must say what it prevented: {}", e);
4124 }
4125
4126 #[test]
4127 fn handshake_queries_are_answered_so_clients_can_connect() {
4128 assert!(matches!(translate("SELECT version()"), Ok(Stmt::Canned { .. })));
4129 assert!(matches!(translate("SHOW transaction_isolation"), Ok(Stmt::Canned { .. })));
4130 assert!(matches!(translate("SELECT current_schema()"), Ok(Stmt::Canned { .. })));
4131 assert!(matches!(translate("SET extra_float_digits = 3"), Ok(Stmt::Ok(_))));
4132 assert!(matches!(translate("BEGIN"), Ok(Stmt::Ok(_))));
4133 assert!(matches!(translate(""), Ok(Stmt::Ok(_))));
4134 }
4135
4136 #[test]
4139 fn unsupported_sql_is_refused_with_a_reason() {
4140 for (sql, expect) in [
4141 ("INSERT INTO t VALUES (1)", "explicit column list"),
4142 ("CREATE TABLE t (a int)", "DDL"),
4143 ("TRUNCATE t", "append-only"),
4144 ("GRANT ALL ON t TO x", "privilege system"),
4145 ("SELECT * FROM a JOIN b ON a.x = b.x", "JOIN is not supported"),
4146 ("SELECT * FROM a UNION SELECT * FROM b", "UNION"),
4147 ("SELECT DISTINCT region FROM orders", "GROUP BY"),
4148 ("SELECT * FROM (SELECT 1) x", "subqueries in FROM"),
4149 ("SELECT * FROM a, b", "more than one collection"),
4150 ("SELECT lower(status) FROM orders", "expressions in the select list"),
4151 ("VACUUM", "only SELECT"),
4152 ] {
4153 let e = translate(sql).unwrap_err();
4154 assert!(e.contains(expect), "for {:?} expected {:?} in {:?}", sql, expect, e);
4155 }
4156 }
4157
4158 fn ins(sql: &str) -> (String, Vec<InsertRow>, Vec<Col>) {
4166 match translate(sql) {
4167 Ok(Stmt::Insert { coll, rows, returning }) => (coll, rows, returning),
4168 other => panic!("expected INSERT for {:?}, got {:?}", sql, other),
4169 }
4170 }
4171
4172 #[test]
4173 fn insert_becomes_a_put_per_row() {
4174 let (coll, rows, ret) = ins("INSERT INTO orders (_id, status, total) VALUES ('o1', 'paid', 120)");
4175 assert_eq!(coll, "orders");
4176 assert_eq!(rows.len(), 1);
4177 assert_eq!(rows[0].id.as_deref(), Some("o1"));
4178 assert_eq!(rows[0].doc.get("status"), Some(&json!("paid")));
4179 assert_eq!(rows[0].doc.get("total"), Some(&json!(120)));
4180 assert!(!rows[0].doc.contains_key("_id"));
4182 assert!(ret.is_empty());
4183 }
4184
4185 #[test]
4186 fn a_multi_row_insert_yields_one_row_each() {
4187 let (_, rows, _) = ins(
4188 "INSERT INTO t (id, n) VALUES ('a', 1), ('b', 2), ('c', 3)");
4189 assert_eq!(rows.len(), 3);
4190 assert_eq!(rows[1].id.as_deref(), Some("b"));
4191 assert_eq!(rows[2].doc.get("n"), Some(&json!(3)));
4192 }
4193
4194 #[test]
4195 fn an_insert_without_an_id_column_lets_the_server_assign_one() {
4196 let (_, rows, _) = ins("INSERT INTO t (n) VALUES (1)");
4197 assert_eq!(rows[0].id, None, "the executor mints a unique key");
4198 assert_eq!(rows[0].doc.get("n"), Some(&json!(1)));
4199 }
4200
4201 #[test]
4204 fn insert_lifts_provenance_out_of_reserved_columns() {
4205 let (_, rows, _) = ins(
4206 "INSERT INTO audit (_id, _caused_by, _valid_from, kind) \
4207 VALUES ('e1', 'abc123', '2026-01-01', 'reprice')");
4208 assert_eq!(rows[0].caused_by, vec!["abc123".to_string()]);
4209 assert_eq!(rows[0].valid_from.as_deref(), Some("2026-01-01"));
4210 assert_eq!(rows[0].doc.get("kind"), Some(&json!("reprice")));
4211 for k in ["_id", "_caused_by", "_valid_from"] {
4213 assert!(!rows[0].doc.contains_key(k), "{} leaked into the doc", k);
4214 }
4215 }
4216
4217 #[test]
4218 fn insert_values_cover_the_scalar_types() {
4219 let (_, rows, _) = ins(
4220 "INSERT INTO t (s, i, f, b, n) VALUES ('x', 42, 1.5, TRUE, NULL)");
4221 assert_eq!(rows[0].doc.get("s"), Some(&json!("x")));
4222 assert_eq!(rows[0].doc.get("i"), Some(&json!(42)));
4223 assert_eq!(rows[0].doc.get("f"), Some(&json!(1.5)));
4224 assert_eq!(rows[0].doc.get("b"), Some(&json!(true)));
4225 assert_eq!(rows[0].doc.get("n"), Some(&Value::Null));
4226 }
4227
4228 #[test]
4231 fn insert_literals_survive_quotes_and_commas() {
4232 let (_, rows, _) = ins("INSERT INTO t (a, b) VALUES ('it''s', 'x,y')");
4233 assert_eq!(rows[0].doc.get("a"), Some(&json!("it's")));
4234 assert_eq!(rows[0].doc.get("b"), Some(&json!("x,y")));
4235 }
4236
4237 #[test]
4238 fn insert_refuses_what_it_cannot_store_faithfully() {
4239 assert!(translate("INSERT INTO t (a) VALUES (1 + 1)").is_err());
4241 assert!(translate("INSERT INTO t (a) VALUES (now())").is_err());
4242 let e = translate("INSERT INTO t (a, b) VALUES (1)").unwrap_err();
4244 assert!(e.contains("values for"), "{}", e);
4245 let e2 = translate("INSERT INTO t VALUES (1)").unwrap_err();
4247 assert!(e2.contains("explicit column list"), "{}", e2);
4248 }
4249
4250 #[test]
4251 fn update_finds_rows_with_the_full_predicate_surface() {
4252 match translate("UPDATE orders SET status = 'void' WHERE total < 50 AND region IN ('eu')") {
4253 Ok(Stmt::Update { coll, set, nql, .. }) => {
4254 assert_eq!(coll, "orders");
4255 assert_eq!(set, vec![("status".to_string(), json!("void"))]);
4256 assert_eq!(nql, r#"FROM orders WHERE total < 50 AND region IN ("eu")"#);
4258 }
4259 other => panic!("expected UPDATE, got {:?}", other),
4260 }
4261 }
4262
4263 #[test]
4264 fn update_without_where_targets_the_whole_collection() {
4265 match translate("UPDATE t SET a = 1") {
4267 Ok(Stmt::Update { nql, .. }) => assert_eq!(nql, "FROM t"),
4268 other => panic!("expected UPDATE, got {:?}", other),
4269 }
4270 }
4271
4272 #[test]
4273 fn update_handles_several_assignments() {
4274 match translate("UPDATE t SET a = 1, b = 'x,y', c = NULL WHERE id = 'k'") {
4275 Ok(Stmt::Update { set, .. }) => {
4276 assert_eq!(set.len(), 3);
4277 assert_eq!(set[1], ("b".to_string(), json!("x,y")));
4278 assert_eq!(set[2], ("c".to_string(), Value::Null));
4279 }
4280 other => panic!("expected UPDATE, got {:?}", other),
4281 }
4282 assert!(translate("UPDATE t SET").is_err());
4283 assert!(translate("UPDATE t SET a").is_err());
4284 }
4285
4286 #[test]
4287 fn delete_becomes_a_predicate_over_the_collection() {
4288 match translate("DELETE FROM orders WHERE status = 'void'") {
4289 Ok(Stmt::Delete { coll, nql, .. }) => {
4290 assert_eq!(coll, "orders");
4291 assert_eq!(nql, r#"FROM orders WHERE status = "void""#);
4292 }
4293 other => panic!("expected DELETE, got {:?}", other),
4294 }
4295 match translate("DELETE FROM t") {
4296 Ok(Stmt::Delete { nql, .. }) => assert_eq!(nql, "FROM t"),
4297 other => panic!("expected DELETE, got {:?}", other),
4298 }
4299 }
4300
4301 #[test]
4302 fn returning_is_parsed_off_every_write() {
4303 let (_, _, ret) = ins("INSERT INTO t (a) VALUES (1) RETURNING a, _id");
4304 assert_eq!(ret.iter().map(|c| c.out.clone()).collect::<Vec<_>>(), vec!["a", "_id"]);
4305 let (_, _, star) = ins("INSERT INTO t (a) VALUES (1) RETURNING *");
4308 assert!(star.is_empty());
4309 assert!(wants_returning("INSERT INTO t (a) VALUES (1) RETURNING *"));
4310 assert!(!wants_returning("INSERT INTO t (a) VALUES (1)"));
4311
4312 match translate("UPDATE t SET a = 1 WHERE id = 'k' RETURNING a") {
4313 Ok(Stmt::Update { nql, returning, .. }) => {
4314 assert_eq!(returning.len(), 1);
4315 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql);
4317 }
4318 other => panic!("expected UPDATE, got {:?}", other),
4319 }
4320 match translate("DELETE FROM t WHERE id = 'k' RETURNING *") {
4321 Ok(Stmt::Delete { nql, .. }) =>
4322 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql),
4323 other => panic!("expected DELETE, got {:?}", other),
4324 }
4325 }
4326
4327 #[test]
4328 fn a_keyword_inside_a_value_is_not_a_clause() {
4329 match translate("UPDATE t SET note = 'where returning from' WHERE id = 'k'") {
4330 Ok(Stmt::Update { set, nql, .. }) => {
4331 assert_eq!(set[0].1, json!("where returning from"));
4332 assert_eq!(nql, r#"FROM t WHERE id = "k""#);
4333 }
4334 other => panic!("expected UPDATE, got {:?}", other),
4335 }
4336 }
4337
4338 #[test]
4339 fn split_top_respects_quotes_and_nesting() {
4340 assert_eq!(split_top("a, b, c", ',').len(), 3);
4341 assert_eq!(split_top("(1, 2), (3, 4)", ',').len(), 2);
4342 assert_eq!(split_top("'a,b', c", ',').len(), 2);
4343 assert_eq!(split_top("'it''s, fine', c", ',').len(), 2);
4344 }
4345
4346 #[test]
4347 fn comments_and_whitespace_do_not_confuse_the_translator() {
4348 assert_eq!(q("SELECT *\n FROM orders -- trailing note\n"), "FROM orders");
4349 assert_eq!(q("SELECT /* inline */ * FROM orders"), "FROM orders");
4350 assert_eq!(q("SELECT * FROM t WHERE note = 'from here to JOIN'"),
4352 r#"FROM t WHERE note = "from here to JOIN""#);
4353 }
4354
4355 #[test]
4356 fn find_kw_ignores_quotes_parens_and_substrings() {
4357 assert_eq!(find_kw("SELECT A FROM B", "FROM"), Some(9));
4358 assert_eq!(find_kw("SELECT 'FROM' FROM B", "FROM"), Some(14));
4359 assert_eq!(find_kw("SELECT F(x FROM y) FROM B", "FROM"), Some(19));
4360 assert_eq!(find_kw("SELECT FROMAGE", "FROM"), None);
4361 assert_eq!(find_kw("SELECT X_FROM", "FROM"), None);
4362 }
4363
4364 #[test]
4367 fn provenance_columns_sort_after_the_users_own_fields() {
4368 let rows = vec![json!({"_id":"1","_hash":"ab","status":"paid","total":9})];
4369 assert_eq!(names(&columns_for(&rows, &[])),
4370 vec!["status", "total", "_hash", "_id"]);
4371 }
4372
4373 #[test]
4374 fn an_explicit_projection_sets_the_column_order() {
4375 let rows = vec![json!({"a":1,"b":2})];
4376 let p = vec![Col::same("b"), Col::same("a")];
4377 assert_eq!(names(&columns_for(&rows, &p)), vec!["b", "a"]);
4378 }
4379
4380 #[test]
4381 fn columns_are_the_union_across_sparse_rows() {
4382 let rows = vec![json!({"a":1}), json!({"b":2})];
4384 assert_eq!(names(&columns_for(&rows, &[])), vec!["a", "b"]);
4385 }
4386
4387 #[test]
4388 fn type_oids_follow_the_first_non_null_value() {
4389 let rows = vec![json!({"i":1,"f":1.5,"b":true,"s":"x","n":null})];
4390 assert_eq!(oid_for(&rows, "i"), OID_INT8);
4391 assert_eq!(oid_for(&rows, "f"), OID_FLOAT8);
4392 assert_eq!(oid_for(&rows, "b"), OID_BOOL);
4393 assert_eq!(oid_for(&rows, "s"), OID_TEXT);
4394 assert_eq!(oid_for(&rows, "n"), OID_TEXT);
4396 assert_eq!(oid_for(&rows, "absent"), OID_TEXT);
4397 }
4398
4399 #[test]
4400 fn a_column_that_is_null_in_the_first_row_still_gets_its_type() {
4401 let rows = vec![json!({"v": null}), json!({"v": 7})];
4402 assert_eq!(oid_for(&rows, "v"), OID_INT8);
4403 }
4404
4405 #[test]
4406 fn cells_render_in_postgres_text_format() {
4407 assert_eq!(cell(Some(&json!("x"))), Some("x".to_string()));
4408 assert_eq!(cell(Some(&json!(true))), Some("t".to_string()));
4409 assert_eq!(cell(Some(&json!(false))), Some("f".to_string()));
4410 assert_eq!(cell(Some(&json!(42))), Some("42".to_string()));
4411 assert_eq!(cell(Some(&json!(null))), None);
4412 assert_eq!(cell(None), None);
4413 assert_eq!(cell(Some(&json!({"a":1}))), Some("{\"a\":1}".to_string()));
4415 }
4416
4417 #[test]
4420 fn message_framing_length_excludes_the_tag() {
4421 let mut m = Out::msg(b'Z');
4422 m.bytes(b"I");
4423 let bytes = m.finish();
4424 assert_eq!(bytes[0], b'Z');
4425 assert_eq!(i32::from_be_bytes([bytes[1], bytes[2], bytes[3], bytes[4]]), 5);
4426 assert_eq!(bytes.len(), 6);
4427 }
4428
4429 #[test]
4430 fn a_result_set_encodes_as_description_then_rows_then_complete() {
4431 let rows = vec![json!({"a": 1}), json!({"a": 2})];
4432 let out = encode_result(&rows, &[]);
4433 assert_eq!(out[0], b'T');
4434 let tags: Vec<u8> = {
4435 let mut t = vec![];
4437 let mut i = 0usize;
4438 while i < out.len() {
4439 t.push(out[i]);
4440 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
4441 i += 1 + len;
4442 }
4443 t
4444 };
4445 assert_eq!(tags, vec![b'T', b'D', b'D', b'C'],
4446 "one description, one row each, one completion");
4447 }
4448
4449 #[test]
4454 fn a_write_with_returning_emits_exactly_one_command_complete() {
4455 let rows = vec![json!({"_id": "o1", "total": 9})];
4456 let mut out = encode_rows(&rows, &[Col::same("_id")]);
4457 out.extend_from_slice(&command_complete("INSERT 0 1"));
4458 let mut tags = vec![];
4459 let mut i = 0usize;
4460 while i < out.len() {
4461 tags.push(out[i]);
4462 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
4463 i += 1 + len;
4464 }
4465 assert_eq!(tags, vec![b'T', b'D', b'C'], "one description, one row, ONE tag");
4466 assert_eq!(tags.iter().filter(|t| **t == b'C').count(), 1);
4467 assert!(!encode_rows(&rows, &[]).contains(&b'C')
4469 || encode_rows(&rows, &[]).iter().filter(|b| **b == b'C').count() > 0);
4470 let bare = encode_rows(&rows, &[Col::same("_id")]);
4471 let mut bare_tags = vec![];
4472 let mut j = 0usize;
4473 while j < bare.len() {
4474 bare_tags.push(bare[j]);
4475 let len = i32::from_be_bytes([bare[j+1], bare[j+2], bare[j+3], bare[j+4]]) as usize;
4476 j += 1 + len;
4477 }
4478 assert_eq!(bare_tags, vec![b'T', b'D'], "encode_rows never appends a tag");
4479 }
4480
4481 #[test]
4482 fn an_empty_result_still_sends_a_description() {
4483 let out = encode_result(&[], &[Col::same("a")]);
4484 assert_eq!(out[0], b'T', "clients need the shape even with no rows");
4485 }
4486
4487 #[test]
4488 fn statements_split_on_top_level_semicolons_only() {
4489 assert_eq!(split_statements("SELECT 1; SELECT 2").len(), 2);
4490 assert_eq!(split_statements("SELECT ';'").len(), 1);
4491 assert_eq!(split_statements("SELECT 1;").len(), 1);
4492 assert_eq!(split_statements(" ").len(), 0);
4493 }
4494
4495 #[test]
4496 fn an_error_names_its_sqlstate() {
4497 let e = String::from_utf8_lossy(&err_msg("0A000", "x")).to_string();
4498 assert!(e.contains("ERROR"));
4499 assert!(e.contains("0A000"));
4500 }
4501
4502 #[test]
4505 fn placeholders_are_counted_outside_string_literals() {
4506 assert_eq!(param_count("SELECT a FROM t WHERE b = $1 AND c = $2"), 2);
4507 assert_eq!(param_count("SELECT a FROM t"), 0);
4508 assert_eq!(param_count("WHERE a = $2 OR b = $2 OR c = $1"), 2);
4510 assert_eq!(param_count("SELECT a FROM t WHERE b = '$1'"), 0,
4511 "a placeholder inside a literal is data, not a parameter");
4512 assert_eq!(param_count("WHERE a = $10 AND b = $1"), 10,
4513 "two-digit indexes must not be read as $1 followed by 0");
4514 }
4515
4516 #[test]
4517 fn parameters_are_spliced_as_literals() {
4518 let out = substitute_params("WHERE a = $1 AND b = $2 AND c = $3",
4519 &[Some("'x'".into()), Some("42".into()), None]).unwrap();
4520 assert_eq!(out, "WHERE a = 'x' AND b = 42 AND c = NULL");
4521 }
4522
4523 #[test]
4524 fn substitution_leaves_string_literals_alone() {
4525 let out = substitute_params("WHERE a = '$1' AND b = $1", &[Some("9".into())]).unwrap();
4526 assert_eq!(out, "WHERE a = '$1' AND b = 9");
4527 }
4528
4529 #[test]
4530 fn too_few_parameters_is_an_error_not_a_silent_null() {
4531 let e = substitute_params("WHERE a = $2", &[Some("1".into())]).unwrap_err();
4534 assert!(e.contains("$2"), "{}", e);
4535 }
4536
4537 #[test]
4538 fn a_quote_in_a_parameter_cannot_escape_its_literal() {
4539 let lit = decode_param(Some(b"it's"), OID_TEXT, 0).unwrap().unwrap();
4540 assert_eq!(lit, "'it''s'");
4541 let out = substitute_params("WHERE a = $1", &[Some(lit)]).unwrap();
4543 assert_eq!(out, "WHERE a = 'it''s'");
4544 }
4545
4546 #[test]
4547 fn binary_parameters_decode_in_every_width_psycopg_sends() {
4548 assert_eq!(decode_param(Some(&[0x00, 0x2a]), OID_INT2, 1).unwrap().unwrap(), "42");
4551 assert_eq!(decode_param(Some(&[0, 0, 0, 7]), OID_INT4, 1).unwrap().unwrap(), "7");
4552 assert_eq!(
4553 decode_param(Some(&[0, 0, 0, 0, 0, 0, 0, 9]), OID_INT8, 1).unwrap().unwrap(), "9");
4554 assert_eq!(
4555 decode_param(Some(&0x400c_0000_0000_0000u64.to_be_bytes()), OID_FLOAT8, 1)
4556 .unwrap().unwrap(), "3.5");
4557 assert_eq!(decode_param(Some(&[1]), OID_BOOL, 1).unwrap().unwrap(), "TRUE");
4558 assert_eq!(decode_param(Some(&[0]), OID_BOOL, 1).unwrap().unwrap(), "FALSE");
4559 }
4560
4561 #[test]
4562 fn a_negative_binary_integer_keeps_its_sign() {
4563 assert_eq!(decode_param(Some(&(-5i32).to_be_bytes()), OID_INT4, 1).unwrap().unwrap(), "-5");
4564 assert_eq!(decode_param(Some(&(-5i16).to_be_bytes()), OID_INT2, 1).unwrap().unwrap(), "-5");
4565 }
4566
4567 #[test]
4568 fn a_binary_parameter_of_the_wrong_width_is_refused() {
4569 let e = decode_param(Some(&[0x2a]), OID_INT4, 1).unwrap_err();
4572 assert!(e.contains("4 bytes"), "{}", e);
4573 }
4574
4575 #[test]
4576 fn an_unspecified_text_parameter_is_treated_as_a_string() {
4577 assert_eq!(decode_param(Some(b"hello"), 0, 0).unwrap().unwrap(), "'hello'");
4580 }
4581
4582 #[test]
4583 fn a_null_parameter_decodes_to_none_in_every_format() {
4584 assert_eq!(decode_param(None, OID_TEXT, 0).unwrap(), None);
4585 assert_eq!(decode_param(None, OID_INT8, 1).unwrap(), None);
4586 }
4587
4588 #[test]
4589 fn an_unsupported_binary_type_says_so_by_name() {
4590 let e = decode_param(Some(&[0u8; 8]), 1114, 1).unwrap_err();
4591 assert!(e.contains("1114"), "{}", e);
4592 assert!(e.contains("text"), "the error should point at the way out: {}", e);
4593 }
4594
4595 #[test]
4596 fn a_text_number_that_is_not_a_number_gets_quoted() {
4597 assert_eq!(decode_param(Some(b"oops"), OID_INT8, 0).unwrap().unwrap(), "'oops'");
4600 }
4601
4602 #[test]
4603 fn a_client_declared_type_is_believed_over_inference() {
4604 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[OID_INT4, 0], None);
4607 assert_eq!(oids, vec![OID_INT4, OID_TEXT]);
4608 }
4609
4610 #[test]
4611 fn parameter_arity_is_taken_from_the_sql_when_the_client_declares_none() {
4612 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[], None);
4615 assert_eq!(oids.len(), 2);
4616 }
4617
4618 #[test]
4619 fn the_field_behind_each_placeholder_is_identified() {
4620 assert_eq!(
4621 param_fields("SELECT a FROM t WHERE qty > $1 AND status = $2", 2),
4622 vec![Some("qty".to_string()), Some("status".to_string())]);
4623 }
4624
4625 #[test]
4626 fn word_operators_do_not_hide_the_field() {
4627 assert_eq!(param_fields("SELECT a FROM t WHERE name LIKE $1", 1),
4628 vec![Some("name".to_string())]);
4629 assert_eq!(param_fields("SELECT a FROM t WHERE qty BETWEEN $1 AND $2", 2),
4630 vec![Some("qty".to_string()), Some("qty".to_string())]);
4631 assert_eq!(param_fields("SELECT a FROM t WHERE region IN ($1, $2)", 2),
4632 vec![Some("region".to_string()), Some("region".to_string())]);
4633 }
4634
4635 #[test]
4636 fn a_clause_position_types_from_the_grammar_not_from_a_column() {
4637 assert_eq!(
4641 infer_param_oids("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", &[], None),
4642 vec![OID_INT8, OID_TEXT]);
4643 assert_eq!(infer_param_oids("SELECT a FROM t AS OF $1", &[], None), vec![OID_INT8]);
4644 assert_eq!(
4647 infer_param_oids("SELECT a FROM t VALID AS OF $1", &[], None), vec![OID_TEXT]);
4648 assert_eq!(
4649 infer_param_oids("SELECT a FROM t LIMIT $1 OFFSET $2", &[], None),
4650 vec![OID_INT8, OID_INT8]);
4651 }
4652
4653 #[test]
4654 fn an_aggregate_column_types_from_what_the_aggregate_means() {
4655 assert_eq!(aggregate_oid("count", None, "t"), Some(OID_INT8));
4659 assert_eq!(aggregate_oid("avg_fee", None, "t"), Some(OID_FLOAT8),
4660 "an average is fractional even over integers");
4661 assert_eq!(aggregate_oid("max__seq", None, "t"), Some(OID_INT8));
4664 assert_eq!(aggregate_oid("total", None, "t"), None, "not an aggregate");
4665 }
4666
4667 #[test]
4668 fn the_parse_probe_uses_a_literal_that_every_clause_accepts() {
4669 let probe = probe_sql("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", 2);
4673 assert!(!probe.contains("NULL"), "{}", probe);
4674 assert!(translate(&probe).is_ok(), "the probe must parse: {}", probe);
4675 }
4676
4677 #[test]
4678 fn a_column_with_mixed_types_across_documents_is_advertised_as_text() {
4679 let rows = vec![json!({"x": 3}), json!({"x": "n/a"})];
4683 assert_eq!(oid_for(&rows, "x"), OID_TEXT);
4684 let rows = vec![json!({"x": 3}), json!({"x": 1.5})];
4686 assert_eq!(oid_for(&rows, "x"), OID_FLOAT8);
4687 let rows = vec![json!({"x": Value::Null}), json!({"x": 7})];
4689 assert_eq!(oid_for(&rows, "x"), OID_INT8);
4690 }
4691
4692 #[test]
4693 fn binary_output_encodes_each_advertised_type() {
4694 assert_eq!(cell_binary(Some(&json!(true)), OID_BOOL).unwrap().unwrap(), vec![1]);
4695 assert_eq!(cell_binary(Some(&json!(42)), OID_INT8).unwrap().unwrap(),
4696 42i64.to_be_bytes().to_vec());
4697 assert_eq!(cell_binary(Some(&json!(3.5)), OID_FLOAT8).unwrap().unwrap(),
4698 3.5f64.to_be_bytes().to_vec());
4699 assert_eq!(cell_binary(Some(&json!("hi")), OID_TEXT).unwrap().unwrap(), b"hi".to_vec());
4701 assert_eq!(cell_binary(Some(&Value::Null), OID_INT8).unwrap(), None);
4702 assert_eq!(cell(Some(&json!(true))).unwrap(), "t");
4704 }
4705
4706 #[test]
4707 fn a_value_that_does_not_fit_its_advertised_binary_type_is_refused() {
4708 let e = cell_binary(Some(&json!("nope")), OID_INT8).unwrap_err();
4713 assert!(e.contains("a string"), "{}", e);
4714 assert!(e.contains("more than one type"), "the error should explain WHY: {}", e);
4715 }
4716
4717 #[test]
4718 fn a_row_description_carries_the_requested_format_per_column() {
4719 let cols = [Col::same("a"), Col::same("b")];
4720 let m = row_description_fmt(&cols, &[OID_INT8, OID_TEXT], &[1, 0]);
4721 assert_eq!(m[0], b'T');
4722 assert_eq!(m[m.len() - 1], 0, "the last column was requested as text");
4724 }
4725
4726 #[test]
4727 fn a_qualified_column_resolves_to_its_bare_name() {
4728 assert_eq!(param_fields("SELECT a FROM t WHERE t.qty = $1", 1),
4729 vec![Some("qty".to_string())]);
4730 }
4731
4732 #[test]
4733 fn insert_placeholders_map_positionally_to_the_column_list() {
4734 assert_eq!(
4735 param_fields("INSERT INTO t (_id, qty, status) VALUES ($1, $2, $3)", 3),
4736 vec![Some("_id".to_string()), Some("qty".to_string()), Some("status".to_string())]);
4737 }
4738
4739 #[test]
4740 fn a_set_clause_placeholder_finds_its_column() {
4741 assert_eq!(param_fields("UPDATE t SET status = $1 WHERE _id = $2", 2),
4742 vec![Some("status".to_string()), Some("_id".to_string())]);
4743 }
4744
4745 #[test]
4746 fn the_target_collection_is_found_for_every_statement_kind() {
4747 assert_eq!(stmt_collection("SELECT a FROM inv WHERE b = $1"), "inv");
4748 assert_eq!(stmt_collection("UPDATE inv SET a = $1"), "inv");
4749 assert_eq!(stmt_collection("DELETE FROM inv WHERE a = $1"), "inv");
4750 assert_eq!(stmt_collection("INSERT INTO inv (a) VALUES ($1)"), "inv");
4751 assert_eq!(stmt_collection("SELECT a FROM public.inv"), "inv");
4753 assert_eq!(stmt_collection("INSERT INTO inv(a) VALUES ($1)"), "inv");
4754 }
4755
4756 #[test]
4757 fn engine_metadata_fields_type_without_touching_storage() {
4758 assert_eq!(infer_field_oid(None, "t", "_seq"), OID_INT8);
4759 assert_eq!(infer_field_oid(None, "t", "_id"), OID_TEXT);
4760 }
4761
4762 #[test]
4763 fn the_protocol_acknowledgements_are_single_empty_messages() {
4764 for (m, tag) in [
4766 (parse_complete(), b'1'), (bind_complete(), b'2'),
4767 (close_complete(), b'3'), (no_data(), b'n'), (portal_suspended(), b's'),
4768 ] {
4769 assert_eq!(m.len(), 5, "{:?}", tag as char);
4770 assert_eq!(m[0], tag);
4771 assert_eq!(i32::from_be_bytes([m[1], m[2], m[3], m[4]]), 4);
4772 }
4773 }
4774
4775 #[test]
4776 fn parameter_description_reports_its_arity_and_types() {
4777 let m = parameter_description(&[OID_TEXT, OID_INT8]);
4778 assert_eq!(m[0], b't');
4779 assert_eq!(i16::from_be_bytes([m[5], m[6]]), 2);
4780 assert_eq!(i32::from_be_bytes([m[7], m[8], m[9], m[10]]), OID_TEXT);
4781 assert_eq!(i32::from_be_bytes([m[11], m[12], m[13], m[14]]), OID_INT8);
4782 }
4783
4784 #[test]
4785 fn a_cstring_is_taken_without_its_terminator() {
4786 let body = b"one\0two\0".to_vec();
4787 let mut at = 0usize;
4788 assert_eq!(take_cstr(&body, &mut at), "one");
4789 assert_eq!(take_cstr(&body, &mut at), "two");
4790 assert_eq!(at, body.len());
4791 }
4792
4793 #[test]
4794 fn truncated_integers_are_reported_rather_than_read_past_the_end() {
4795 let body = vec![0u8, 1];
4796 let mut at = 0usize;
4797 assert!(take_i32(&body, &mut at).is_err());
4798 let mut at = 0usize;
4799 assert!(take_i16(&body, &mut at).is_ok());
4800 }
4801
4802 #[test]
4803 fn a_binary_result_format_request_is_refused_rather_than_faked() {
4804 let out = encode_rows(&[], &[Col::same("a")]);
4807 let desc_format = &out[out.len() - 2..];
4808 assert_eq!(i16::from_be_bytes([desc_format[0], desc_format[1]]), 0,
4809 "every column is advertised as text format");
4810 }
4811
4812 #[test]
4813 fn a_float_parameter_does_not_render_as_rust_infinity() {
4814 assert_eq!(fmt_float(f64::INFINITY), "'Infinity'");
4815 assert_eq!(fmt_float(f64::NEG_INFINITY), "'-Infinity'");
4816 assert_eq!(fmt_float(f64::NAN), "'NaN'");
4817 assert_eq!(fmt_float(3.0), "3", "a whole float should not gain a .0 tail");
4818 assert_eq!(fmt_float(3.5), "3.5");
4819 }
4820}