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 nql_string(s: &str) -> String {
679 format!("\"{}\"", s.replace('\\', "\\\\").replace('"', "\\\""))
680}
681
682fn sql_literals_to_nql(s: &str) -> String {
683 let mut out = String::with_capacity(s.len());
684 let mut it = s.chars().peekable();
685 while let Some(c) = it.next() {
686 match c {
687 '\'' => {
688 out.push('"');
689 while let Some(ch) = it.next() {
690 if ch == '\'' {
691 if it.peek() == Some(&'\'') {
692 it.next();
693 out.push('\''); } else {
695 break;
696 }
697 } else if ch == '"' {
698 out.push('\\');
701 out.push('"');
702 } else {
703 out.push(ch);
704 }
705 }
706 out.push('"');
707 }
708 '<' if it.peek() == Some(&'>') => { it.next(); out.push_str("!="); }
709 _ => out.push(c),
710 }
711 }
712 out
713}
714
715fn strip_prefix_ci(s: &str, prefix: &str) -> Option<String> {
716 if s.len() >= prefix.len() && s[..prefix.len()].eq_ignore_ascii_case(prefix) {
717 Some(s[prefix.len()..].trim_start().to_string())
718 } else {
719 None
720 }
721}
722
723fn find_kw(s: &str, kw: &str) -> Option<usize> {
726 let bytes = s.as_bytes();
727 let k = kw.as_bytes();
728 let mut depth = 0i32;
729 let mut in_s = false;
730 let mut in_d = false;
731 let mut i = 0usize;
732 while i < bytes.len() {
733 let c = bytes[i];
734 if in_s { if c == b'\'' { in_s = false; } i += 1; continue; }
735 if in_d { if c == b'"' { in_d = false; } i += 1; continue; }
736 match c {
737 b'\'' => { in_s = true; i += 1; continue; }
738 b'"' => { in_d = true; i += 1; continue; }
739 b'(' => { depth += 1; i += 1; continue; }
740 b')' => { depth -= 1; i += 1; continue; }
741 _ => {}
742 }
743 if depth == 0 && i + k.len() <= bytes.len()
744 && bytes[i..i + k.len()].eq_ignore_ascii_case(k)
745 {
746 let before_ok = i == 0 || !(bytes[i - 1] as char).is_alphanumeric() && bytes[i - 1] != b'_';
747 let after = i + k.len();
748 let after_ok = after >= bytes.len()
749 || !(bytes[after] as char).is_alphanumeric() && bytes[after] != b'_';
750 if before_ok && after_ok {
751 return Some(i);
752 }
753 }
754 i += 1;
755 }
756 None
757}
758
759fn split_top(s: &str, sep: char) -> Vec<String> {
763 let mut out = vec![];
764 let mut cur = String::new();
765 let mut depth = 0i32;
766 let mut in_s = false;
767 let mut it = s.chars().peekable();
768 while let Some(c) = it.next() {
769 if in_s {
770 cur.push(c);
771 if c == '\'' {
772 if it.peek() == Some(&'\'') { cur.push(it.next().unwrap()); } else { in_s = false; }
774 }
775 continue;
776 }
777 match c {
778 '\'' => { in_s = true; cur.push(c); }
779 '(' => { depth += 1; cur.push(c); }
780 ')' => { depth -= 1; cur.push(c); }
781 x if x == sep && depth == 0 => { out.push(cur.trim().to_string()); cur.clear(); }
782 _ => cur.push(c),
783 }
784 }
785 if !cur.trim().is_empty() { out.push(cur.trim().to_string()); }
786 out
787}
788
789fn sql_value(raw: &str) -> Result<Value, String> {
795 let t = raw.trim();
796 if t.is_empty() {
797 return Err("empty value".into());
798 }
799 let up = t.to_uppercase();
800 if up == "NULL" { return Ok(Value::Null); }
801 if up == "TRUE" { return Ok(Value::Bool(true)); }
802 if up == "FALSE" { return Ok(Value::Bool(false)); }
803 if t.starts_with('\'') && t.ends_with('\'') && t.len() >= 2 {
804 let inner = &t[1..t.len() - 1];
806 return Ok(Value::String(inner.replace("''", "'")));
807 }
808 if let Ok(i) = t.parse::<i64>() { return Ok(Value::from(i)); }
809 if let Ok(f) = t.parse::<f64>() { return Ok(Value::from(f)); }
810 Err(format!(
811 "cannot use {:?} as a value — this endpoint accepts string literals, \
812 numbers, TRUE/FALSE and NULL. Expressions, casts and function calls \
813 are not evaluated, because storing an unevaluated expression as text \
814 would be worse than refusing it", t))
815}
816
817fn split_returning(tail: &str) -> (String, Vec<Col>) {
819 let tu = tail.to_uppercase();
820 match find_kw(&tu, "RETURNING") {
821 None => (tail.to_string(), vec![]),
822 Some(at) => {
823 let head = tail[..at].trim().to_string();
824 let list = tail[at + "RETURNING".len()..].trim();
825 if list == "*" {
826 return (head, vec![]); }
828 let cols = split_top(list, ',')
829 .into_iter()
830 .map(|p| {
831 let raw = p.split_whitespace().next().unwrap_or(&p).to_string();
832 let name = raw.rsplit('.').next().unwrap_or(&raw).trim_matches('"').to_string();
833 Col::same(&name)
834 })
835 .collect();
836 (head, cols)
837 }
838 }
839}
840
841fn take_reserved(doc: &mut serde_json::Map<String, Value>) -> (Option<String>, Vec<String>, Option<String>, Option<String>) {
843 let id = doc.remove("_id").or_else(|| doc.remove("id"))
844 .and_then(|v| match v {
845 Value::String(s) => Some(s),
846 Value::Null => None,
847 other => Some(other.to_string()), });
849 let caused_by = match doc.remove("_caused_by") {
850 Some(Value::String(s)) => vec![s],
851 Some(Value::Array(a)) => a.into_iter()
852 .filter_map(|v| v.as_str().map(str::to_string)).collect(),
853 _ => vec![],
854 };
855 let vf = doc.remove("_valid_from").and_then(|v| v.as_str().map(str::to_string));
856 let vt = doc.remove("_valid_to").and_then(|v| v.as_str().map(str::to_string));
857 (id, caused_by, vf, vt)
858}
859
860fn translate_insert(sql: &str) -> Result<Stmt, String> {
862 let rest = strip_prefix_ci(sql, "INSERT")
863 .and_then(|r| strip_prefix_ci(&r, "INTO"))
864 .ok_or("expected INSERT INTO")?;
865 let ru = rest.to_uppercase();
869 let values_at = find_kw(&ru, "VALUES").ok_or(
870 "expected VALUES — `INSERT … SELECT` is not supported on this endpoint")?;
871 let head = rest[..values_at].trim().to_string();
872 let open = head.find('(').ok_or(
873 "INSERT needs an explicit column list — `INSERT INTO t (a, b) VALUES (…)`. \
874 NEDB is schemaless, so there is no declared column order to infer from")?;
875 let coll = head[..open].trim().trim_matches('"');
876 let coll = coll.rsplit('.').next().unwrap_or(coll).to_string();
877 if coll.is_empty() {
878 return Err("expected a collection name after INSERT INTO".into());
879 }
880 let close = head.rfind(')').ok_or("unterminated column list")?;
881 if close < open {
882 return Err("malformed column list".into());
883 }
884 let tail_from_values = rest[values_at..].to_string();
885 let cols: Vec<String> = split_top(&head[open + 1..close], ',')
886 .into_iter()
887 .map(|c| c.trim().trim_matches('"').to_string())
888 .collect();
889 if cols.is_empty() {
890 return Err("the column list is empty".into());
891 }
892
893 let after = strip_prefix_ci(&tail_from_values, "VALUES")
894 .ok_or("expected VALUES after the column list")?;
895 let (values_part, returning) = split_returning(&after);
896
897 let mut rows = vec![];
898 for group in split_top(&values_part, ',') {
899 let g = group.trim();
900 if !(g.starts_with('(') && g.ends_with(')')) {
901 return Err(format!("expected a parenthesised row of values, got {:?}", g));
902 }
903 let vals = split_top(&g[1..g.len() - 1], ',');
904 if vals.len() != cols.len() {
905 return Err(format!(
906 "{} values for {} columns — every row must match the column list",
907 vals.len(), cols.len()));
908 }
909 let mut doc = serde_json::Map::new();
910 for (c, v) in cols.iter().zip(vals.iter()) {
911 doc.insert(c.clone(), sql_value(v)?);
912 }
913 let (id, caused_by, valid_from, valid_to) = take_reserved(&mut doc);
914 rows.push(InsertRow { id, doc, caused_by, valid_from, valid_to });
915 }
916 if rows.is_empty() {
917 return Err("INSERT with no rows".into());
918 }
919 Ok(Stmt::Insert { coll, rows, returning })
920}
921
922fn translate_update(sql: &str) -> Result<Stmt, String> {
924 let rest = strip_prefix_ci(sql, "UPDATE").ok_or("expected UPDATE")?;
925 let ru = rest.to_uppercase();
926 let set_at = find_kw(&ru, "SET").ok_or("expected SET in UPDATE")?;
927 let target = rest[..set_at].trim();
930 let mut parts = target.split_whitespace();
931 let coll = parts.next().unwrap_or("").trim_matches('"');
932 let coll = coll.rsplit('.').next().unwrap_or(coll).to_string();
933 let upd_alias: Option<String> = match parts.next() {
934 Some(w) if w.eq_ignore_ascii_case("AS") => {
935 parts.next().map(|a| a.trim_matches('"').to_string())
936 }
937 Some(w) => Some(w.trim_matches('"').to_string()),
938 None => None,
939 };
940 if coll.is_empty() {
941 return Err("expected a collection name after UPDATE".into());
942 }
943 let after_set = rest[set_at + 3..].trim().to_string();
944 let (after_set, returning) = split_returning(&after_set);
945
946 let au = after_set.to_uppercase();
948 let (assigns_raw, where_raw) = match find_kw(&au, "WHERE") {
949 Some(at) => (after_set[..at].to_string(), after_set[at..].to_string()),
950 None => (after_set.clone(), String::new()),
951 };
952
953 let mut set = vec![];
954 for a in split_top(&assigns_raw, ',') {
955 let eq = a.find('=').ok_or(format!("expected `col = value` in SET, got {:?}", a))?;
956 let col = a[..eq].trim().trim_matches('"').to_string();
957 if col.is_empty() {
958 return Err("empty column name in SET".into());
959 }
960 set.push((col, sql_value(&a[eq + 1..])?));
961 }
962 if set.is_empty() {
963 return Err("UPDATE with no assignments".into());
964 }
965 let where_raw = strip_column_qualifiers(where_raw.trim(), &coll, upd_alias.as_deref())?;
968 let nql = format!("FROM {} {}", coll, sql_literals_to_nql(&where_raw))
969 .trim().to_string();
970 Ok(Stmt::Update { coll, set, nql, returning })
971}
972
973fn translate_delete(sql: &str) -> Result<Stmt, String> {
975 let rest = strip_prefix_ci(sql, "DELETE")
976 .and_then(|r| strip_prefix_ci(&r, "FROM"))
977 .ok_or("expected DELETE FROM")?;
978 let (rest, returning) = split_returning(&rest);
979 let end = rest.find(' ').unwrap_or(rest.len());
980 let coll = rest[..end].trim().trim_matches('"');
981 let coll = coll.rsplit('.').next().unwrap_or(coll).to_string();
982 if coll.is_empty() {
983 return Err("expected a collection name after DELETE FROM".into());
984 }
985 let (del_alias, where_raw) = split_table_alias(rest[end..].trim());
986 let where_raw = strip_column_qualifiers(where_raw, &coll, del_alias.as_deref())?;
987 let nql = format!("FROM {} {}", coll, sql_literals_to_nql(&where_raw))
988 .trim().to_string();
989 Ok(Stmt::Delete { coll, nql, returning })
990}
991
992pub fn translate(sql_raw: &str) -> Result<Stmt, String> {
994 let sql = normalise(sql_raw);
995 let sql = sql.trim().trim_end_matches(';').trim();
996 if sql.is_empty() {
997 return Ok(Stmt::Ok(""));
998 }
999 let upper = sql.to_uppercase();
1000
1001 if upper.starts_with("SET ") || upper.starts_with("BEGIN") || upper.starts_with("COMMIT")
1006 || upper.starts_with("ROLLBACK") || upper.starts_with("DISCARD")
1007 || upper.starts_with("LISTEN ") || upper.starts_with("UNLISTEN ")
1008 {
1009 return Ok(Stmt::Ok(if upper.starts_with("SET") { "SET" } else { "OK" }));
1011 }
1012 if upper.starts_with("SHOW ") {
1013 let name = sql[5..].trim().to_lowercase();
1014 let val = match name.as_str() {
1015 "transaction_isolation" | "default_transaction_isolation" => "read committed",
1016 "server_version" => SERVER_VERSION,
1017 "server_encoding" | "client_encoding" => "UTF8",
1018 "standard_conforming_strings" => "on",
1019 "is_superuser" => "off",
1020 _ => "",
1021 };
1022 return Ok(Stmt::Canned { cols: vec![name], row: vec![val.to_string()] });
1023 }
1024 if upper == "SELECT VERSION()" {
1025 return Ok(Stmt::Canned {
1026 cols: vec!["version".into()],
1027 row: vec![full_version_string()],
1028 });
1029 }
1030 if upper == "SELECT 1" || upper == "SELECT 1;" {
1031 return Ok(Stmt::Canned { cols: vec!["?column?".into()], row: vec!["1".into()] });
1032 }
1033 if upper.starts_with("SELECT CURRENT_SCHEMA") {
1034 return Ok(Stmt::Canned { cols: vec!["current_schema".into()], row: vec!["public".into()] });
1035 }
1036 if upper.starts_with("SELECT CURRENT_DATABASE") {
1037 return Ok(Stmt::Canned { cols: vec!["current_database".into()], row: vec!["nedb".into()] });
1038 }
1039 if upper.starts_with("SELECT CURRENT_USER") || upper.starts_with("SELECT USER") {
1040 return Ok(Stmt::Canned { cols: vec!["current_user".into()], row: vec!["nedb".into()] });
1041 }
1042
1043 if upper.starts_with("INSERT") { return translate_insert(sql); }
1047 if upper.starts_with("UPDATE") { return translate_update(sql); }
1048 if upper.starts_with("DELETE") { return translate_delete(sql); }
1049
1050 for (kw, why) in [
1052 ("CREATE", "DDL is not supported — collections are created implicitly by the first write to them, because NEDB is schemaless"),
1053 ("ALTER", "DDL is not supported — there is no schema to alter"),
1054 ("DROP", "DDL is not supported; drop a database with DELETE /v1/databases/<db>"),
1055 ("TRUNCATE", "not supported, and not an oversight: NEDB is append-only so that history cannot be discarded. That is the product"),
1056 ("COPY", "not supported; use GET /v1/databases/<db>/since for bulk export"),
1057 ("GRANT", "there is no SQL-level privilege system; auth is the bearer token"),
1058 ("REVOKE", "there is no SQL-level privilege system; auth is the bearer token"),
1059 ] {
1060 if upper.starts_with(kw) {
1061 return Err(format!("{} is not supported — {}", kw, why));
1062 }
1063 }
1064 if !upper.starts_with("SELECT") {
1065 return Err(format!(
1066 "only SELECT, INSERT, UPDATE and DELETE are supported on the Postgres \
1067 endpoint (got {:?})",
1068 sql.split_whitespace().next().unwrap_or("")
1069 ));
1070 }
1071 for (kw, why) in [
1072 (" JOIN ", "JOIN is not supported — NQL is single-collection; join in your client or model the relation with LINK/TRAVERSE"),
1073 (" UNION ", "UNION is not supported"),
1074 (" INTERSECT ", "INTERSECT is not supported"),
1075 (" EXCEPT ", "EXCEPT is not supported"),
1076 (" OVER (", "window functions are not supported"),
1077 ("DISTINCT ", "DISTINCT is not supported — GROUP BY <col> gives the distinct values with counts"),
1078 ] {
1079 if upper.contains(kw) {
1080 return Err(why.to_string());
1081 }
1082 }
1083 if find_kw(&upper, "FROM").is_none() {
1084 return Err("SELECT without FROM is not supported on this endpoint".into());
1085 }
1086
1087 let after_select = strip_prefix_ci(sql, "SELECT").ok_or("expected SELECT")?;
1089 let from_at = find_kw(&after_select.to_uppercase(), "FROM")
1090 .ok_or("expected FROM after the select list")?;
1091 let projection = after_select[..from_at].trim().to_string();
1092 let rest = after_select[from_at + 4..].trim().to_string();
1093 if rest.is_empty() {
1094 return Err("expected a collection name after FROM".into());
1095 }
1096 if rest.starts_with('(') {
1115 if let Some(flat) = flatten_count_of_subquery(&projection, &rest) {
1116 return translate(&flat);
1120 }
1121 return Err("subqueries in FROM are not supported — except \
1122 `SELECT count(*) FROM (…)`, which is rewritten to a flat \
1123 count when the inner query has no LIMIT, OFFSET, DISTINCT, \
1124 GROUP BY or aggregate of its own (any of those would make the \
1125 two counts different numbers)".into());
1126 }
1127 let coll_end = rest.find(' ').unwrap_or(rest.len());
1128 let coll = &rest[..coll_end];
1129 if coll.contains(',') {
1130 return Err("selecting from more than one collection is not supported (no JOIN)".into());
1131 }
1132 let bare = coll.rsplit('.').next().unwrap_or(coll).trim_matches('"');
1139 let qualified = coll
1140 .split('.')
1141 .map(|p| p.trim_matches('"'))
1142 .collect::<Vec<_>>()
1143 .join(".");
1144 let coll = if qualified.starts_with("information_schema.") {
1145 qualified.as_str()
1146 } else {
1147 bare
1148 };
1149 let tail = rest[coll_end..].trim();
1150
1151 let mut agg_clause = String::new();
1166 let mut agg_srcs: Vec<String> = vec![];
1167 let mut project: Vec<Col> = vec![];
1168
1169 if projection == "*" {
1170 } else {
1172 for part in split_top_level(&projection, ',') {
1173 let p = part.trim();
1174 if p.is_empty() {
1175 return Err("empty column in the select list".into());
1176 }
1177 let (expr, alias) = split_output_alias(p);
1178 let eu = expr.to_uppercase();
1179
1180 if eu.starts_with("COUNT(") {
1183 if agg_clause.is_empty() {
1184 agg_clause = " COUNT".to_string();
1185 }
1186 agg_srcs.push("count".to_string());
1187 project.push(Col::renamed("count", alias.unwrap_or("count")));
1188 continue;
1189 }
1190 if let Some(agg) = ["SUM", "AVG", "MIN", "MAX"]
1191 .iter()
1192 .find(|a| eu.starts_with(&format!("{}(", a)))
1193 {
1194 let inner = expr[agg.len() + 1..].trim_end_matches(')').trim();
1195 if inner.is_empty() || inner == "*" {
1196 return Err(format!("{}() needs a column", agg));
1197 }
1198 let inner = inner.rsplit('.').next().unwrap_or(inner).trim_matches('"');
1199 let named = format!("{} {}", agg, inner);
1200 if !agg_clause.is_empty() && agg_clause.trim() != "COUNT" && agg_clause.trim() != named {
1201 return Err(format!(
1202 "only one of SUM/AVG/MIN/MAX is supported per statement \
1203 (already have {:?}, then {:?}) — NQL's grouped row carries \
1204 the group key, `count`, and ONE named aggregate",
1205 agg_clause.trim(), named));
1206 }
1207 agg_clause = format!(" {}", named);
1208 let src = format!("{}_{}", agg.to_lowercase(), inner);
1211 project.push(Col::renamed(&src, alias.unwrap_or(&agg.to_lowercase())));
1212 agg_srcs.push(src);
1213 continue;
1214 }
1215 let bare = expr.rsplit('.').next().unwrap_or(expr).trim_matches('"');
1223 let is_column = !bare.is_empty()
1224 && !bare.starts_with(|c: char| c.is_ascii_digit())
1225 && bare.chars().all(|c| c.is_alphanumeric() || c == '_' || c == '$');
1226 if !is_column {
1227 return Err(format!(
1228 "expressions in the select list are not supported ({:?}) — \
1229 supported: *, a column list, COUNT(*), or SUM/AVG/MIN/MAX(col). \
1230 Compute it in your client, or read the column and map it there",
1231 p));
1232 }
1233 let name = bare;
1234 project.push(Col::renamed(name, alias.unwrap_or(name)));
1235 }
1236 }
1237
1238 let (alias, tail) = split_table_alias(tail);
1247 let mut tail = strip_column_qualifiers(tail, coll, alias.as_deref())?;
1248 let tu = tail.to_uppercase();
1249 if let Some(at) = find_kw(&tu, "AS OF SYSTEM TIME") {
1250 let before = tail[..at].to_string();
1251 let after = tail[at + "AS OF SYSTEM TIME".len()..].trim_start().to_string();
1252 let end = after.find(' ').unwrap_or(after.len());
1254 let seq = after[..end].trim().trim_matches('\'').trim_matches('"').to_string();
1255 if seq.parse::<u64>().is_err() {
1256 return Err(format!(
1257 "AS OF SYSTEM TIME takes a NEDB sequence number here, not a timestamp (got {:?}). \
1258 NEDB's history is sequence-addressed and never garbage-collected, so a seq is \
1259 exact where a wall-clock time would be approximate", seq));
1260 }
1261 tail = format!("{} AS OF {} {}", before.trim(), seq, after[end..].trim())
1262 .trim()
1263 .to_string();
1264 }
1265
1266 let tu_ord = tail.to_uppercase();
1279 if let Some(ob_at) = find_kw(&tu_ord, "ORDER BY") {
1280 let start = ob_at + "ORDER BY".len();
1281 let end = ["LIMIT", "OFFSET", "GROUP BY", "TRACE", "TRAVERSE", "SEARCH"]
1283 .iter()
1284 .filter_map(|k| find_kw(&tu_ord[start..], k).map(|at| start + at))
1285 .min()
1286 .unwrap_or(tail.len());
1287 let mut keys = vec![];
1288 for item in split_top_level(&tail[start..end], ',') {
1289 let item = item.trim();
1290 if item.is_empty() {
1291 continue;
1292 }
1293 let mut parts = item.split_whitespace();
1294 let first = parts.next().unwrap_or("");
1295 let rest: Vec<&str> = parts.collect();
1296 match first.parse::<usize>() {
1297 Ok(n) if n >= 1 => {
1298 let col = project.get(n - 1).ok_or_else(|| {
1299 if project.is_empty() {
1300 format!(
1301 "ORDER BY {} is a select-list POSITION, and `SELECT *` \
1302 has no list to index — name the column instead", n)
1303 } else {
1304 format!(
1305 "ORDER BY {} is out of range: the select list has {} \
1306 column(s)", n, project.len())
1307 }
1308 })?;
1309 keys.push(
1310 std::iter::once(col.src.as_str())
1311 .chain(rest.iter().copied())
1312 .collect::<Vec<_>>()
1313 .join(" "),
1314 );
1315 }
1316 _ => keys.push(item.to_string()),
1319 }
1320 }
1321 tail = format!("{} ORDER BY {} {}", &tail[..ob_at], keys.join(", "), &tail[end..])
1322 .split_whitespace()
1323 .collect::<Vec<_>>()
1324 .join(" ");
1325 }
1326
1327 let mut gkey: Option<String> = None;
1335 let tu_all = tail.to_uppercase();
1336 if let Some(gb_at) = find_kw(&tu_all, "GROUP BY") {
1337 let head = tail[..gb_at].trim_end().to_string();
1338 let after = tail[gb_at + "GROUP BY".len()..].trim_start();
1339 let key_end = after.find(|c: char| c == ' ' || c == ',').unwrap_or(after.len());
1340 let group_key = after[..key_end].trim().trim_matches('"').to_string();
1341 let after_key = after[key_end..].trim_start();
1342 gkey = Some(group_key.clone());
1343
1344 if after_key.starts_with(',') {
1348 return Err(format!(
1349 "GROUP BY takes one key here (got {:?} and more) — NQL groups by a \
1350 single field, and grouping by only the first would aggregate over \
1351 rows the query meant to keep apart",
1352 group_key));
1353 }
1354
1355 for c in &project {
1356 let ok = c.src == group_key
1357 || c.src == "count"
1358 || agg_srcs.contains(&c.src);
1359 if !ok {
1360 return Err(format!(
1361 "column {:?} must appear in the GROUP BY clause or be used in an \
1362 aggregate function — a grouped row carries the group key, `count`, \
1363 and the aggregate, nothing else",
1364 c.src));
1365 }
1366 }
1367
1368 tail = format!("{} GROUP BY {}{} {}", head, group_key, agg_clause, after_key)
1379 .split_whitespace()
1380 .collect::<Vec<_>>()
1381 .join(" ");
1382 agg_clause.clear();
1383 }
1384
1385 let tu_hav = tail.to_uppercase();
1400 if let Some(h_at) = find_kw(&tu_hav, "HAVING") {
1401 let start = h_at + "HAVING".len();
1402 let end = ["ORDER BY", "LIMIT", "OFFSET"]
1403 .iter()
1404 .filter_map(|k| find_kw(&tu_hav[start..], k).map(|at| start + at))
1405 .min()
1406 .unwrap_or(tail.len());
1407 let clause = tail[start..end].to_string();
1408 let lhs_end = clause
1410 .find(|c: char| "<>=!".contains(c))
1411 .unwrap_or(clause.len());
1412 let lhs = clause[..lhs_end].trim();
1413 if !lhs.is_empty() {
1414 let lu = lhs.to_uppercase();
1415 let is_count = lu == "COUNT" || lu.replace(' ', "") == "COUNT(*)"
1422 || project.iter().any(|c| c.out.eq_ignore_ascii_case(lhs) && c.src == "count");
1423 let mapped = if is_count {
1424 Some("count".to_string())
1425 } else {
1426 agg_srcs.iter().find(|s| s.eq_ignore_ascii_case(lhs)).cloned().or_else(|| {
1428 project.iter()
1429 .find(|c| c.out.eq_ignore_ascii_case(lhs) && agg_srcs.contains(&c.src))
1430 .map(|c| c.src.clone())
1431 })
1432 };
1433 match mapped {
1434 Some(m) => {
1435 let rewritten = format!("{} {}", m, clause[lhs_end..].trim());
1439 tail = format!("{} HAVING {} {}",
1440 tail[..h_at].trim(), rewritten.trim(), tail[end..].trim())
1441 .trim().to_string();
1442 }
1443 None if gkey.as_deref().map(|g| g.eq_ignore_ascii_case(lhs)) == Some(true) => {}
1444 None => {
1445 return Err(format!(
1446 "HAVING names {:?}, which this grouped row does not carry. \
1447 It has the group key{}{}. Filtering on anything else would \
1448 answer zero rows rather than report a mistake",
1449 lhs,
1450 gkey.as_deref().map(|g| format!(" ({:?})", g)).unwrap_or_default(),
1451 if agg_srcs.is_empty() { String::new() }
1452 else { format!(", plus {}", agg_srcs.join(", ")) }));
1453 }
1454 }
1455 }
1456 }
1457
1458
1459 let tail = sql_literals_to_nql(&tail);
1460 let nql = format!("FROM {}{}{}", coll,
1461 if agg_clause.is_empty() { String::new() } else { agg_clause },
1462 if tail.is_empty() { String::new() } else { format!(" {}", tail) });
1463
1464 Ok(Stmt::Query { nql: nql.trim().to_string(), project })
1465}
1466
1467const SERVER_VERSION: &str = "15.0";
1468
1469pub fn version_string() -> String {
1471 full_version_string()
1472}
1473
1474fn full_version_string() -> String {
1475 format!(
1476 "PostgreSQL {} (NEDB {}) — tamper-evident, append-only, permanent \
1477 history. SELECT + INSERT/UPDATE/DELETE; an UPDATE is a new version, \
1478 so prior values stay readable with AS OF SYSTEM TIME.",
1479 SERVER_VERSION,
1480 env!("CARGO_PKG_VERSION")
1481 )
1482}
1483
1484fn columns_for(rows: &[Value], project: &[Col]) -> Vec<Col> {
1492 if !project.is_empty() {
1493 return project.to_vec();
1494 }
1495 let mut plain: Vec<String> = vec![];
1496 let mut meta: Vec<String> = vec![];
1497 for r in rows {
1498 if let Value::Object(m) = r {
1499 for k in m.keys() {
1500 let target = if k.starts_with('_') { &mut meta } else { &mut plain };
1501 if !target.contains(k) {
1502 target.push(k.clone());
1503 }
1504 }
1505 }
1506 }
1507 meta.sort();
1514 plain.extend(meta);
1515 plain.into_iter().map(|k| Col::same(&k)).collect()
1516}
1517
1518fn oid_of_value(v: &Value) -> Option<i32> {
1520 match v {
1521 Value::Null => None,
1522 Value::Bool(_) => Some(OID_BOOL),
1523 Value::Number(n) => Some(if n.is_i64() || n.is_u64() { OID_INT8 } else { OID_FLOAT8 }),
1524 Value::String(_) => Some(OID_TEXT),
1525 _ => Some(OID_TEXT),
1527 }
1528}
1529
1530fn unify_oid(a: i32, b: i32) -> i32 {
1537 if a == b {
1538 return a;
1539 }
1540 match (a, b) {
1541 (OID_INT8, OID_FLOAT8) | (OID_FLOAT8, OID_INT8) => OID_FLOAT8,
1542 _ => OID_TEXT,
1543 }
1544}
1545
1546pub fn oid_for_column(rows: &[Value], col: &str) -> i32 {
1558 oid_for(rows, col)
1559}
1560
1561fn oid_for(rows: &[Value], col: &str) -> i32 {
1562 let mut acc: Option<i32> = None;
1563 for r in rows {
1564 if let Some(o) = r.get(col).and_then(oid_of_value) {
1565 acc = Some(match acc {
1566 None => o,
1567 Some(prev) => unify_oid(prev, o),
1568 });
1569 if acc == Some(OID_TEXT) {
1570 break; }
1572 }
1573 }
1574 acc.unwrap_or(OID_TEXT)
1575}
1576
1577fn cell(v: Option<&Value>) -> Option<String> {
1579 match v {
1580 None | Some(Value::Null) => None, Some(Value::String(s)) => Some(s.clone()),
1582 Some(Value::Bool(b)) => Some(if *b { "t".into() } else { "f".into() }),
1583 Some(other) => Some(other.to_string()),
1584 }
1585}
1586
1587fn cell_binary(v: Option<&Value>, oid: i32) -> Result<Option<Vec<u8>>, String> {
1599 let v = match v {
1600 None | Some(Value::Null) => return Ok(None),
1601 Some(v) => v,
1602 };
1603 let as_f64 = |n: &serde_json::Number| n.as_f64()
1604 .ok_or_else(|| "a number too large to send as float8".to_string());
1605 Ok(Some(match (oid, v) {
1606 (OID_BOOL, Value::Bool(b)) => vec![u8::from(*b)],
1607 (OID_INT2, Value::Number(n)) => {
1608 let i = n.as_i64().ok_or("not an integer")?;
1609 i16::try_from(i).map_err(|_| format!("{} does not fit in int2", i))?
1610 .to_be_bytes().to_vec()
1611 }
1612 (OID_INT4, Value::Number(n)) => {
1613 let i = n.as_i64().ok_or("not an integer")?;
1614 i32::try_from(i).map_err(|_| format!("{} does not fit in int4", i))?
1615 .to_be_bytes().to_vec()
1616 }
1617 (OID_INT8, Value::Number(n)) => {
1618 n.as_i64().ok_or("not an integer")?.to_be_bytes().to_vec()
1619 }
1620 (OID_FLOAT4, Value::Number(n)) => (as_f64(n)? as f32).to_be_bytes().to_vec(),
1621 (OID_FLOAT8, Value::Number(n)) => as_f64(n)?.to_be_bytes().to_vec(),
1622 (OID_TEXT | OID_VARCHAR | OID_NAME | OID_UNKNOWN | OID_JSON, _) => {
1624 cell(Some(v)).unwrap_or_default().into_bytes()
1625 }
1626 (OID_JSONB, _) => {
1628 let mut b = vec![1u8];
1629 b.extend_from_slice(cell(Some(v)).unwrap_or_default().as_bytes());
1630 b
1631 }
1632 (oid, val) => {
1633 let kind = match val {
1634 Value::Bool(_) => "a boolean",
1635 Value::Number(_) => "a number",
1636 Value::String(_) => "a string",
1637 Value::Array(_) => "an array",
1638 _ => "an object",
1639 };
1640 return Err(format!(
1641 "cannot send {} in binary format as type OID {} — the field holds \
1642 more than one type across documents, so it cannot be described \
1643 by a single Postgres type. Select it with a text cast, or use a \
1644 text-format client",
1645 kind, oid
1646 ));
1647 }
1648 }))
1649}
1650
1651fn row_description_fmt(cols: &[Col], oids: &[i32], fmts: &[i16]) -> Vec<u8> {
1653 let mut m = Out::msg(b'T');
1654 m.i16(cols.len() as i16);
1655 for (i, c) in cols.iter().enumerate() {
1656 m.cstr(&c.out);
1657 m.i32(0); m.i16((i + 1) as i16); m.i32(oids.get(i).copied().unwrap_or(OID_TEXT));
1660 m.i16(-1); m.i32(-1); m.i16(fmts.get(i).copied().unwrap_or(0));
1663 }
1664 m.finish()
1665}
1666
1667fn row_description(cols: &[Col], oids: &[i32]) -> Vec<u8> {
1668 row_description_fmt(cols, oids, &[])
1669}
1670
1671fn data_row_bytes(vals: &[Option<Vec<u8>>]) -> Vec<u8> {
1672 let mut m = Out::msg(b'D');
1673 m.i16(vals.len() as i16);
1674 for v in vals {
1675 match v {
1676 None => m.i32(-1),
1677 Some(b) => {
1678 m.i32(b.len() as i32);
1679 m.bytes(b);
1680 }
1681 }
1682 }
1683 m.finish()
1684}
1685
1686fn data_row(vals: &[Option<String>]) -> Vec<u8> {
1687 let owned: Vec<Option<Vec<u8>>> =
1688 vals.iter().map(|v| v.as_ref().map(|s| s.as_bytes().to_vec())).collect();
1689 data_row_bytes(&owned)
1690}
1691
1692pub fn encode_rows(rows: &[Value], project: &[Col]) -> Vec<u8> {
1702 let cols = columns_for(rows, project);
1703 let oids: Vec<i32> = cols.iter().map(|c| oid_for(rows, &c.src)).collect();
1704 let mut out = row_description(&cols, &oids);
1705 for r in rows {
1706 let vals: Vec<Option<String>> = cols.iter().map(|c| cell(r.get(&c.src))).collect();
1707 out.extend_from_slice(&data_row(&vals));
1708 }
1709 out
1710}
1711
1712pub fn encode_result(rows: &[Value], project: &[Col]) -> Vec<u8> {
1714 let mut out = encode_rows(rows, project);
1715 out.extend_from_slice(&command_complete(&format!("SELECT {}", rows.len())));
1716 out
1717}
1718
1719const OID_INT2: i32 = 21;
1746const OID_INT4: i32 = 23;
1747const OID_OID: i32 = 26;
1748const OID_FLOAT4: i32 = 700;
1749const OID_VARCHAR: i32 = 1043;
1750const OID_NAME: i32 = 19;
1751const OID_UNKNOWN: i32 = 705;
1752const OID_JSON: i32 = 114;
1753const OID_JSONB: i32 = 3802;
1754
1755fn param_count(sql: &str) -> usize {
1761 let b = sql.as_bytes();
1762 let mut i = 0usize;
1763 let mut in_s = false;
1764 let mut max = 0usize;
1765 while i < b.len() {
1766 let c = b[i];
1767 if in_s {
1768 if c == b'\'' {
1769 in_s = false;
1770 }
1771 i += 1;
1772 continue;
1773 }
1774 if c == b'\'' {
1775 in_s = true;
1776 i += 1;
1777 continue;
1778 }
1779 if c == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1780 let mut j = i + 1;
1781 let mut n = 0usize;
1782 while j < b.len() && b[j].is_ascii_digit() {
1783 n = n * 10 + (b[j] - b'0') as usize;
1784 j += 1;
1785 }
1786 max = max.max(n);
1787 i = j;
1788 continue;
1789 }
1790 i += 1;
1791 }
1792 max
1793}
1794
1795fn infer_field_oid(db: Option<&Arc<Db>>, coll: &str, field: &str) -> i32 {
1804 match field {
1808 "_seq" => return OID_INT8,
1809 "_id" | "_hash" | "_prev" | "_collection" | "_valid_from" | "_valid_to" => return OID_TEXT,
1810 _ => {}
1811 }
1812 if !field.is_empty() && crate::pgcatalog::is_catalog(coll) {
1818 if let Some(rows) = crate::pgcatalog::rows(coll, db) {
1819 return oid_for(&rows, field);
1820 }
1821 }
1822 let db = match db {
1823 Some(db) => db,
1824 None => return OID_TEXT,
1825 };
1826 if coll.is_empty() || field.is_empty() {
1827 return OID_TEXT;
1828 }
1829 let rows = match crate::nql::query(db, &format!("FROM {} LIMIT {}", coll, TYPE_SAMPLE)) {
1830 Ok((rows, _)) => rows,
1831 Err(_) => return OID_TEXT,
1832 };
1833 oid_for(&rows, field)
1837}
1838
1839fn aggregate_oid(src: &str, db: Option<&Arc<Db>>, coll: &str) -> Option<i32> {
1851 if src == "count" {
1852 return Some(OID_INT8);
1853 }
1854 for (prefix, fixed) in [
1855 ("count_", Some(OID_INT8)),
1856 ("avg_", Some(OID_FLOAT8)),
1857 ("sum_", None),
1858 ("min_", None),
1859 ("max_", None),
1860 ] {
1861 if let Some(field) = src.strip_prefix(prefix) {
1862 return Some(match fixed {
1863 Some(oid) => oid,
1864 None => match infer_field_oid(db, coll, field) {
1867 OID_INT8 => OID_INT8,
1868 OID_FLOAT8 => OID_FLOAT8,
1869 other => other,
1872 },
1873 });
1874 }
1875 }
1876 None
1877}
1878
1879const TYPE_SAMPLE: usize = 200;
1885
1886fn stmt_collection(sql: &str) -> String {
1888 let s = normalise(sql);
1889 let up = s.to_uppercase();
1890 let after = if let Some(at) = find_kw(&up, "FROM") {
1891 &s[at + 4..]
1892 } else if let Some(rest) = strip_prefix_ci(&s, "UPDATE") {
1893 return rest
1894 .split_whitespace()
1895 .next()
1896 .unwrap_or("")
1897 .rsplit('.')
1898 .next()
1899 .unwrap_or("")
1900 .trim_matches('"')
1901 .to_string();
1902 } else if let Some(rest) = strip_prefix_ci(&s, "INSERT INTO") {
1903 return rest
1904 .split(|c: char| c.is_whitespace() || c == '(')
1905 .find(|t| !t.is_empty())
1906 .unwrap_or("")
1907 .rsplit('.')
1908 .next()
1909 .unwrap_or("")
1910 .trim_matches('"')
1911 .to_string();
1912 } else {
1913 return String::new();
1914 };
1915 after
1916 .trim()
1917 .split(|c: char| c.is_whitespace())
1918 .find(|t| !t.is_empty())
1919 .unwrap_or("")
1920 .rsplit('.')
1921 .next()
1922 .unwrap_or("")
1923 .trim_matches('"')
1924 .to_string()
1925}
1926
1927fn param_fields(sql: &str, n_params: usize) -> Vec<Option<String>> {
1938 let s = normalise(sql);
1939 let mut out = vec![None; n_params];
1940
1941 let up = s.to_uppercase();
1944 if up.starts_with("INSERT") {
1945 if let (Some(open), Some(vals_at)) = (s.find('('), find_kw(&up, "VALUES")) {
1946 if open < vals_at {
1947 if let Some(close) = s[open..vals_at].rfind(')') {
1948 let cols: Vec<String> = split_top(&s[open + 1..open + close], ',')
1949 .into_iter()
1950 .map(|c| c.trim().trim_matches('"').to_string())
1951 .collect();
1952 let tail = &s[vals_at..];
1954 let mut seen = 0usize;
1955 let b = tail.as_bytes();
1956 let mut i = 0usize;
1957 let mut in_s = false;
1958 while i < b.len() {
1959 if in_s {
1960 if b[i] == b'\'' { in_s = false; }
1961 i += 1;
1962 continue;
1963 }
1964 if b[i] == b'\'' { in_s = true; i += 1; continue; }
1965 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1966 let mut j = i + 1;
1967 let mut num = 0usize;
1968 while j < b.len() && b[j].is_ascii_digit() {
1969 num = num * 10 + (b[j] - b'0') as usize;
1970 j += 1;
1971 }
1972 if num >= 1 && num <= n_params {
1973 if let Some(c) = cols.get(seen % cols.len().max(1)) {
1974 out[num - 1] = Some(c.clone());
1975 }
1976 }
1977 seen += 1;
1978 i = j;
1979 continue;
1980 }
1981 i += 1;
1982 }
1983 return out;
1984 }
1985 }
1986 }
1987 }
1988
1989 let b = s.as_bytes();
1991 let mut i = 0usize;
1992 let mut in_s = false;
1993 while i < b.len() {
1994 if in_s {
1995 if b[i] == b'\'' { in_s = false; }
1996 i += 1;
1997 continue;
1998 }
1999 if b[i] == b'\'' { in_s = true; i += 1; continue; }
2000 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
2001 let mut j = i + 1;
2002 let mut num = 0usize;
2003 while j < b.len() && b[j].is_ascii_digit() {
2004 num = num * 10 + (b[j] - b'0') as usize;
2005 j += 1;
2006 }
2007 if num >= 1 && num <= n_params {
2008 let left = &s[..i];
2009 let trimmed = left.trim_end_matches(|c: char| {
2012 c.is_whitespace() || "=<>!+-*/%(,".contains(c)
2013 });
2014 let mut tok = trimmed
2017 .rsplit(|c: char| c.is_whitespace() || c == '(' || c == ',')
2018 .find(|t| !t.is_empty())
2019 .unwrap_or("")
2020 .trim_matches('"');
2021 let mut before = trimmed;
2022 for _ in 0..4 {
2023 let upper_tok = tok.to_uppercase();
2024 if upper_tok.starts_with('$')
2030 || matches!(upper_tok.as_str(),
2031 "LIKE" | "ILIKE" | "IN" | "BETWEEN" | "AND" | "OR" | "NOT" | "IS") {
2032 before = before[..before.len() - tok.len()].trim_end_matches(|c: char| {
2033 c.is_whitespace() || "=<>!(,".contains(c)
2034 });
2035 tok = before
2036 .rsplit(|c: char| c.is_whitespace() || c == '(' || c == ',')
2037 .find(|t| !t.is_empty())
2038 .unwrap_or("")
2039 .trim_matches('"');
2040 } else {
2041 break;
2042 }
2043 }
2044 if !tok.is_empty()
2045 && tok.chars().all(|c| c.is_alphanumeric() || c == '_' || c == '.')
2046 && !tok.chars().next().map(|c| c.is_ascii_digit()).unwrap_or(true)
2047 {
2048 out[num - 1] = Some(tok.rsplit('.').next().unwrap_or(tok).to_string());
2049 }
2050 }
2051 i = j;
2052 continue;
2053 }
2054 i += 1;
2055 }
2056 out
2057}
2058
2059fn clause_param_oids(sql: &str, n_params: usize) -> Vec<Option<i32>> {
2069 let s = normalise(sql);
2070 let mut out = vec![None; n_params];
2071 let b = s.as_bytes();
2072 let mut i = 0usize;
2073 let mut in_s = false;
2074 while i < b.len() {
2075 if in_s {
2076 if b[i] == b'\'' { in_s = false; }
2077 i += 1;
2078 continue;
2079 }
2080 if b[i] == b'\'' { in_s = true; i += 1; continue; }
2081 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
2082 let mut j = i + 1;
2083 let mut num = 0usize;
2084 while j < b.len() && b[j].is_ascii_digit() {
2085 num = num * 10 + (b[j] - b'0') as usize;
2086 j += 1;
2087 }
2088 if num >= 1 && num <= n_params {
2089 let left = s[..i].trim_end().to_uppercase();
2090 out[num - 1] = if left.ends_with("VALID AS OF") {
2093 Some(OID_TEXT)
2094 } else if left.ends_with("AS OF SYSTEM TIME")
2095 || left.ends_with("FOR SYSTEM_TIME AS OF")
2096 || left.ends_with("AS OF")
2097 || left.ends_with("LIMIT")
2098 || left.ends_with("OFFSET")
2099 {
2100 Some(OID_INT8)
2101 } else {
2102 None
2103 };
2104 }
2105 i = j;
2106 continue;
2107 }
2108 i += 1;
2109 }
2110 out
2111}
2112
2113fn infer_param_oids(sql: &str, declared: &[i32], db: Option<&Arc<Db>>) -> Vec<i32> {
2119 let n = param_count(sql).max(declared.len());
2120 if n == 0 {
2121 return vec![];
2122 }
2123 let coll = stmt_collection(sql);
2124 let fields = param_fields(sql, n);
2125 let clauses = clause_param_oids(sql, n);
2126 (0..n)
2127 .map(|i| match declared.get(i) {
2128 Some(&oid) if oid != 0 => oid,
2129 _ => match clauses[i] {
2132 Some(oid) => oid,
2133 None => match &fields[i] {
2134 Some(f) => infer_field_oid(db, &coll, f),
2135 None => OID_TEXT,
2136 },
2137 },
2138 })
2139 .collect()
2140}
2141
2142fn decode_param(raw: Option<&[u8]>, oid: i32, format: i16) -> Result<Option<String>, String> {
2148 let bytes = match raw {
2149 None => return Ok(None),
2150 Some(b) => b,
2151 };
2152 let quote = |s: &str| format!("'{}'", s.replace('\'', "''"));
2153
2154 if format == 0 {
2155 let s = String::from_utf8_lossy(bytes).to_string();
2156 return Ok(Some(match oid {
2157 OID_BOOL => {
2158 let t = matches!(s.as_str(), "t" | "true" | "TRUE" | "1" | "yes" | "on");
2159 if t { "TRUE".into() } else { "FALSE".into() }
2160 }
2161 OID_INT2 | OID_INT4 | OID_INT8 | OID_OID | OID_FLOAT4 | OID_FLOAT8 => {
2162 if s.parse::<f64>().is_ok() { s } else { quote(&s) }
2166 }
2167 _ => quote(&s),
2172 }));
2173 }
2174 if format != 1 {
2175 return Err(format!("unsupported parameter format code {}", format));
2176 }
2177
2178 let need = |n: usize| -> Result<(), String> {
2180 if bytes.len() == n {
2181 Ok(())
2182 } else {
2183 Err(format!(
2184 "binary parameter of type OID {} should be {} bytes, got {}",
2185 oid, n, bytes.len()
2186 ))
2187 }
2188 };
2189 Ok(Some(match oid {
2190 OID_BOOL => {
2191 need(1)?;
2192 if bytes[0] != 0 { "TRUE".into() } else { "FALSE".into() }
2193 }
2194 OID_INT2 => {
2195 need(2)?;
2196 i16::from_be_bytes([bytes[0], bytes[1]]).to_string()
2197 }
2198 OID_INT4 => {
2199 need(4)?;
2200 i32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]).to_string()
2201 }
2202 OID_OID => {
2203 need(4)?;
2204 u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]).to_string()
2205 }
2206 OID_INT8 => {
2207 need(8)?;
2208 i64::from_be_bytes(bytes[..8].try_into().unwrap()).to_string()
2209 }
2210 OID_FLOAT4 => {
2211 need(4)?;
2212 let f = f32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]);
2213 fmt_float(f as f64)
2214 }
2215 OID_FLOAT8 => {
2216 need(8)?;
2217 fmt_float(f64::from_be_bytes(bytes[..8].try_into().unwrap()))
2218 }
2219 OID_TEXT | OID_VARCHAR | OID_NAME | OID_UNKNOWN | OID_JSON | 0 => {
2220 quote(&String::from_utf8_lossy(bytes))
2221 }
2222 OID_JSONB => {
2223 let body = if bytes.first() == Some(&1) { &bytes[1..] } else { bytes };
2225 quote(&String::from_utf8_lossy(body))
2226 }
2227 other => {
2228 return Err(format!(
2229 "parameter type OID {} is not supported in binary format — \
2230 the supported set is bool, int2/int4/int8, float4/float8, \
2231 text/varchar/json/jsonb. Send it as text, or cast it in the \
2232 statement",
2233 other
2234 ))
2235 }
2236 }))
2237}
2238
2239fn fmt_float(f: f64) -> String {
2241 if f.is_nan() {
2242 "'NaN'".into()
2243 } else if f.is_infinite() {
2244 if f > 0.0 { "'Infinity'".into() } else { "'-Infinity'".into() }
2245 } else if f.fract() == 0.0 && f.abs() < 1e15 {
2246 format!("{:.0}", f)
2247 } else {
2248 f.to_string()
2249 }
2250}
2251
2252fn substitute_params(sql: &str, params: &[Option<String>]) -> Result<String, String> {
2260 let b = sql.as_bytes();
2261 let mut out = String::with_capacity(sql.len() + 16);
2262 let mut i = 0usize;
2263 let mut in_s = false;
2264 while i < b.len() {
2265 let c = b[i];
2266 if in_s {
2267 out.push(c as char);
2268 if c == b'\'' { in_s = false; }
2269 i += 1;
2270 continue;
2271 }
2272 if c == b'\'' {
2273 in_s = true;
2274 out.push('\'');
2275 i += 1;
2276 continue;
2277 }
2278 if c == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
2279 let mut j = i + 1;
2280 let mut n = 0usize;
2281 while j < b.len() && b[j].is_ascii_digit() {
2282 n = n * 10 + (b[j] - b'0') as usize;
2283 j += 1;
2284 }
2285 match params.get(n.wrapping_sub(1)) {
2286 Some(Some(lit)) => out.push_str(lit),
2287 Some(None) => out.push_str("NULL"),
2288 None => {
2289 return Err(format!(
2290 "bind message supplies {} parameter(s) but the statement uses ${}",
2291 params.len(), n
2292 ))
2293 }
2294 }
2295 i = j;
2296 continue;
2297 }
2298 out.push(c as char);
2299 i += 1;
2300 }
2301 Ok(out)
2302}
2303
2304struct Prepared {
2306 sql: String,
2307 param_oids: Vec<i32>,
2310 out_shape: Option<Option<(Vec<Col>, Vec<i32>)>>,
2318}
2319
2320fn prepared_shape<'a>(
2322 p: &'a mut Prepared,
2323 db: Option<&Arc<Db>>,
2324) -> &'a Option<(Vec<Col>, Vec<i32>)> {
2325 if p.out_shape.is_none() {
2326 p.out_shape = Some(describe_shape(&p.sql, db, p.param_oids.len()));
2327 }
2328 p.out_shape.as_ref().expect("just filled")
2329}
2330
2331struct Portal {
2333 sql: String,
2334 result: Option<PortalResult>,
2340 frozen: Option<Vec<Col>>,
2350 formats: Vec<i16>,
2352 declared: Option<(Vec<Col>, Vec<i32>)>,
2359}
2360
2361impl Portal {
2362 fn format_of(&self, i: usize) -> i16 {
2365 match self.formats.len() {
2366 0 => 0,
2367 1 => self.formats[0],
2368 _ => self.formats.get(i).copied().unwrap_or(0),
2369 }
2370 }
2371 fn shape(&self, r: &PortalResult) -> (Vec<Col>, Vec<i32>) {
2373 match &self.declared {
2374 Some((cols, oids)) if self.formats.iter().any(|f| *f == 1) => {
2375 (cols.clone(), oids.clone())
2376 }
2377 _ => {
2378 let cols = columns_for(&r.rows, &r.project);
2379 let oids = cols.iter().map(|c| oid_for(&r.rows, &c.src)).collect();
2380 (cols, oids)
2381 }
2382 }
2383 }
2384}
2385
2386struct PortalResult {
2387 rows: Vec<Value>,
2388 project: Vec<Col>,
2389 has_rows: bool,
2390 tag: String,
2391 tag_counts_rows: bool,
2392 sent: usize,
2394}
2395
2396fn parse_complete() -> Vec<u8> { Out::msg(b'1').finish() }
2397fn bind_complete() -> Vec<u8> { Out::msg(b'2').finish() }
2398fn close_complete() -> Vec<u8> { Out::msg(b'3').finish() }
2399fn no_data() -> Vec<u8> { Out::msg(b'n').finish() }
2400fn portal_suspended() -> Vec<u8> { Out::msg(b's').finish() }
2401
2402fn parameter_description(oids: &[i32]) -> Vec<u8> {
2403 let mut m = Out::msg(b't');
2404 m.i16(oids.len() as i16);
2405 for o in oids {
2406 m.i32(*o);
2407 }
2408 m.finish()
2409}
2410
2411fn take_cstr(body: &[u8], at: &mut usize) -> String {
2413 let start = *at;
2414 while *at < body.len() && body[*at] != 0 {
2415 *at += 1;
2416 }
2417 let s = String::from_utf8_lossy(&body[start..*at]).to_string();
2418 if *at < body.len() {
2419 *at += 1; }
2421 s
2422}
2423
2424fn take_i16(body: &[u8], at: &mut usize) -> Result<i16, String> {
2425 if *at + 2 > body.len() {
2426 return Err("truncated message".into());
2427 }
2428 let v = i16::from_be_bytes([body[*at], body[*at + 1]]);
2429 *at += 2;
2430 Ok(v)
2431}
2432
2433fn take_i32(body: &[u8], at: &mut usize) -> Result<i32, String> {
2434 if *at + 4 > body.len() {
2435 return Err("truncated message".into());
2436 }
2437 let v = i32::from_be_bytes([body[*at], body[*at + 1], body[*at + 2], body[*at + 3]]);
2438 *at += 4;
2439 Ok(v)
2440}
2441
2442fn sample_columns(db: Option<&Arc<Db>>, coll: &str) -> Vec<Col> {
2448 let db = match db {
2449 Some(db) => db,
2450 None => return vec![],
2451 };
2452 let rows = match crate::nql::query(db, &format!("FROM {} LIMIT 25", coll)) {
2453 Ok((rows, _)) => rows,
2454 Err(_) => return vec![],
2455 };
2456 let mut names: Vec<String> = vec![];
2457 for r in &rows {
2458 if let Value::Object(m) = r {
2459 for k in m.keys() {
2460 if !names.iter().any(|n| n == k) {
2461 names.push(k.clone());
2462 }
2463 }
2464 }
2465 }
2466 names.sort();
2467 names.iter().map(|n| Col::same(n)).collect()
2468}
2469
2470fn describe_shape(
2478 sql: &str,
2479 db: Option<&Arc<Db>>,
2480 n_params: usize,
2481) -> Option<(Vec<Col>, Vec<i32>)> {
2482 let probe = probe_sql(sql, n_params);
2483
2484 if sql_engine_owns(&probe) {
2495 if let Ok(Some((done, _))) = try_catalog_select(&probe, db) {
2496 if done.project.is_empty() {
2497 return None;
2498 }
2499 let oids = done
2500 .project
2501 .iter()
2502 .map(|c| oid_for(&done.rows, &c.src))
2503 .collect();
2504 return Some((done.project, oids));
2505 }
2506 }
2507
2508 let stmt = translate(&probe).ok()?;
2509 let coll = stmt_collection(sql);
2510
2511 let cols = match stmt {
2512 Stmt::Ok(_) => return None,
2513 Stmt::Canned { cols, .. } => cols.iter().map(|c| Col::same(c)).collect(),
2514 Stmt::Query { project, .. } => {
2515 if project.is_empty() { sample_columns(db, &coll) } else { project }
2516 }
2517 Stmt::Insert { returning, .. } | Stmt::Update { returning, .. } | Stmt::Delete { returning, .. } => {
2518 if !wants_returning(sql) {
2519 return None;
2520 }
2521 if returning.is_empty() { sample_columns(db, &coll) } else { returning }
2522 }
2523 };
2524 if cols.is_empty() {
2525 return None;
2529 }
2530 let oids = cols
2531 .iter()
2532 .map(|c| {
2533 aggregate_oid(&c.src, db, &coll)
2534 .unwrap_or_else(|| infer_field_oid(db, &coll, &c.src))
2535 })
2536 .collect();
2537 Some((cols, oids))
2538}
2539
2540fn probe_sql(sql: &str, n_params: usize) -> String {
2548 let stub: Vec<Option<String>> = vec![Some("0".to_string()); n_params];
2549 substitute_params(sql, &stub).unwrap_or_else(|_| sql.to_string())
2550}
2551
2552fn ensure_executed(
2554 portal: &mut Portal,
2555 db_name: &str,
2556 db: Option<&Arc<Db>>,
2557 read_only: bool,
2558) -> Result<(), Vec<u8>> {
2559 if portal.result.is_some() {
2560 return Ok(());
2561 }
2562 let ex = execute_stmt(&portal.sql, db_name, db, read_only)?;
2563 let project = if let Some(f) = &portal.frozen {
2566 f.clone()
2567 } else {
2568 let p = if ex.project.is_empty() {
2569 columns_for(&ex.rows, &[])
2570 } else {
2571 ex.project.clone()
2572 };
2573 portal.frozen = Some(p.clone());
2574 p
2575 };
2576 portal.result = Some(PortalResult {
2577 rows: ex.rows,
2578 project,
2579 has_rows: ex.has_rows,
2580 tag: ex.tag,
2581 tag_counts_rows: ex.tag_counts_rows,
2582 sent: 0,
2583 });
2584 Ok(())
2585}
2586
2587async fn read_exact(sock: &mut TcpStream, n: usize) -> std::io::Result<Vec<u8>> {
2590 let mut buf = vec![0u8; n];
2591 sock.read_exact(&mut buf).await?;
2592 Ok(buf)
2593}
2594
2595async fn read_i32(sock: &mut TcpStream) -> std::io::Result<i32> {
2596 let b = read_exact(sock, 4).await?;
2597 Ok(i32::from_be_bytes([b[0], b[1], b[2], b[3]]))
2598}
2599
2600fn parse_startup_params(body: &[u8]) -> HashMap<String, String> {
2601 let mut out = HashMap::new();
2602 let mut parts = body.split(|b| *b == 0).map(|s| String::from_utf8_lossy(s).to_string());
2603 while let (Some(k), Some(v)) = (parts.next(), parts.next()) {
2604 if k.is_empty() {
2605 break;
2606 }
2607 out.insert(k, v);
2608 }
2609 out
2610}
2611
2612async fn handle(mut sock: TcpStream, resolver: Arc<dyn DbResolver>, read_only: bool) -> std::io::Result<()> {
2614 let params = loop {
2616 let len = read_i32(&mut sock).await?;
2617 if len < 8 || len > 1 << 20 {
2618 return Ok(()); }
2620 let code = read_i32(&mut sock).await?;
2621 let body = read_exact(&mut sock, (len - 8) as usize).await?;
2622 match code {
2623 SSL_REQUEST | GSS_REQUEST => {
2624 sock.write_all(b"N").await?;
2626 continue;
2627 }
2628 CANCEL_REQUEST => return Ok(()), PROTO_V3 => break parse_startup_params(&body),
2630 other => {
2631 let major = other >> 16;
2632 sock.write_all(&err_msg(
2633 "0A000",
2634 &format!("unsupported frontend protocol {}.{} — this endpoint speaks 3.0",
2635 major, other & 0xffff),
2636 )).await?;
2637 return Ok(());
2638 }
2639 }
2640 };
2641
2642 let db_name = params.get("database").cloned().unwrap_or_default();
2643
2644 let resolved: Option<Arc<Db>> = {
2650 let r = Arc::clone(&resolver);
2651 let name = db_name.clone();
2652 tokio::task::spawn_blocking(move || r.resolve(&name))
2653 .await
2654 .unwrap_or(None)
2655 };
2656
2657 if let Some(expected) = resolver.token() {
2659 let mut m = Out::msg(b'R');
2661 m.i32(3);
2662 sock.write_all(&m.finish()).await?;
2663
2664 let tag = read_exact(&mut sock, 1).await?;
2665 if tag[0] != b'p' {
2666 sock.write_all(&err_msg("28000", "expected a password message")).await?;
2667 return Ok(());
2668 }
2669 let len = read_i32(&mut sock).await?;
2670 if len < 4 || len > 1 << 16 {
2671 return Ok(());
2672 }
2673 let body = read_exact(&mut sock, (len - 4) as usize).await?;
2674 let supplied = String::from_utf8_lossy(&body).trim_end_matches('\0').to_string();
2675 let ok = supplied.len() == expected.len()
2677 && supplied.bytes().zip(expected.bytes()).fold(0u8, |a, (x, y)| a | (x ^ y)) == 0;
2678 if !ok {
2679 sock.write_all(&err_msg("28P01", "password authentication failed")).await?;
2680 return Ok(());
2681 }
2682 }
2683
2684 let mut m = Out::msg(b'R');
2685 m.i32(0); sock.write_all(&m.finish()).await?;
2687
2688 for (k, v) in [
2689 ("server_version", SERVER_VERSION),
2690 ("server_encoding", "UTF8"),
2691 ("client_encoding", "UTF8"),
2692 ("DateStyle", "ISO, MDY"),
2693 ("integer_datetimes", "on"),
2694 ("standard_conforming_strings", "on"),
2695 ("application_name", "nedbd"),
2696 ] {
2697 let mut p = Out::msg(b'S');
2698 p.cstr(k);
2699 p.cstr(v);
2700 sock.write_all(&p.finish()).await?;
2701 }
2702 let mut k = Out::msg(b'K');
2703 k.i32(std::process::id() as i32);
2704 k.i32(0);
2705 sock.write_all(&k.finish()).await?;
2706 sock.write_all(&ready()).await?;
2707
2708 let mut prepared: HashMap<String, Prepared> = HashMap::new();
2714 let mut portals: HashMap<String, Portal> = HashMap::new();
2715 let mut failed = false;
2719
2720 loop {
2721 let mut tag = [0u8; 1];
2722 if sock.read_exact(&mut tag).await.is_err() {
2723 return Ok(()); }
2725 let len = read_i32(&mut sock).await?;
2726 if len < 4 || len > 64 << 20 {
2727 return Ok(());
2728 }
2729 let body = read_exact(&mut sock, (len - 4) as usize).await?;
2730
2731 if failed && tag[0] != b'S' && tag[0] != b'X' {
2733 continue;
2734 }
2735
2736 match tag[0] {
2737 b'X' => return Ok(()), b'Q' => {
2740 let sql = String::from_utf8_lossy(&body).trim_end_matches('\0').to_string();
2741 let out = run_simple_query(&sql, &db_name, resolved.as_ref(), read_only);
2742 sock.write_all(&out).await?;
2743 sock.write_all(&ready()).await?;
2744 portals.remove("");
2746 }
2747
2748 b'P' => {
2750 let mut at = 0usize;
2751 let name = take_cstr(&body, &mut at);
2752 let sql = take_cstr(&body, &mut at);
2753 let n = take_i16(&body, &mut at).unwrap_or(0).max(0) as usize;
2754 let mut declared = Vec::with_capacity(n);
2755 let mut bad = false;
2756 for _ in 0..n {
2757 match take_i32(&body, &mut at) {
2758 Ok(o) => declared.push(o),
2759 Err(_) => { bad = true; break; }
2760 }
2761 }
2762 if bad {
2763 sock.write_all(&err_msg("08P01", "malformed Parse message")).await?;
2764 failed = true;
2765 continue;
2766 }
2767 let probe = probe_sql(&sql, param_count(&sql));
2777 if !sql_engine_owns(&probe) {
2778 if let Err(why) = translate(&probe) {
2779 sock.write_all(&err_msg("0A000", &why)).await?;
2780 failed = true;
2781 continue;
2782 }
2783 }
2784 let param_oids = infer_param_oids(&sql, &declared, resolved.as_ref());
2785 prepared.insert(name, Prepared { sql, param_oids, out_shape: None });
2786 sock.write_all(&parse_complete()).await?;
2787 }
2788
2789 b'B' => {
2791 let mut at = 0usize;
2792 let portal_name = take_cstr(&body, &mut at);
2793 let stmt_name = take_cstr(&body, &mut at);
2794 if !prepared.contains_key(&stmt_name) {
2795 sock.write_all(&err_msg("26000", &format!(
2796 "prepared statement {:?} does not exist", stmt_name))).await?;
2797 failed = true;
2798 continue;
2799 }
2800 let p = &prepared[&stmt_name];
2801 let mut want_formats: Vec<i16> = vec![];
2802 let res: Result<String, String> = (|| {
2803 let nfmt = take_i16(&body, &mut at)? .max(0) as usize;
2804 let mut fmts = Vec::with_capacity(nfmt);
2805 for _ in 0..nfmt {
2806 fmts.push(take_i16(&body, &mut at)?);
2807 }
2808 let nparam = take_i16(&body, &mut at)?.max(0) as usize;
2809 let mut vals: Vec<Option<String>> = Vec::with_capacity(nparam);
2810 for i in 0..nparam {
2811 let l = take_i32(&body, &mut at)?;
2812 let raw: Option<Vec<u8>> = if l < 0 {
2813 None
2814 } else {
2815 let l = l as usize;
2816 if at + l > body.len() {
2817 return Err("truncated Bind parameter".into());
2818 }
2819 let v = body[at..at + l].to_vec();
2820 at += l;
2821 Some(v)
2822 };
2823 let f = match fmts.len() {
2826 0 => 0,
2827 1 => fmts[0],
2828 _ => *fmts.get(i).unwrap_or(&0),
2829 };
2830 let oid = *p.param_oids.get(i).unwrap_or(&OID_TEXT);
2831 vals.push(decode_param(raw.as_deref(), oid, f)?);
2832 }
2833 let nres = take_i16(&body, &mut at)?.max(0) as usize;
2838 for _ in 0..nres {
2839 let f = take_i16(&body, &mut at)?;
2840 if f != 0 && f != 1 {
2841 return Err(format!("unknown result format code {}", f));
2842 }
2843 want_formats.push(f);
2844 }
2845 substitute_params(&p.sql, &vals)
2846 })();
2847 match res {
2848 Ok(sql) => {
2849 let declared = if want_formats.iter().any(|f| *f == 1) {
2852 let p = prepared.get_mut(&stmt_name).expect("checked above");
2853 prepared_shape(p, resolved.as_ref()).clone()
2854 } else {
2855 None
2856 };
2857 portals.insert(portal_name, Portal {
2858 sql, result: None, frozen: None,
2859 formats: want_formats, declared,
2860 });
2861 sock.write_all(&bind_complete()).await?;
2862 }
2863 Err(why) => {
2864 sock.write_all(&err_msg("08P01", &why)).await?;
2865 failed = true;
2866 }
2867 }
2868 }
2869
2870 b'D' => {
2872 let kind = body.first().copied().unwrap_or(b'S');
2873 let mut at = 1usize;
2874 let name = take_cstr(&body, &mut at);
2875 if kind == b'S' {
2876 if !prepared.contains_key(&name) {
2877 sock.write_all(&err_msg("26000", &format!(
2878 "prepared statement {:?} does not exist", name))).await?;
2879 failed = true;
2880 continue;
2881 }
2882 let p = prepared.get_mut(&name).expect("checked above");
2883 let oids = p.param_oids.clone();
2884 sock.write_all(¶meter_description(&oids)).await?;
2887 let out = match prepared_shape(p, resolved.as_ref()) {
2891 Some((cols, col_oids)) => row_description(cols, col_oids),
2892 None => no_data(),
2893 };
2894 sock.write_all(&out).await?;
2895 } else {
2896 let portal = match portals.get_mut(&name) {
2897 Some(p) => p,
2898 None => {
2899 sock.write_all(&err_msg("34000", &format!(
2900 "portal {:?} does not exist", name))).await?;
2901 failed = true;
2902 continue;
2903 }
2904 };
2905 match ensure_executed(portal, &db_name, resolved.as_ref(), read_only) {
2910 Err(encoded) => {
2911 sock.write_all(&encoded).await?;
2912 failed = true;
2913 }
2914 Ok(()) => {
2915 let r = portal.result.as_ref().expect("just executed");
2916 if !r.has_rows {
2917 sock.write_all(&no_data()).await?;
2918 } else {
2919 let (cols, oids) = portal.shape(r);
2920 let fmts: Vec<i16> =
2921 (0..cols.len()).map(|i| portal.format_of(i)).collect();
2922 sock.write_all(&row_description_fmt(&cols, &oids, &fmts)).await?;
2923 }
2924 }
2925 }
2926 }
2927 }
2928
2929 b'E' => {
2931 let mut at = 0usize;
2932 let name = take_cstr(&body, &mut at);
2933 let max_rows = take_i32(&body, &mut at).unwrap_or(0);
2934 let portal = match portals.get_mut(&name) {
2935 Some(p) => p,
2936 None => {
2937 sock.write_all(&err_msg("34000", &format!(
2938 "portal {:?} does not exist", name))).await?;
2939 failed = true;
2940 continue;
2941 }
2942 };
2943 if let Err(encoded) = ensure_executed(portal, &db_name, resolved.as_ref(), read_only) {
2944 sock.write_all(&encoded).await?;
2945 failed = true;
2946 continue;
2947 }
2948 let r = portal.result.as_ref().expect("just executed");
2949 if !r.has_rows {
2950 let tag = r.tag.clone();
2951 sock.write_all(&command_complete(&tag)).await?;
2952 continue;
2953 }
2954 let (cols, oids) = portal.shape(r);
2955 let limit = if max_rows > 0 {
2956 (r.sent + max_rows as usize).min(r.rows.len())
2957 } else {
2958 r.rows.len()
2959 };
2960 let mut encoded: Vec<Vec<u8>> = Vec::with_capacity(limit - r.sent);
2965 let mut fail: Option<String> = None;
2966 for row in &r.rows[r.sent..limit] {
2967 let mut vals: Vec<Option<Vec<u8>>> = Vec::with_capacity(cols.len());
2968 for (i, c) in cols.iter().enumerate() {
2969 let v = row.get(&c.src);
2970 let got = if portal.format_of(i) == 1 {
2971 cell_binary(v, oids.get(i).copied().unwrap_or(OID_TEXT))
2972 .map_err(|e| format!("column {:?}: {}", c.out, e))
2973 } else {
2974 Ok(cell(v).map(|s| s.into_bytes()))
2975 };
2976 match got {
2977 Ok(b) => vals.push(b),
2978 Err(e) => { fail = Some(e); break; }
2979 }
2980 }
2981 if fail.is_some() {
2982 break;
2983 }
2984 encoded.push(data_row_bytes(&vals));
2985 }
2986 if let Some(why) = fail {
2987 sock.write_all(&err_msg("22P03", &why)).await?;
2988 failed = true;
2989 continue;
2990 }
2991 let mut out = vec![];
2992 for e in &encoded {
2993 out.extend_from_slice(e);
2994 }
2995 let r = portal.result.as_mut().expect("just executed");
2996 r.sent = limit;
2997 if max_rows > 0 && r.sent < r.rows.len() {
3001 out.extend_from_slice(&portal_suspended());
3002 } else {
3003 let tag = if r.tag_counts_rows {
3004 format!("{} {}", r.tag, r.sent)
3005 } else {
3006 r.tag.clone()
3007 };
3008 out.extend_from_slice(&command_complete(&tag));
3009 }
3010 sock.write_all(&out).await?;
3011 }
3012
3013 b'C' => {
3015 let kind = body.first().copied().unwrap_or(b'S');
3016 let mut at = 1usize;
3017 let name = take_cstr(&body, &mut at);
3018 if kind == b'S' {
3019 prepared.remove(&name);
3020 } else {
3021 portals.remove(&name);
3022 }
3023 sock.write_all(&close_complete()).await?;
3026 }
3027
3028 b'H' => {}
3032
3033 b'S' => {
3034 failed = false;
3035 sock.write_all(&ready()).await?;
3036 }
3037
3038 other => {
3039 sock.write_all(&err_msg(
3040 "08P01",
3041 &format!("unexpected frontend message {:?}", other as char),
3042 )).await?;
3043 failed = true;
3044 }
3045 }
3046 }
3047}
3048
3049const READ_ONLY_MSG: &str =
3050 "this endpoint is running read-only (NEDBD_PG_READ_ONLY=1). Writes are \
3051 implemented but disabled on this server — unset the flag to allow them.";
3052
3053fn no_db(db_name: &str) -> Vec<u8> {
3054 err_msg("3D000", &format!(
3055 "database {:?} is not open on this server — create it first \
3056 (POST /v1/databases), or connect with -d <name>", db_name))
3057}
3058
3059fn catalog_name(n: &str) -> String {
3063 let joined: Vec<&str> = n.split('.').collect();
3064 if joined.len() >= 2 && joined[joined.len() - 2] == "information_schema" {
3065 format!("information_schema.{}", joined[joined.len() - 1])
3066 } else {
3067 joined[joined.len() - 1].to_string()
3068 }
3069}
3070
3071fn sql_engine_for_collections() -> bool {
3109 use std::sync::OnceLock;
3110 static ON: OnceLock<bool> = OnceLock::new();
3111 *ON.get_or_init(|| {
3112 matches!(std::env::var("NEDBD_SQL_ENGINE").as_deref(), Ok("1") | Ok("true") | Ok("on"))
3113 })
3114}
3115
3116fn sql_engine_owns(sql: &str) -> bool {
3117 let Ok(sel) = crate::sqlselect::parse(sql) else { return false };
3118 let touched = sel.base_relations();
3119 if touched.is_empty() {
3120 return translate(sql).is_err();
3121 }
3122 if touched.iter().any(|t| crate::pgcatalog::is_catalog(&catalog_name(t))) {
3123 return true;
3124 }
3125 sql_engine_for_collections()
3131}
3132
3133fn try_catalog_select(
3146 sql: &str,
3147 db: Option<&Arc<Db>>,
3148) -> Result<Option<(Executed, crate::sqlplan::Plan)>, Vec<u8>> {
3149 let sel = match crate::sqlselect::parse(sql) {
3150 Ok(sel) => sel,
3151 Err(why) => {
3152 if mentions_catalog(sql) {
3161 return Err(err_msg("0A000", &format!(
3162 "this catalogue query uses SQL this endpoint does not \
3163 implement: {}", why)));
3164 }
3165 return Ok(None);
3166 }
3167 };
3168
3169 if !sql_engine_owns(sql) {
3175 return Ok(None);
3176 }
3177
3178 let temporal: std::collections::HashMap<String, u64> = {
3205 let mut seen: std::collections::HashMap<String, Vec<Option<u64>>> =
3210 std::collections::HashMap::new();
3211 for t in sel.from.iter().chain(sel.joins.iter().map(|j| &j.table)) {
3212 seen.entry(catalog_name(&t.name).to_ascii_lowercase())
3213 .or_default()
3214 .push(t.as_of);
3215 }
3216 let mut out: std::collections::HashMap<String, u64> = std::collections::HashMap::new();
3217 for (key, ats) in &seen {
3218 let mut distinct: Vec<Option<u64>> = ats.clone();
3219 distinct.sort();
3220 distinct.dedup();
3221 match distinct.as_slice() {
3222 [Some(seq)] => {
3224 out.insert(key.clone(), *seq);
3225 }
3226 [None] => {}
3227 _ => {
3230 let mixed_tip = distinct.contains(&None);
3231 let seqs: Vec<String> =
3232 distinct.iter().flatten().map(|s| s.to_string()).collect();
3233 let detail = if mixed_tip {
3234 format!(
3235 "at the tip and AS OF {}",
3236 seqs.join(" and "))
3237 } else {
3238 format!("AS OF {}", seqs.join(" and "))
3239 };
3240 return Err(err_msg("0A000", &format!(
3241 "{:?} is read {} in one statement. This endpoint reads each \
3242 collection once per statement, so it cannot serve both — and \
3243 answering from either one would silently return the same rows for \
3244 both arms, which is the comparison failing to be a comparison. Ask \
3245 the two questions separately.",
3246 key, detail)));
3247 }
3248 }
3249 }
3250 out
3251 };
3252
3253 let nql_verbs: std::collections::HashMap<String, (Option<String>, Option<String>)> = {
3260 let mut out: std::collections::HashMap<String, (Option<String>, Option<String>)> =
3261 std::collections::HashMap::new();
3262 for t in sel.from.iter().chain(sel.joins.iter().map(|j| &j.table)) {
3263 let k = catalog_name(&t.name).to_ascii_lowercase();
3264 let e = out.entry(k.clone()).or_default();
3265 for (slot, incoming, verb) in [
3266 (&mut e.0, &t.valid_as_of, "VALID AS OF"),
3267 (&mut e.1, &t.search, "SEARCH"),
3268 ] {
3269 match (slot.as_deref(), incoming.as_deref()) {
3270 (Some(a), Some(b)) if a != b => {
3271 return Err(err_msg("0A000", &format!(
3272 "{:?} is read with two different {} arguments in one statement \
3273 ({:?} and {:?}). This endpoint reads each collection once, so \
3274 it cannot serve both. Ask the two questions separately.",
3275 k, verb, a, b)));
3276 }
3277 (None, Some(b)) => *slot = Some(b.to_string()),
3278 _ => {}
3279 }
3280 }
3281 }
3282 out
3283 };
3284
3285 let pushdown_prefilters: std::collections::HashMap<String, String> = {
3286 let refs: Vec<&crate::sqlselect::TableRef> = sel
3287 .from
3288 .iter()
3289 .chain(sel.joins.iter().map(|j| &j.table))
3290 .collect();
3291 let bindings: Vec<String> = refs.iter().map(|t| t.binding()).collect();
3292 let nullable = crate::sqlpush::nullable_bindings(&sel);
3293 let mut out = std::collections::HashMap::new();
3294 let mut ambiguous: Vec<String> = vec![];
3295 for t in &refs {
3296 let key = catalog_name(&t.name).to_ascii_lowercase();
3297 if out.contains_key(&key) || ambiguous.contains(&key) {
3298 out.remove(&key);
3299 ambiguous.push(key);
3300 continue;
3301 }
3302 if let Some(p) = crate::sqlpush::nql_prefilter(
3303 sel.where_.as_ref(), &t.binding(), &bindings, &nullable) {
3304 out.insert(key, p);
3305 }
3306 }
3307 out
3308 };
3309
3310 let resolve = |name: &str| -> anyhow::Result<Option<Box<dyn crate::sqlselect::Relation>>> {
3311 let cname = catalog_name(name);
3312 {
3325 let k = cname.to_ascii_lowercase();
3326 let bad = if temporal.contains_key(&k) {
3327 Some("AS OF SYSTEM TIME")
3328 } else {
3329 match nql_verbs.get(&k) {
3330 Some((Some(_), _)) => Some("VALID AS OF"),
3331 Some((_, Some(_))) => Some("SEARCH"),
3332 _ => None,
3333 }
3334 };
3335 if let Some(clause) = bad {
3336 if crate::pgcatalog::is_catalog(&cname) {
3337 anyhow::bail!(
3338 "{} is not supported on the catalogue relation {:?} — a catalogue is \
3339 synthesised from the store's current shape rather than read from the \
3340 log, so it has no history to reach and no document text to search. \
3341 Ignoring the clause would answer your question with present-day rows \
3342 and look like it worked",
3343 clause, cname);
3344 }
3345 }
3346 }
3347 if let Some(rows) = crate::pgcatalog::rows(&cname, db) {
3348 return Ok(Some(crate::sqlselect::from_vec(rows)));
3352 }
3353 let key = cname.to_ascii_lowercase();
3373 let pre = pushdown_prefilters.get(&key);
3374 let nql = {
3384 let mut q = format!("FROM {}", cname);
3385 if let Some(seq) = temporal.get(&key) {
3386 q.push_str(&format!(" AS OF {}", seq));
3387 }
3388 if let Some(d) = nql_verbs.get(&key).and_then(|v| v.0.as_deref()) {
3389 q.push_str(&format!(" VALID AS OF {}", nql_string(d)));
3390 }
3391 if let Some(p) = pre {
3392 q.push_str(&format!(" WHERE {}", p));
3393 }
3394 if let Some(t) = nql_verbs.get(&key).and_then(|v| v.1.as_deref()) {
3395 q.push_str(&format!(" SEARCH {}", nql_string(t)));
3396 }
3397 q
3398 };
3399 match db {
3400 Some(db) => match crate::nql::query(db, &nql) {
3401 Ok((rows, _)) => Ok(Some(crate::sqlselect::from_vec(rows))),
3402 Err(_) if pre.is_some() => {
3408 let bare = match temporal.get(&key) {
3415 Some(seq) => format!("FROM {} AS OF {}", cname, seq),
3416 None => format!("FROM {}", cname),
3417 };
3418 match crate::nql::query(db, &bare) {
3419 Ok((rows, _)) => Ok(Some(crate::sqlselect::from_vec(rows))),
3420 Err(_) => Ok(None),
3421 }
3422 }
3423 Err(_) => Ok(None),
3424 },
3425 None => Ok(None),
3426 }
3427 };
3428
3429 let (cols, rows, plan) = crate::sqlselect::execute_explain(
3430 &sel,
3431 &resolve,
3432 crate::sqljoin::JoinExec::Auto,
3433 )
3434 .map_err(|e| err_msg("42601", &e.to_string()))?;
3435
3436 Ok(Some((
3437 Executed {
3438 rows,
3439 project: cols
3443 .iter()
3444 .map(|c| Col::renamed(&c.key, &c.name))
3445 .collect(),
3446 has_rows: true,
3447 tag: "SELECT".into(),
3448 tag_counts_rows: true,
3449 },
3450 plan,
3451 )))
3452}
3453
3454fn strip_explain(sql: &str) -> Option<&str> {
3461 let t = sql.trim().trim_end_matches(';').trim();
3462 let mut rest = t.strip_prefix("EXPLAIN").or_else(|| t.strip_prefix("explain"))?;
3463 if !rest.starts_with(char::is_whitespace) {
3465 return None;
3466 }
3467 rest = rest.trim_start();
3468 loop {
3469 let low = rest.to_lowercase();
3470 if let Some(r) = low.strip_prefix("analyze").or_else(|| low.strip_prefix("analyse")) {
3471 if r.starts_with(char::is_whitespace) || r.is_empty() {
3472 rest = rest[rest.len() - r.len()..].trim_start();
3473 continue;
3474 }
3475 }
3476 if let Some(r) = low.strip_prefix("verbose") {
3477 if r.starts_with(char::is_whitespace) || r.is_empty() {
3478 rest = rest[rest.len() - r.len()..].trim_start();
3479 continue;
3480 }
3481 }
3482 break;
3483 }
3484 Some(rest)
3485}
3486
3487fn plan_result(lines: Vec<String>) -> Executed {
3490 Executed {
3491 rows: lines
3492 .into_iter()
3493 .map(|l| serde_json::json!({ "QUERY PLAN": l }))
3494 .collect(),
3495 project: vec![Col::same("QUERY PLAN")],
3496 has_rows: true,
3497 tag: "EXPLAIN".into(),
3498 tag_counts_rows: false,
3499 }
3500}
3501
3502fn mentions_catalog(sql: &str) -> bool {
3509 let low = sql.to_lowercase();
3510 low.contains("pg_catalog.")
3511 || low.contains("information_schema.")
3512 || low.contains("from pg_")
3513 || low.contains("join pg_")
3514}
3515
3516fn catalog_target(nql: &str) -> Option<String> {
3521 let coll = crate::nql::parse(nql).ok()?.coll;
3522 if crate::pgcatalog::is_catalog(&coll) {
3523 Some(coll)
3524 } else {
3525 None
3526 }
3527}
3528
3529fn wants_returning(sql: &str) -> bool {
3533 find_kw(&sql.to_uppercase(), "RETURNING").is_some()
3534}
3535
3536fn next_row_id() -> String {
3538 use std::sync::atomic::{AtomicU64, Ordering};
3539 static N: AtomicU64 = AtomicU64::new(0);
3540 let n = N.fetch_add(1, Ordering::Relaxed);
3541 let ts = std::time::SystemTime::now()
3542 .duration_since(std::time::UNIX_EPOCH)
3543 .map(|d| d.as_micros())
3544 .unwrap_or(0);
3545 format!("r{}{}", ts, n)
3546}
3547
3548pub struct Executed {
3556 pub rows: Vec<Value>,
3558 pub project: Vec<Col>,
3560 pub has_rows: bool,
3564 pub tag: String,
3567 pub tag_counts_rows: bool,
3569}
3570
3571impl Executed {
3572 fn nothing(tag: &str) -> Self {
3573 Executed { rows: vec![], project: vec![], has_rows: false, tag: tag.to_string(), tag_counts_rows: false }
3574 }
3575 fn tag_for(&self, sent: usize) -> String {
3577 if self.tag_counts_rows { format!("{} {}", self.tag, sent) } else { self.tag.clone() }
3578 }
3579}
3580
3581fn execute_stmt(
3587 stmt_sql: &str,
3588 db_name: &str,
3589 db: Option<&Arc<Db>>,
3590 read_only: bool,
3591) -> Result<Executed, Vec<u8>> {
3592 if let Some(inner) = strip_explain(stmt_sql) {
3601 if let Some((_, plan)) = try_catalog_select(inner, db)? {
3602 return Ok(plan_result(plan.render()));
3603 }
3604 let mut lines = vec![];
3605 match translate(inner) {
3606 Ok(_) => {
3607 lines.push(
3608 "NQL path — this statement is translated to NQL and \
3609 executed by the storage engine, not by the SQL evaluator."
3610 .to_string(),
3611 );
3612 lines.push(
3613 "No plan is reported, because the SQL evaluator is not \
3614 what runs it. Reporting one would describe a pipeline \
3615 that never executed."
3616 .to_string(),
3617 );
3618 lines.push(
3619 "The SQL evaluator (joins, CASE, scalar functions, a \
3620 hash-join planner) currently serves catalogue queries."
3621 .to_string(),
3622 );
3623 }
3624 Err(why) => lines.push(format!("cannot be executed: {why}")),
3625 }
3626 return Ok(plan_result(lines));
3627 }
3628
3629 if let Some((done, _plan)) = try_catalog_select(stmt_sql, db)? {
3630 return Ok(done);
3631 }
3632
3633 let stmt = translate(stmt_sql).map_err(|why| err_msg("0A000", &why))?;
3634
3635 macro_rules! need_db {
3638 () => {
3639 match db {
3640 Some(db) => db,
3641 None => return Err(no_db(db_name)),
3642 }
3643 };
3644 }
3645 macro_rules! need_write {
3646 () => {
3647 if read_only {
3648 return Err(err_msg("25006", READ_ONLY_MSG));
3649 }
3650 };
3651 }
3652
3653 match stmt {
3654 Stmt::Ok(tag) => Ok(Executed::nothing(if tag.is_empty() { "SELECT 0" } else { tag })),
3655
3656 Stmt::Canned { cols, row } => {
3657 let mut obj = serde_json::Map::new();
3660 for (c, v) in cols.iter().zip(row.iter()) {
3661 obj.insert(c.clone(), Value::String(v.clone()));
3662 }
3663 Ok(Executed {
3664 rows: vec![Value::Object(obj)],
3665 project: cols.iter().map(|c| Col::same(c)).collect(),
3666 has_rows: true,
3667 tag: "SELECT".into(),
3668 tag_counts_rows: true,
3669 })
3670 }
3671
3672 Stmt::Query { nql, project } => {
3673 if let Some(coll) = catalog_target(&nql) {
3683 let rows = crate::pgcatalog::rows(&coll, db)
3684 .expect("catalog_target only returns names pgcatalog serves");
3685 let rows = crate::nql::query_rows(rows, &nql)
3686 .map_err(|e| err_msg("42601", &e.to_string()))?;
3687 return Ok(Executed {
3688 rows, project, has_rows: true,
3689 tag: "SELECT".into(), tag_counts_rows: true,
3690 });
3691 }
3692 let db = need_db!();
3693 let (rows, _) = crate::nql::query(db, &nql).map_err(|e| {
3694 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3695 })?;
3696 Ok(Executed { rows, project, has_rows: true, tag: "SELECT".into(), tag_counts_rows: true })
3697 }
3698
3699 Stmt::Insert { coll, rows, returning } => {
3700 let db = need_db!();
3701 need_write!();
3702 let mut written: Vec<Value> = vec![];
3703 for (i, r) in rows.iter().enumerate() {
3704 let id = match &r.id {
3708 Some(id) => id.clone(),
3709 None => format!("{}-{}", next_row_id(), i),
3710 };
3711 let node = db
3712 .put(&coll, &id, Value::Object(r.doc.clone()),
3713 r.caused_by.clone(), r.valid_from.clone(), r.valid_to.clone())
3714 .map_err(|e| err_msg("XX000", &format!("INSERT failed: {}", e)))?;
3715 written.push(crate::nql::node_to_json(&node));
3716 }
3717 let n = written.len();
3718 let has_rows = wants_returning(stmt_sql);
3719 Ok(Executed {
3720 rows: if has_rows { written } else { vec![] },
3721 project: returning,
3722 has_rows,
3723 tag: format!("INSERT 0 {}", n),
3725 tag_counts_rows: false,
3726 })
3727 }
3728
3729 Stmt::Update { coll, set, nql, returning } => {
3730 let db = need_db!();
3731 need_write!();
3732 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
3735 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3736 })?;
3737 let mut written: Vec<Value> = vec![];
3738 for row in &matched {
3739 let id = match row.get("_id").and_then(|v| v.as_str()) {
3740 Some(id) => id.to_string(),
3741 None => continue,
3742 };
3743 let mut doc = match db.get(&coll, &id) {
3747 Some(n) => match n.data {
3748 Value::Object(m) => m,
3749 _ => serde_json::Map::new(),
3750 },
3751 None => continue,
3752 };
3753 for (k, v) in &set {
3754 doc.insert(k.clone(), v.clone());
3755 }
3756 let node = db
3759 .put(&coll, &id, Value::Object(doc), vec![], None, None)
3760 .map_err(|e| err_msg("XX000", &format!("UPDATE failed: {}", e)))?;
3761 written.push(crate::nql::node_to_json(&node));
3762 }
3763 let n = written.len();
3764 let has_rows = wants_returning(stmt_sql);
3765 Ok(Executed {
3766 rows: if has_rows { written } else { vec![] },
3767 project: returning,
3768 has_rows,
3769 tag: format!("UPDATE {}", n),
3770 tag_counts_rows: false,
3771 })
3772 }
3773
3774 Stmt::Delete { coll, nql, returning } => {
3775 let db = need_db!();
3776 need_write!();
3777 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
3778 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3779 })?;
3780 let returned = matched.clone();
3783 let mut n = 0usize;
3784 for row in &matched {
3785 if let Some(id) = row.get("_id").and_then(|v| v.as_str()) {
3786 match db.delete(&coll, id) {
3787 Ok(true) => n += 1,
3788 Ok(false) => {}
3789 Err(e) => return Err(err_msg("XX000", &format!("DELETE failed: {}", e))),
3790 }
3791 }
3792 }
3793 let has_rows = wants_returning(stmt_sql);
3794 Ok(Executed {
3795 rows: if has_rows { returned } else { vec![] },
3796 project: returning,
3797 has_rows,
3798 tag: format!("DELETE {}", n),
3799 tag_counts_rows: false,
3800 })
3801 }
3802 }
3803}
3804
3805fn run_simple_query(sql: &str, db_name: &str, db: Option<&Arc<Db>>, read_only: bool) -> Vec<u8> {
3807 let mut out = vec![];
3808 let statements = split_statements(sql);
3809 if statements.is_empty() {
3810 return Out::msg(b'I').finish();
3812 }
3813 for stmt_sql in statements {
3814 match execute_stmt(&stmt_sql, db_name, db, read_only) {
3815 Err(encoded) => {
3817 out.extend_from_slice(&encoded);
3818 return out;
3819 }
3820 Ok(ex) => {
3821 if ex.has_rows {
3822 out.extend_from_slice(&encode_rows(&ex.rows, &ex.project));
3823 }
3824 out.extend_from_slice(&command_complete(&ex.tag_for(ex.rows.len())));
3825 }
3826 }
3827 }
3828 out
3829}
3830
3831fn split_statements(sql: &str) -> Vec<String> {
3833 let mut out = vec![];
3834 let mut cur = String::new();
3835 let mut in_s = false;
3836 for c in sql.chars() {
3837 match c {
3838 '\'' => { in_s = !in_s; cur.push(c); }
3839 ';' if !in_s => {
3840 if !cur.trim().is_empty() { out.push(cur.clone()); }
3841 cur.clear();
3842 }
3843 _ => cur.push(c),
3844 }
3845 }
3846 if !cur.trim().is_empty() {
3847 out.push(cur);
3848 }
3849 out
3850}
3851
3852pub async fn run(host: &str, port: u16, resolver: Arc<dyn DbResolver>) -> anyhow::Result<()> {
3854 let read_only = std::env::var("NEDBD_PG_READ_ONLY")
3858 .map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
3859 .unwrap_or(false);
3860 let listener = TcpListener::bind((host, port)).await?;
3861 println!(" pgwire postgres endpoint on {}:{} — psql / DBeaver / psycopg ({})",
3862 host, port,
3863 if read_only { "SELECT only — read-only mode" } else { "SELECT + INSERT/UPDATE/DELETE" });
3864 loop {
3865 let (sock, _peer) = match listener.accept().await {
3866 Ok(v) => v,
3867 Err(e) => {
3868 eprintln!(" [pgwire] accept failed: {}", e);
3869 continue;
3870 }
3871 };
3872 let r = Arc::clone(&resolver);
3873 tokio::spawn(async move {
3874 let _ = sock.set_nodelay(true);
3875 if let Err(e) = handle(sock, r, read_only).await {
3876 if e.kind() != std::io::ErrorKind::UnexpectedEof
3878 && e.kind() != std::io::ErrorKind::ConnectionReset
3879 {
3880 eprintln!(" [pgwire] connection error: {}", e);
3881 }
3882 }
3883 });
3884 }
3885}
3886
3887#[cfg(test)]
3890mod explain_tests {
3891 use super::*;
3892
3893 #[test]
3894 fn a_bare_explain_is_stripped() {
3895 assert_eq!(strip_explain("EXPLAIN SELECT 1"), Some("SELECT 1"));
3896 assert_eq!(strip_explain("explain select 1"), Some("select 1"));
3897 assert_eq!(strip_explain(" EXPLAIN SELECT 1 ; "), Some("SELECT 1"));
3898 }
3899
3900 #[test]
3901 fn analyze_and_verbose_are_accepted_and_ignored() {
3902 assert_eq!(strip_explain("EXPLAIN ANALYZE SELECT 1"), Some("SELECT 1"));
3906 assert_eq!(strip_explain("EXPLAIN ANALYSE SELECT 1"), Some("SELECT 1"));
3907 assert_eq!(strip_explain("EXPLAIN VERBOSE SELECT 1"), Some("SELECT 1"));
3908 assert_eq!(strip_explain("EXPLAIN ANALYZE VERBOSE SELECT 1"), Some("SELECT 1"));
3909 assert_eq!(strip_explain("explain analyze verbose select 1"), Some("select 1"));
3910 }
3911
3912 #[test]
3913 fn a_word_merely_starting_with_explain_is_not_a_keyword() {
3914 assert_eq!(strip_explain("EXPLAINED SELECT 1"), None);
3915 assert_eq!(strip_explain("SELECT 1"), None);
3916 assert_eq!(strip_explain("SELECT explain FROM t"), None);
3917 }
3918
3919 #[test]
3920 fn a_column_named_analyze_is_not_eaten() {
3921 assert_eq!(strip_explain("EXPLAIN analyzed_view"), Some("analyzed_view"));
3924 }
3925
3926 #[test]
3927 fn the_plan_result_has_postgres_shape() {
3928 let e = plan_result(vec!["Seq Scan on t".into(), "note".into()]);
3929 assert_eq!(e.project.len(), 1);
3930 assert_eq!(e.project[0].out, "QUERY PLAN");
3931 assert_eq!(e.rows.len(), 2);
3932 assert_eq!(e.rows[0]["QUERY PLAN"], "Seq Scan on t");
3933 assert_eq!(e.tag, "EXPLAIN");
3934 assert!(!e.tag_counts_rows);
3936 }
3937}
3938
3939#[cfg(test)]
3940mod tests {
3941 use super::*;
3942 use serde_json::json;
3943
3944 fn q(sql: &str) -> String {
3945 match translate(sql) {
3946 Ok(Stmt::Query { nql, .. }) => nql,
3947 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3948 }
3949 }
3950 fn proj(sql: &str) -> Vec<String> {
3952 match translate(sql) {
3953 Ok(Stmt::Query { project, .. }) => project.iter().map(|c| c.out.clone()).collect(),
3954 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3955 }
3956 }
3957 fn proj_pairs(sql: &str) -> Vec<(String, String)> {
3959 match translate(sql) {
3960 Ok(Stmt::Query { project, .. }) =>
3961 project.iter().map(|c| (c.src.clone(), c.out.clone())).collect(),
3962 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3963 }
3964 }
3965 fn names(cols: &[Col]) -> Vec<String> { cols.iter().map(|c| c.out.clone()).collect() }
3966
3967 fn cols_of(sql: &str) -> Vec<Col> {
3971 match translate(sql).unwrap() {
3972 Stmt::Query { project, .. } => project,
3973 other => panic!("{:?}", other),
3974 }
3975 }
3976
3977 #[test]
3978 fn select_star_becomes_bare_from() {
3979 assert_eq!(q("SELECT * FROM orders"), "FROM orders");
3980 assert_eq!(q("select * from orders;"), "FROM orders");
3981 assert_eq!(proj("SELECT * FROM orders"), Vec::<String>::new());
3982 }
3983
3984 #[test]
3985 fn a_column_list_becomes_a_projection_not_a_clause() {
3986 assert_eq!(q("SELECT status, total FROM orders"), "FROM orders");
3989 assert_eq!(proj("SELECT status, total FROM orders"), vec!["status", "total"]);
3990 }
3991
3992 #[test]
3993 fn a_qualifier_reduces_to_the_field_while_an_ALIAS_is_the_name_the_client_sees() {
3994 let cols = cols_of("SELECT o.status AS s, o.total total, o.region FROM orders o");
4002 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
4003 vec!["status", "total", "region"]);
4004 assert_eq!(cols.iter().map(|c| c.out.clone()).collect::<Vec<_>>(),
4005 vec!["s", "total", "region"]);
4006 assert_eq!(q("SELECT * FROM public.orders"), "FROM orders");
4007 assert_eq!(q("SELECT * FROM \"orders\""), "FROM orders");
4008 }
4009
4010 #[test]
4011 fn a_select_list_may_MIX_columns_with_an_aggregate() {
4012 assert_eq!(q("SELECT status, count(*) AS count_1 FROM orders GROUP BY status"),
4020 "FROM orders GROUP BY status COUNT");
4021 assert_eq!(q("SELECT status, count(*) FROM orders WHERE total > 1 GROUP BY status ORDER BY status LIMIT 5"),
4024 "FROM orders WHERE total > 1 GROUP BY status COUNT ORDER BY status LIMIT 5");
4025 assert_eq!(q("SELECT count(*) FROM orders"), "FROM orders COUNT");
4027 assert_eq!(q("SELECT sum(total) FROM orders"), "FROM orders SUM total");
4028 let e = translate("SELECT status, count(*) FROM orders GROUP BY status, region").unwrap_err();
4032 assert!(e.contains("GROUP BY takes one key"), "{}", e);
4033 let cols = cols_of("SELECT status, count(*) AS count_1 FROM orders GROUP BY status");
4034 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
4035 vec!["status", "count"]);
4036 assert_eq!(cols.iter().map(|c| c.out.clone()).collect::<Vec<_>>(),
4037 vec!["status", "count_1"]);
4038
4039 let cols = cols_of("SELECT status, count(*), sum(total) FROM orders GROUP BY status");
4042 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
4043 vec!["status", "count", "sum_total"]);
4044 assert_eq!(q("SELECT status, count(*), sum(total) FROM orders GROUP BY status"),
4045 "FROM orders GROUP BY status SUM total");
4046
4047 assert_eq!(q("SELECT o.status, sum(o.total) FROM orders o GROUP BY o.status"),
4049 "FROM orders GROUP BY status SUM total");
4050
4051 let e = translate("SELECT status, sum(total), avg(total) FROM orders GROUP BY status")
4054 .unwrap_err();
4055 assert!(e.contains("only one of SUM/AVG/MIN/MAX"), "{}", e);
4056
4057 let e = translate("SELECT status, total, count(*) FROM orders GROUP BY status")
4059 .unwrap_err();
4060 assert!(e.contains("must appear in the GROUP BY clause"), "{}", e);
4061 }
4062
4063 #[test]
4064 fn ORDER_BY_an_ordinal_resolves_to_that_select_list_column() {
4065 assert_eq!(q("SELECT status, total FROM orders ORDER BY 1"),
4070 "FROM orders ORDER BY status");
4071 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 DESC"),
4072 "FROM orders ORDER BY total DESC");
4073 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 DESC, 1"),
4075 "FROM orders ORDER BY total DESC, status");
4076 assert_eq!(q("SELECT status, total FROM orders ORDER BY 1, total DESC"),
4077 "FROM orders ORDER BY status, total DESC");
4078 assert_eq!(q("SELECT status, count(*) AS n FROM orders GROUP BY status ORDER BY 1"),
4082 "FROM orders GROUP BY status COUNT ORDER BY status");
4083 assert_eq!(q("SELECT status, count(*) AS n FROM orders GROUP BY status ORDER BY 2 DESC"),
4085 "FROM orders GROUP BY status COUNT ORDER BY count DESC");
4086 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 LIMIT 1"),
4089 "FROM orders ORDER BY total LIMIT 1");
4090 assert_eq!(q("SELECT status FROM orders WHERE total > 1 ORDER BY 1"),
4092 "FROM orders WHERE total > 1 ORDER BY status");
4093
4094 let e = translate("SELECT status FROM orders ORDER BY 4").unwrap_err();
4098 assert!(e.contains("out of range") && e.contains("1 column"), "{}", e);
4099 let e = translate("SELECT * FROM orders ORDER BY 1").unwrap_err();
4100 assert!(e.contains("no list to index"), "{}", e);
4101 }
4102
4103 #[test]
4104 fn count_of_a_subquery_flattens_only_when_the_two_counts_MUST_agree() {
4105 assert_eq!(
4109 q("SELECT count(*) AS count_1 FROM (SELECT orders._id AS a, orders.status AS b \
4110 FROM orders WHERE orders.status = 'paid') AS anon_1"),
4111 r#"FROM orders COUNT WHERE status = "paid""#);
4114 assert_eq!(q("SELECT count(*) FROM (SELECT orders._id FROM orders) AS anon_1"),
4116 "FROM orders COUNT");
4117 assert_eq!(q("SELECT count(*) FROM (SELECT _id FROM orders ORDER BY total DESC) AS a"),
4119 "FROM orders COUNT");
4120 let cols = cols_of("SELECT count(*) AS count_1 FROM (SELECT _id FROM orders) AS a");
4122 assert_eq!(cols[0].src, "count");
4123 assert_eq!(cols[0].out, "count_1");
4124
4125 for sql in [
4128 "SELECT count(*) FROM (SELECT _id FROM orders LIMIT 1) AS a",
4130 "SELECT count(*) FROM (SELECT _id FROM orders OFFSET 1) AS a",
4131 "SELECT count(*) FROM (SELECT status FROM orders GROUP BY status) AS a",
4133 "SELECT count(*) FROM (SELECT count(*) FROM orders) AS a",
4135 "SELECT count(*) FROM (SELECT sum(total) FROM orders) AS a",
4136 "SELECT count(*), status FROM (SELECT status FROM orders) AS a",
4138 "SELECT status FROM (SELECT status FROM orders) AS a",
4139 "SELECT count(*) FROM (SELECT x FROM (SELECT _id AS x FROM orders) AS b) AS a",
4141 ] {
4142 let e = translate(sql).unwrap_err();
4143 assert!(e.contains("subqueries in FROM"), "{} -> {}", sql, e);
4144 }
4145
4146 for (sql, needle) in [
4151 ("SELECT count(*) FROM (SELECT DISTINCT status FROM orders) AS a", "DISTINCT"),
4152 ("SELECT count(*) FROM (SELECT a FROM t UNION SELECT b FROM u) AS x", "UNION"),
4153 ] {
4154 let e = translate(sql).unwrap_err();
4155 assert!(e.contains(needle), "{} -> {}", sql, e);
4156 }
4157 }
4158
4159 #[test]
4160 fn a_QUALIFIED_column_in_WHERE_finds_its_field_instead_of_ZERO_ROWS() {
4161 assert_eq!(q("SELECT _id FROM orders WHERE orders.status = 'paid'"),
4169 r#"FROM orders WHERE status = "paid""#);
4170 assert_eq!(q("SELECT _id FROM orders WHERE orders.total > 50"),
4171 "FROM orders WHERE total > 50");
4172 assert_eq!(q("SELECT _id FROM orders ORDER BY orders.total DESC LIMIT 2"),
4174 "FROM orders ORDER BY total DESC LIMIT 2");
4175 assert_eq!(q("SELECT status, count(*) FROM orders GROUP BY orders.status"),
4176 "FROM orders GROUP BY status COUNT");
4177
4178 assert_eq!(q("SELECT o.status FROM orders o WHERE o.status = 'paid'"),
4182 r#"FROM orders WHERE status = "paid""#);
4183 assert_eq!(q("SELECT o.status FROM orders AS o WHERE o.total > 1"),
4184 "FROM orders WHERE total > 1");
4185
4186 let e = translate("SELECT _id FROM orders WHERE nosuch.status = 'paid'").unwrap_err();
4191 assert!(e.contains("no table or alias named \"nosuch\""), "{}", e);
4192 let e = translate("SELECT _id FROM orders o WHERE p.status = 'paid'").unwrap_err();
4193 assert!(e.contains("aliased \"o\""), "the message names the alias in scope: {}", e);
4194
4195 assert_eq!(q("SELECT _id FROM orders WHERE status = 'pa.id'"),
4197 r#"FROM orders WHERE status = "pa.id""#);
4198 assert_eq!(q("SELECT _id FROM orders WHERE total > 1.5"),
4200 "FROM orders WHERE total > 1.5");
4201
4202 match translate("UPDATE orders o SET status = 'x' WHERE o.total > 5").unwrap() {
4204 Stmt::Update { coll, nql, .. } => {
4205 assert_eq!(coll, "orders", "the alias is not part of the collection name");
4206 assert_eq!(nql, "FROM orders WHERE total > 5");
4207 }
4208 other => panic!("{:?}", other),
4209 }
4210 match translate("DELETE FROM orders o WHERE o.status = 'paid'").unwrap() {
4211 Stmt::Delete { coll, nql, .. } => {
4212 assert_eq!(coll, "orders");
4213 assert_eq!(nql, r#"FROM orders WHERE status = "paid""#);
4214 }
4215 other => panic!("{:?}", other),
4216 }
4217
4218 assert_eq!(q("SELECT _id FROM orders AS OF SYSTEM TIME 3 WHERE orders.total > 1"),
4220 "FROM orders AS OF 3 WHERE total > 1");
4221 }
4222
4223 #[test]
4224 fn where_clauses_pass_through_with_sql_literals_rewritten() {
4225 assert_eq!(q("SELECT * FROM orders WHERE status = 'paid'"),
4226 r#"FROM orders WHERE status = "paid""#);
4227 assert_eq!(q("SELECT * FROM orders WHERE status <> 'paid'"),
4228 r#"FROM orders WHERE status != "paid""#);
4229 assert_eq!(q("SELECT * FROM orders WHERE status IN ('paid','open')"),
4230 r#"FROM orders WHERE status IN ("paid","open")"#);
4231 }
4232
4233 #[test]
4236 fn a_doubled_sql_quote_is_one_literal_character() {
4237 assert_eq!(q("SELECT * FROM t WHERE name = 'it''s'"),
4238 r#"FROM t WHERE name = "it's""#);
4239 }
4240
4241 #[test]
4244 fn a_double_quote_inside_a_sql_literal_is_escaped_for_nql() {
4245 assert_eq!(q(r#"SELECT * FROM t WHERE name = 'say "hi"'"#),
4246 r#"FROM t WHERE name = "say \"hi\"""#);
4247 }
4248
4249 #[test]
4250 fn the_shared_clauses_are_handed_to_nql_unchanged() {
4251 assert_eq!(q("SELECT * FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5"),
4252 "FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5");
4253 assert_eq!(q("SELECT * FROM orders GROUP BY region"), "FROM orders GROUP BY region");
4254 assert_eq!(q("SELECT * FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC"),
4255 "FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC");
4256 }
4257
4258 #[test]
4265 fn an_aggregate_is_one_column_named_as_sql_names_it() {
4266 assert_eq!(proj_pairs("SELECT COUNT(*) FROM orders"),
4267 vec![("count".to_string(), "count".to_string())]);
4268 assert_eq!(proj_pairs("SELECT SUM(total) FROM orders"),
4269 vec![("sum_total".to_string(), "sum".to_string())]);
4270 assert_eq!(proj_pairs("SELECT avg(total) FROM orders"),
4271 vec![("avg_total".to_string(), "avg".to_string())]);
4272 assert_eq!(proj_pairs("SELECT MIN(total) FROM orders"),
4273 vec![("min_total".to_string(), "min".to_string())]);
4274 let rows = vec![json!({"count": 4, "sum_total": 420, "value": 420})];
4276 let p = vec![Col::renamed("sum_total", "sum")];
4277 let cols = columns_for(&rows, &p);
4278 assert_eq!(names(&cols), vec!["sum"], "one column, SQL's name");
4279 assert_eq!(cell(rows[0].get(&cols[0].src)), Some("420".to_string()));
4280 }
4281
4282 #[test]
4287 fn a_bare_column_with_group_by_is_refused_not_nulled() {
4288 let e = translate("SELECT region, total FROM orders GROUP BY region").unwrap_err();
4289 assert!(e.contains("must appear in the GROUP BY clause"), "{}", e);
4290 assert!(e.contains("total"), "the message names the offending column: {}", e);
4291
4292 assert!(translate("SELECT region FROM orders GROUP BY region").is_ok());
4294 assert!(translate("SELECT region, count FROM orders GROUP BY region").is_ok());
4295 assert!(translate("SELECT SUM(total) FROM orders GROUP BY region").is_ok());
4297 assert!(translate("SELECT * FROM orders GROUP BY region").is_ok());
4299 }
4300
4301 #[test]
4302 fn count_star_becomes_nql_count() {
4303 assert_eq!(q("SELECT COUNT(*) FROM orders"), "FROM orders COUNT");
4304 assert_eq!(q("SELECT count(*) FROM orders WHERE total > 5"),
4305 "FROM orders COUNT WHERE total > 5");
4306 }
4307
4308 #[test]
4309 fn aggregates_carry_their_target_column() {
4310 assert_eq!(q("SELECT SUM(total) FROM orders"), "FROM orders SUM total");
4311 assert_eq!(q("SELECT avg(total) FROM orders WHERE region = 'eu'"),
4312 r#"FROM orders AVG total WHERE region = "eu""#);
4313 assert!(translate("SELECT SUM(*) FROM orders").is_err());
4314 }
4315
4316 #[test]
4319 fn as_of_system_time_bridges_to_nql_as_of() {
4320 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42"),
4321 "FROM orders AS OF 42");
4322 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42 WHERE total > 1"),
4323 "FROM orders AS OF 42 WHERE total > 1");
4324 let e = translate("SELECT * FROM orders AS OF SYSTEM TIME '2026-01-01'").unwrap_err();
4326 assert!(e.contains("sequence number"), "{}", e);
4327 }
4328
4329 #[test]
4338 fn a_select_list_expression_is_refused_rather_than_answered_blank() {
4339 for sql in [
4340 "SELECT total * 2 FROM orders",
4341 "SELECT total, total*2 AS doubled FROM orders",
4342 "SELECT total + 1 FROM orders",
4343 "SELECT status || 'x' FROM orders",
4344 "SELECT -total FROM orders",
4345 "SELECT lower(status) FROM orders",
4346 ] {
4347 let e = translate(sql).unwrap_err();
4348 assert!(e.contains("expressions in the select list"), "{} -> {}", sql, e);
4349 }
4350 assert_eq!(q("SELECT _id, status FROM orders"), "FROM orders");
4353 assert_eq!(q("SELECT \"status\" FROM orders"), "FROM orders");
4354 assert_eq!(q("SELECT orders.status FROM orders"), "FROM orders");
4355 assert_eq!(q("SELECT o.status FROM orders o"), "FROM orders");
4356 assert_eq!(q("SELECT total AS t FROM orders"), "FROM orders");
4357 assert!(translate("SELECT count(*) FROM orders").is_ok());
4358 assert!(translate("SELECT sum(total) FROM orders").is_ok());
4359 }
4360
4361 #[test]
4369 fn having_is_translated_to_nqls_spelling_and_refuses_an_unknown_key() {
4370 for sql in [
4372 "SELECT status, count(*) AS n FROM orders GROUP BY status HAVING count(*) > 1",
4373 "SELECT status, count(*) AS n FROM orders GROUP BY status HAVING n > 1",
4374 "SELECT status, count(*) FROM orders GROUP BY status HAVING COUNT > 1",
4375 "SELECT status, count(*) FROM orders GROUP BY status HAVING count > 1",
4376 ] {
4377 let got = q(sql);
4378 assert_eq!(got, "FROM orders GROUP BY status COUNT HAVING count > 1",
4379 "{} -> {}", sql, got);
4380 }
4381 assert_eq!(q("SELECT status, sum(total) AS s FROM orders GROUP BY status HAVING s > 100"),
4383 "FROM orders GROUP BY status SUM total HAVING sum_total > 100");
4384 assert_eq!(q("SELECT status, sum(total) FROM orders GROUP BY status HAVING sum_total > 100"),
4386 "FROM orders GROUP BY status SUM total HAVING sum_total > 100");
4387 assert_eq!(q("SELECT status, count(*) FROM orders GROUP BY status HAVING status > 'a'"),
4390 "FROM orders GROUP BY status COUNT HAVING status > \"a\"");
4391 let e = translate(
4393 "SELECT status, count(*) FROM orders GROUP BY status HAVING nosuch > 1").unwrap_err();
4394 assert!(e.contains("HAVING names") && e.contains("nosuch"), "{}", e);
4395 assert!(e.contains("zero rows"), "the message must say what it prevented: {}", e);
4396 }
4397
4398 #[test]
4399 fn handshake_queries_are_answered_so_clients_can_connect() {
4400 assert!(matches!(translate("SELECT version()"), Ok(Stmt::Canned { .. })));
4401 assert!(matches!(translate("SHOW transaction_isolation"), Ok(Stmt::Canned { .. })));
4402 assert!(matches!(translate("SELECT current_schema()"), Ok(Stmt::Canned { .. })));
4403 assert!(matches!(translate("SET extra_float_digits = 3"), Ok(Stmt::Ok(_))));
4404 assert!(matches!(translate("BEGIN"), Ok(Stmt::Ok(_))));
4405 assert!(matches!(translate(""), Ok(Stmt::Ok(_))));
4406 }
4407
4408 #[test]
4411 fn unsupported_sql_is_refused_with_a_reason() {
4412 for (sql, expect) in [
4413 ("INSERT INTO t VALUES (1)", "explicit column list"),
4414 ("CREATE TABLE t (a int)", "DDL"),
4415 ("TRUNCATE t", "append-only"),
4416 ("GRANT ALL ON t TO x", "privilege system"),
4417 ("SELECT * FROM a JOIN b ON a.x = b.x", "JOIN is not supported"),
4418 ("SELECT * FROM a UNION SELECT * FROM b", "UNION"),
4419 ("SELECT DISTINCT region FROM orders", "GROUP BY"),
4420 ("SELECT * FROM (SELECT 1) x", "subqueries in FROM"),
4421 ("SELECT * FROM a, b", "more than one collection"),
4422 ("SELECT lower(status) FROM orders", "expressions in the select list"),
4423 ("VACUUM", "only SELECT"),
4424 ] {
4425 let e = translate(sql).unwrap_err();
4426 assert!(e.contains(expect), "for {:?} expected {:?} in {:?}", sql, expect, e);
4427 }
4428 }
4429
4430 fn ins(sql: &str) -> (String, Vec<InsertRow>, Vec<Col>) {
4438 match translate(sql) {
4439 Ok(Stmt::Insert { coll, rows, returning }) => (coll, rows, returning),
4440 other => panic!("expected INSERT for {:?}, got {:?}", sql, other),
4441 }
4442 }
4443
4444 #[test]
4445 fn insert_becomes_a_put_per_row() {
4446 let (coll, rows, ret) = ins("INSERT INTO orders (_id, status, total) VALUES ('o1', 'paid', 120)");
4447 assert_eq!(coll, "orders");
4448 assert_eq!(rows.len(), 1);
4449 assert_eq!(rows[0].id.as_deref(), Some("o1"));
4450 assert_eq!(rows[0].doc.get("status"), Some(&json!("paid")));
4451 assert_eq!(rows[0].doc.get("total"), Some(&json!(120)));
4452 assert!(!rows[0].doc.contains_key("_id"));
4454 assert!(ret.is_empty());
4455 }
4456
4457 #[test]
4458 fn a_multi_row_insert_yields_one_row_each() {
4459 let (_, rows, _) = ins(
4460 "INSERT INTO t (id, n) VALUES ('a', 1), ('b', 2), ('c', 3)");
4461 assert_eq!(rows.len(), 3);
4462 assert_eq!(rows[1].id.as_deref(), Some("b"));
4463 assert_eq!(rows[2].doc.get("n"), Some(&json!(3)));
4464 }
4465
4466 #[test]
4467 fn an_insert_without_an_id_column_lets_the_server_assign_one() {
4468 let (_, rows, _) = ins("INSERT INTO t (n) VALUES (1)");
4469 assert_eq!(rows[0].id, None, "the executor mints a unique key");
4470 assert_eq!(rows[0].doc.get("n"), Some(&json!(1)));
4471 }
4472
4473 #[test]
4476 fn insert_lifts_provenance_out_of_reserved_columns() {
4477 let (_, rows, _) = ins(
4478 "INSERT INTO audit (_id, _caused_by, _valid_from, kind) \
4479 VALUES ('e1', 'abc123', '2026-01-01', 'reprice')");
4480 assert_eq!(rows[0].caused_by, vec!["abc123".to_string()]);
4481 assert_eq!(rows[0].valid_from.as_deref(), Some("2026-01-01"));
4482 assert_eq!(rows[0].doc.get("kind"), Some(&json!("reprice")));
4483 for k in ["_id", "_caused_by", "_valid_from"] {
4485 assert!(!rows[0].doc.contains_key(k), "{} leaked into the doc", k);
4486 }
4487 }
4488
4489 #[test]
4490 fn insert_values_cover_the_scalar_types() {
4491 let (_, rows, _) = ins(
4492 "INSERT INTO t (s, i, f, b, n) VALUES ('x', 42, 1.5, TRUE, NULL)");
4493 assert_eq!(rows[0].doc.get("s"), Some(&json!("x")));
4494 assert_eq!(rows[0].doc.get("i"), Some(&json!(42)));
4495 assert_eq!(rows[0].doc.get("f"), Some(&json!(1.5)));
4496 assert_eq!(rows[0].doc.get("b"), Some(&json!(true)));
4497 assert_eq!(rows[0].doc.get("n"), Some(&Value::Null));
4498 }
4499
4500 #[test]
4503 fn insert_literals_survive_quotes_and_commas() {
4504 let (_, rows, _) = ins("INSERT INTO t (a, b) VALUES ('it''s', 'x,y')");
4505 assert_eq!(rows[0].doc.get("a"), Some(&json!("it's")));
4506 assert_eq!(rows[0].doc.get("b"), Some(&json!("x,y")));
4507 }
4508
4509 #[test]
4510 fn insert_refuses_what_it_cannot_store_faithfully() {
4511 assert!(translate("INSERT INTO t (a) VALUES (1 + 1)").is_err());
4513 assert!(translate("INSERT INTO t (a) VALUES (now())").is_err());
4514 let e = translate("INSERT INTO t (a, b) VALUES (1)").unwrap_err();
4516 assert!(e.contains("values for"), "{}", e);
4517 let e2 = translate("INSERT INTO t VALUES (1)").unwrap_err();
4519 assert!(e2.contains("explicit column list"), "{}", e2);
4520 }
4521
4522 #[test]
4523 fn update_finds_rows_with_the_full_predicate_surface() {
4524 match translate("UPDATE orders SET status = 'void' WHERE total < 50 AND region IN ('eu')") {
4525 Ok(Stmt::Update { coll, set, nql, .. }) => {
4526 assert_eq!(coll, "orders");
4527 assert_eq!(set, vec![("status".to_string(), json!("void"))]);
4528 assert_eq!(nql, r#"FROM orders WHERE total < 50 AND region IN ("eu")"#);
4530 }
4531 other => panic!("expected UPDATE, got {:?}", other),
4532 }
4533 }
4534
4535 #[test]
4536 fn update_without_where_targets_the_whole_collection() {
4537 match translate("UPDATE t SET a = 1") {
4539 Ok(Stmt::Update { nql, .. }) => assert_eq!(nql, "FROM t"),
4540 other => panic!("expected UPDATE, got {:?}", other),
4541 }
4542 }
4543
4544 #[test]
4545 fn update_handles_several_assignments() {
4546 match translate("UPDATE t SET a = 1, b = 'x,y', c = NULL WHERE id = 'k'") {
4547 Ok(Stmt::Update { set, .. }) => {
4548 assert_eq!(set.len(), 3);
4549 assert_eq!(set[1], ("b".to_string(), json!("x,y")));
4550 assert_eq!(set[2], ("c".to_string(), Value::Null));
4551 }
4552 other => panic!("expected UPDATE, got {:?}", other),
4553 }
4554 assert!(translate("UPDATE t SET").is_err());
4555 assert!(translate("UPDATE t SET a").is_err());
4556 }
4557
4558 #[test]
4559 fn delete_becomes_a_predicate_over_the_collection() {
4560 match translate("DELETE FROM orders WHERE status = 'void'") {
4561 Ok(Stmt::Delete { coll, nql, .. }) => {
4562 assert_eq!(coll, "orders");
4563 assert_eq!(nql, r#"FROM orders WHERE status = "void""#);
4564 }
4565 other => panic!("expected DELETE, got {:?}", other),
4566 }
4567 match translate("DELETE FROM t") {
4568 Ok(Stmt::Delete { nql, .. }) => assert_eq!(nql, "FROM t"),
4569 other => panic!("expected DELETE, got {:?}", other),
4570 }
4571 }
4572
4573 #[test]
4574 fn returning_is_parsed_off_every_write() {
4575 let (_, _, ret) = ins("INSERT INTO t (a) VALUES (1) RETURNING a, _id");
4576 assert_eq!(ret.iter().map(|c| c.out.clone()).collect::<Vec<_>>(), vec!["a", "_id"]);
4577 let (_, _, star) = ins("INSERT INTO t (a) VALUES (1) RETURNING *");
4580 assert!(star.is_empty());
4581 assert!(wants_returning("INSERT INTO t (a) VALUES (1) RETURNING *"));
4582 assert!(!wants_returning("INSERT INTO t (a) VALUES (1)"));
4583
4584 match translate("UPDATE t SET a = 1 WHERE id = 'k' RETURNING a") {
4585 Ok(Stmt::Update { nql, returning, .. }) => {
4586 assert_eq!(returning.len(), 1);
4587 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql);
4589 }
4590 other => panic!("expected UPDATE, got {:?}", other),
4591 }
4592 match translate("DELETE FROM t WHERE id = 'k' RETURNING *") {
4593 Ok(Stmt::Delete { nql, .. }) =>
4594 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql),
4595 other => panic!("expected DELETE, got {:?}", other),
4596 }
4597 }
4598
4599 #[test]
4600 fn a_keyword_inside_a_value_is_not_a_clause() {
4601 match translate("UPDATE t SET note = 'where returning from' WHERE id = 'k'") {
4602 Ok(Stmt::Update { set, nql, .. }) => {
4603 assert_eq!(set[0].1, json!("where returning from"));
4604 assert_eq!(nql, r#"FROM t WHERE id = "k""#);
4605 }
4606 other => panic!("expected UPDATE, got {:?}", other),
4607 }
4608 }
4609
4610 #[test]
4611 fn split_top_respects_quotes_and_nesting() {
4612 assert_eq!(split_top("a, b, c", ',').len(), 3);
4613 assert_eq!(split_top("(1, 2), (3, 4)", ',').len(), 2);
4614 assert_eq!(split_top("'a,b', c", ',').len(), 2);
4615 assert_eq!(split_top("'it''s, fine', c", ',').len(), 2);
4616 }
4617
4618 #[test]
4619 fn comments_and_whitespace_do_not_confuse_the_translator() {
4620 assert_eq!(q("SELECT *\n FROM orders -- trailing note\n"), "FROM orders");
4621 assert_eq!(q("SELECT /* inline */ * FROM orders"), "FROM orders");
4622 assert_eq!(q("SELECT * FROM t WHERE note = 'from here to JOIN'"),
4624 r#"FROM t WHERE note = "from here to JOIN""#);
4625 }
4626
4627 #[test]
4628 fn find_kw_ignores_quotes_parens_and_substrings() {
4629 assert_eq!(find_kw("SELECT A FROM B", "FROM"), Some(9));
4630 assert_eq!(find_kw("SELECT 'FROM' FROM B", "FROM"), Some(14));
4631 assert_eq!(find_kw("SELECT F(x FROM y) FROM B", "FROM"), Some(19));
4632 assert_eq!(find_kw("SELECT FROMAGE", "FROM"), None);
4633 assert_eq!(find_kw("SELECT X_FROM", "FROM"), None);
4634 }
4635
4636 #[test]
4639 fn provenance_columns_sort_after_the_users_own_fields() {
4640 let rows = vec![json!({"_id":"1","_hash":"ab","status":"paid","total":9})];
4641 assert_eq!(names(&columns_for(&rows, &[])),
4642 vec!["status", "total", "_hash", "_id"]);
4643 }
4644
4645 #[test]
4646 fn an_explicit_projection_sets_the_column_order() {
4647 let rows = vec![json!({"a":1,"b":2})];
4648 let p = vec![Col::same("b"), Col::same("a")];
4649 assert_eq!(names(&columns_for(&rows, &p)), vec!["b", "a"]);
4650 }
4651
4652 #[test]
4653 fn columns_are_the_union_across_sparse_rows() {
4654 let rows = vec![json!({"a":1}), json!({"b":2})];
4656 assert_eq!(names(&columns_for(&rows, &[])), vec!["a", "b"]);
4657 }
4658
4659 #[test]
4660 fn type_oids_follow_the_first_non_null_value() {
4661 let rows = vec![json!({"i":1,"f":1.5,"b":true,"s":"x","n":null})];
4662 assert_eq!(oid_for(&rows, "i"), OID_INT8);
4663 assert_eq!(oid_for(&rows, "f"), OID_FLOAT8);
4664 assert_eq!(oid_for(&rows, "b"), OID_BOOL);
4665 assert_eq!(oid_for(&rows, "s"), OID_TEXT);
4666 assert_eq!(oid_for(&rows, "n"), OID_TEXT);
4668 assert_eq!(oid_for(&rows, "absent"), OID_TEXT);
4669 }
4670
4671 #[test]
4672 fn a_column_that_is_null_in_the_first_row_still_gets_its_type() {
4673 let rows = vec![json!({"v": null}), json!({"v": 7})];
4674 assert_eq!(oid_for(&rows, "v"), OID_INT8);
4675 }
4676
4677 #[test]
4678 fn cells_render_in_postgres_text_format() {
4679 assert_eq!(cell(Some(&json!("x"))), Some("x".to_string()));
4680 assert_eq!(cell(Some(&json!(true))), Some("t".to_string()));
4681 assert_eq!(cell(Some(&json!(false))), Some("f".to_string()));
4682 assert_eq!(cell(Some(&json!(42))), Some("42".to_string()));
4683 assert_eq!(cell(Some(&json!(null))), None);
4684 assert_eq!(cell(None), None);
4685 assert_eq!(cell(Some(&json!({"a":1}))), Some("{\"a\":1}".to_string()));
4687 }
4688
4689 #[test]
4692 fn message_framing_length_excludes_the_tag() {
4693 let mut m = Out::msg(b'Z');
4694 m.bytes(b"I");
4695 let bytes = m.finish();
4696 assert_eq!(bytes[0], b'Z');
4697 assert_eq!(i32::from_be_bytes([bytes[1], bytes[2], bytes[3], bytes[4]]), 5);
4698 assert_eq!(bytes.len(), 6);
4699 }
4700
4701 #[test]
4702 fn a_result_set_encodes_as_description_then_rows_then_complete() {
4703 let rows = vec![json!({"a": 1}), json!({"a": 2})];
4704 let out = encode_result(&rows, &[]);
4705 assert_eq!(out[0], b'T');
4706 let tags: Vec<u8> = {
4707 let mut t = vec![];
4709 let mut i = 0usize;
4710 while i < out.len() {
4711 t.push(out[i]);
4712 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
4713 i += 1 + len;
4714 }
4715 t
4716 };
4717 assert_eq!(tags, vec![b'T', b'D', b'D', b'C'],
4718 "one description, one row each, one completion");
4719 }
4720
4721 #[test]
4726 fn a_write_with_returning_emits_exactly_one_command_complete() {
4727 let rows = vec![json!({"_id": "o1", "total": 9})];
4728 let mut out = encode_rows(&rows, &[Col::same("_id")]);
4729 out.extend_from_slice(&command_complete("INSERT 0 1"));
4730 let mut tags = vec![];
4731 let mut i = 0usize;
4732 while i < out.len() {
4733 tags.push(out[i]);
4734 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
4735 i += 1 + len;
4736 }
4737 assert_eq!(tags, vec![b'T', b'D', b'C'], "one description, one row, ONE tag");
4738 assert_eq!(tags.iter().filter(|t| **t == b'C').count(), 1);
4739 assert!(!encode_rows(&rows, &[]).contains(&b'C')
4741 || encode_rows(&rows, &[]).iter().filter(|b| **b == b'C').count() > 0);
4742 let bare = encode_rows(&rows, &[Col::same("_id")]);
4743 let mut bare_tags = vec![];
4744 let mut j = 0usize;
4745 while j < bare.len() {
4746 bare_tags.push(bare[j]);
4747 let len = i32::from_be_bytes([bare[j+1], bare[j+2], bare[j+3], bare[j+4]]) as usize;
4748 j += 1 + len;
4749 }
4750 assert_eq!(bare_tags, vec![b'T', b'D'], "encode_rows never appends a tag");
4751 }
4752
4753 #[test]
4754 fn an_empty_result_still_sends_a_description() {
4755 let out = encode_result(&[], &[Col::same("a")]);
4756 assert_eq!(out[0], b'T', "clients need the shape even with no rows");
4757 }
4758
4759 #[test]
4760 fn statements_split_on_top_level_semicolons_only() {
4761 assert_eq!(split_statements("SELECT 1; SELECT 2").len(), 2);
4762 assert_eq!(split_statements("SELECT ';'").len(), 1);
4763 assert_eq!(split_statements("SELECT 1;").len(), 1);
4764 assert_eq!(split_statements(" ").len(), 0);
4765 }
4766
4767 #[test]
4768 fn an_error_names_its_sqlstate() {
4769 let e = String::from_utf8_lossy(&err_msg("0A000", "x")).to_string();
4770 assert!(e.contains("ERROR"));
4771 assert!(e.contains("0A000"));
4772 }
4773
4774 #[test]
4777 fn placeholders_are_counted_outside_string_literals() {
4778 assert_eq!(param_count("SELECT a FROM t WHERE b = $1 AND c = $2"), 2);
4779 assert_eq!(param_count("SELECT a FROM t"), 0);
4780 assert_eq!(param_count("WHERE a = $2 OR b = $2 OR c = $1"), 2);
4782 assert_eq!(param_count("SELECT a FROM t WHERE b = '$1'"), 0,
4783 "a placeholder inside a literal is data, not a parameter");
4784 assert_eq!(param_count("WHERE a = $10 AND b = $1"), 10,
4785 "two-digit indexes must not be read as $1 followed by 0");
4786 }
4787
4788 #[test]
4789 fn parameters_are_spliced_as_literals() {
4790 let out = substitute_params("WHERE a = $1 AND b = $2 AND c = $3",
4791 &[Some("'x'".into()), Some("42".into()), None]).unwrap();
4792 assert_eq!(out, "WHERE a = 'x' AND b = 42 AND c = NULL");
4793 }
4794
4795 #[test]
4796 fn substitution_leaves_string_literals_alone() {
4797 let out = substitute_params("WHERE a = '$1' AND b = $1", &[Some("9".into())]).unwrap();
4798 assert_eq!(out, "WHERE a = '$1' AND b = 9");
4799 }
4800
4801 #[test]
4802 fn too_few_parameters_is_an_error_not_a_silent_null() {
4803 let e = substitute_params("WHERE a = $2", &[Some("1".into())]).unwrap_err();
4806 assert!(e.contains("$2"), "{}", e);
4807 }
4808
4809 #[test]
4810 fn a_quote_in_a_parameter_cannot_escape_its_literal() {
4811 let lit = decode_param(Some(b"it's"), OID_TEXT, 0).unwrap().unwrap();
4812 assert_eq!(lit, "'it''s'");
4813 let out = substitute_params("WHERE a = $1", &[Some(lit)]).unwrap();
4815 assert_eq!(out, "WHERE a = 'it''s'");
4816 }
4817
4818 #[test]
4819 fn binary_parameters_decode_in_every_width_psycopg_sends() {
4820 assert_eq!(decode_param(Some(&[0x00, 0x2a]), OID_INT2, 1).unwrap().unwrap(), "42");
4823 assert_eq!(decode_param(Some(&[0, 0, 0, 7]), OID_INT4, 1).unwrap().unwrap(), "7");
4824 assert_eq!(
4825 decode_param(Some(&[0, 0, 0, 0, 0, 0, 0, 9]), OID_INT8, 1).unwrap().unwrap(), "9");
4826 assert_eq!(
4827 decode_param(Some(&0x400c_0000_0000_0000u64.to_be_bytes()), OID_FLOAT8, 1)
4828 .unwrap().unwrap(), "3.5");
4829 assert_eq!(decode_param(Some(&[1]), OID_BOOL, 1).unwrap().unwrap(), "TRUE");
4830 assert_eq!(decode_param(Some(&[0]), OID_BOOL, 1).unwrap().unwrap(), "FALSE");
4831 }
4832
4833 #[test]
4834 fn a_negative_binary_integer_keeps_its_sign() {
4835 assert_eq!(decode_param(Some(&(-5i32).to_be_bytes()), OID_INT4, 1).unwrap().unwrap(), "-5");
4836 assert_eq!(decode_param(Some(&(-5i16).to_be_bytes()), OID_INT2, 1).unwrap().unwrap(), "-5");
4837 }
4838
4839 #[test]
4840 fn a_binary_parameter_of_the_wrong_width_is_refused() {
4841 let e = decode_param(Some(&[0x2a]), OID_INT4, 1).unwrap_err();
4844 assert!(e.contains("4 bytes"), "{}", e);
4845 }
4846
4847 #[test]
4848 fn an_unspecified_text_parameter_is_treated_as_a_string() {
4849 assert_eq!(decode_param(Some(b"hello"), 0, 0).unwrap().unwrap(), "'hello'");
4852 }
4853
4854 #[test]
4855 fn a_null_parameter_decodes_to_none_in_every_format() {
4856 assert_eq!(decode_param(None, OID_TEXT, 0).unwrap(), None);
4857 assert_eq!(decode_param(None, OID_INT8, 1).unwrap(), None);
4858 }
4859
4860 #[test]
4861 fn an_unsupported_binary_type_says_so_by_name() {
4862 let e = decode_param(Some(&[0u8; 8]), 1114, 1).unwrap_err();
4863 assert!(e.contains("1114"), "{}", e);
4864 assert!(e.contains("text"), "the error should point at the way out: {}", e);
4865 }
4866
4867 #[test]
4868 fn a_text_number_that_is_not_a_number_gets_quoted() {
4869 assert_eq!(decode_param(Some(b"oops"), OID_INT8, 0).unwrap().unwrap(), "'oops'");
4872 }
4873
4874 #[test]
4875 fn a_client_declared_type_is_believed_over_inference() {
4876 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[OID_INT4, 0], None);
4879 assert_eq!(oids, vec![OID_INT4, OID_TEXT]);
4880 }
4881
4882 #[test]
4883 fn parameter_arity_is_taken_from_the_sql_when_the_client_declares_none() {
4884 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[], None);
4887 assert_eq!(oids.len(), 2);
4888 }
4889
4890 #[test]
4891 fn the_field_behind_each_placeholder_is_identified() {
4892 assert_eq!(
4893 param_fields("SELECT a FROM t WHERE qty > $1 AND status = $2", 2),
4894 vec![Some("qty".to_string()), Some("status".to_string())]);
4895 }
4896
4897 #[test]
4898 fn word_operators_do_not_hide_the_field() {
4899 assert_eq!(param_fields("SELECT a FROM t WHERE name LIKE $1", 1),
4900 vec![Some("name".to_string())]);
4901 assert_eq!(param_fields("SELECT a FROM t WHERE qty BETWEEN $1 AND $2", 2),
4902 vec![Some("qty".to_string()), Some("qty".to_string())]);
4903 assert_eq!(param_fields("SELECT a FROM t WHERE region IN ($1, $2)", 2),
4904 vec![Some("region".to_string()), Some("region".to_string())]);
4905 }
4906
4907 #[test]
4908 fn a_clause_position_types_from_the_grammar_not_from_a_column() {
4909 assert_eq!(
4913 infer_param_oids("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", &[], None),
4914 vec![OID_INT8, OID_TEXT]);
4915 assert_eq!(infer_param_oids("SELECT a FROM t AS OF $1", &[], None), vec![OID_INT8]);
4916 assert_eq!(
4919 infer_param_oids("SELECT a FROM t VALID AS OF $1", &[], None), vec![OID_TEXT]);
4920 assert_eq!(
4921 infer_param_oids("SELECT a FROM t LIMIT $1 OFFSET $2", &[], None),
4922 vec![OID_INT8, OID_INT8]);
4923 }
4924
4925 #[test]
4926 fn an_aggregate_column_types_from_what_the_aggregate_means() {
4927 assert_eq!(aggregate_oid("count", None, "t"), Some(OID_INT8));
4931 assert_eq!(aggregate_oid("avg_fee", None, "t"), Some(OID_FLOAT8),
4932 "an average is fractional even over integers");
4933 assert_eq!(aggregate_oid("max__seq", None, "t"), Some(OID_INT8));
4936 assert_eq!(aggregate_oid("total", None, "t"), None, "not an aggregate");
4937 }
4938
4939 #[test]
4940 fn the_parse_probe_uses_a_literal_that_every_clause_accepts() {
4941 let probe = probe_sql("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", 2);
4945 assert!(!probe.contains("NULL"), "{}", probe);
4946 assert!(translate(&probe).is_ok(), "the probe must parse: {}", probe);
4947 }
4948
4949 #[test]
4950 fn a_column_with_mixed_types_across_documents_is_advertised_as_text() {
4951 let rows = vec![json!({"x": 3}), json!({"x": "n/a"})];
4955 assert_eq!(oid_for(&rows, "x"), OID_TEXT);
4956 let rows = vec![json!({"x": 3}), json!({"x": 1.5})];
4958 assert_eq!(oid_for(&rows, "x"), OID_FLOAT8);
4959 let rows = vec![json!({"x": Value::Null}), json!({"x": 7})];
4961 assert_eq!(oid_for(&rows, "x"), OID_INT8);
4962 }
4963
4964 #[test]
4965 fn binary_output_encodes_each_advertised_type() {
4966 assert_eq!(cell_binary(Some(&json!(true)), OID_BOOL).unwrap().unwrap(), vec![1]);
4967 assert_eq!(cell_binary(Some(&json!(42)), OID_INT8).unwrap().unwrap(),
4968 42i64.to_be_bytes().to_vec());
4969 assert_eq!(cell_binary(Some(&json!(3.5)), OID_FLOAT8).unwrap().unwrap(),
4970 3.5f64.to_be_bytes().to_vec());
4971 assert_eq!(cell_binary(Some(&json!("hi")), OID_TEXT).unwrap().unwrap(), b"hi".to_vec());
4973 assert_eq!(cell_binary(Some(&Value::Null), OID_INT8).unwrap(), None);
4974 assert_eq!(cell(Some(&json!(true))).unwrap(), "t");
4976 }
4977
4978 #[test]
4979 fn a_value_that_does_not_fit_its_advertised_binary_type_is_refused() {
4980 let e = cell_binary(Some(&json!("nope")), OID_INT8).unwrap_err();
4985 assert!(e.contains("a string"), "{}", e);
4986 assert!(e.contains("more than one type"), "the error should explain WHY: {}", e);
4987 }
4988
4989 #[test]
4990 fn a_row_description_carries_the_requested_format_per_column() {
4991 let cols = [Col::same("a"), Col::same("b")];
4992 let m = row_description_fmt(&cols, &[OID_INT8, OID_TEXT], &[1, 0]);
4993 assert_eq!(m[0], b'T');
4994 assert_eq!(m[m.len() - 1], 0, "the last column was requested as text");
4996 }
4997
4998 #[test]
4999 fn a_qualified_column_resolves_to_its_bare_name() {
5000 assert_eq!(param_fields("SELECT a FROM t WHERE t.qty = $1", 1),
5001 vec![Some("qty".to_string())]);
5002 }
5003
5004 #[test]
5005 fn insert_placeholders_map_positionally_to_the_column_list() {
5006 assert_eq!(
5007 param_fields("INSERT INTO t (_id, qty, status) VALUES ($1, $2, $3)", 3),
5008 vec![Some("_id".to_string()), Some("qty".to_string()), Some("status".to_string())]);
5009 }
5010
5011 #[test]
5012 fn a_set_clause_placeholder_finds_its_column() {
5013 assert_eq!(param_fields("UPDATE t SET status = $1 WHERE _id = $2", 2),
5014 vec![Some("status".to_string()), Some("_id".to_string())]);
5015 }
5016
5017 #[test]
5018 fn the_target_collection_is_found_for_every_statement_kind() {
5019 assert_eq!(stmt_collection("SELECT a FROM inv WHERE b = $1"), "inv");
5020 assert_eq!(stmt_collection("UPDATE inv SET a = $1"), "inv");
5021 assert_eq!(stmt_collection("DELETE FROM inv WHERE a = $1"), "inv");
5022 assert_eq!(stmt_collection("INSERT INTO inv (a) VALUES ($1)"), "inv");
5023 assert_eq!(stmt_collection("SELECT a FROM public.inv"), "inv");
5025 assert_eq!(stmt_collection("INSERT INTO inv(a) VALUES ($1)"), "inv");
5026 }
5027
5028 #[test]
5029 fn engine_metadata_fields_type_without_touching_storage() {
5030 assert_eq!(infer_field_oid(None, "t", "_seq"), OID_INT8);
5031 assert_eq!(infer_field_oid(None, "t", "_id"), OID_TEXT);
5032 }
5033
5034 #[test]
5035 fn the_protocol_acknowledgements_are_single_empty_messages() {
5036 for (m, tag) in [
5038 (parse_complete(), b'1'), (bind_complete(), b'2'),
5039 (close_complete(), b'3'), (no_data(), b'n'), (portal_suspended(), b's'),
5040 ] {
5041 assert_eq!(m.len(), 5, "{:?}", tag as char);
5042 assert_eq!(m[0], tag);
5043 assert_eq!(i32::from_be_bytes([m[1], m[2], m[3], m[4]]), 4);
5044 }
5045 }
5046
5047 #[test]
5048 fn parameter_description_reports_its_arity_and_types() {
5049 let m = parameter_description(&[OID_TEXT, OID_INT8]);
5050 assert_eq!(m[0], b't');
5051 assert_eq!(i16::from_be_bytes([m[5], m[6]]), 2);
5052 assert_eq!(i32::from_be_bytes([m[7], m[8], m[9], m[10]]), OID_TEXT);
5053 assert_eq!(i32::from_be_bytes([m[11], m[12], m[13], m[14]]), OID_INT8);
5054 }
5055
5056 #[test]
5057 fn a_cstring_is_taken_without_its_terminator() {
5058 let body = b"one\0two\0".to_vec();
5059 let mut at = 0usize;
5060 assert_eq!(take_cstr(&body, &mut at), "one");
5061 assert_eq!(take_cstr(&body, &mut at), "two");
5062 assert_eq!(at, body.len());
5063 }
5064
5065 #[test]
5066 fn truncated_integers_are_reported_rather_than_read_past_the_end() {
5067 let body = vec![0u8, 1];
5068 let mut at = 0usize;
5069 assert!(take_i32(&body, &mut at).is_err());
5070 let mut at = 0usize;
5071 assert!(take_i16(&body, &mut at).is_ok());
5072 }
5073
5074 #[test]
5075 fn a_binary_result_format_request_is_refused_rather_than_faked() {
5076 let out = encode_rows(&[], &[Col::same("a")]);
5079 let desc_format = &out[out.len() - 2..];
5080 assert_eq!(i16::from_be_bytes([desc_format[0], desc_format[1]]), 0,
5081 "every column is advertised as text format");
5082 }
5083
5084 #[test]
5085 fn a_float_parameter_does_not_render_as_rust_infinity() {
5086 assert_eq!(fmt_float(f64::INFINITY), "'Infinity'");
5087 assert_eq!(fmt_float(f64::NEG_INFINITY), "'-Infinity'");
5088 assert_eq!(fmt_float(f64::NAN), "'NaN'");
5089 assert_eq!(fmt_float(3.0), "3", "a whole float should not gain a .0 tail");
5090 assert_eq!(fmt_float(3.5), "3.5");
5091 }
5092}