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
3581pub fn execute_sql(db: &Arc<Db>, sql: &str, read_only: bool)
3607 -> std::result::Result<Executed, String>
3608{
3609 execute_stmt(sql, "", Some(db), read_only).map_err(|wire| decode_wire_error(&wire))
3610}
3611
3612fn decode_wire_error(buf: &[u8]) -> String {
3619 let mut code: Option<String> = None;
3620 let mut msg: Option<String> = None;
3621 let body = if buf.len() > 5 { &buf[5..] } else { buf };
3623 let mut i = 0usize;
3624 while i < body.len() && body[i] != 0 {
3625 let field = body[i];
3626 i += 1;
3627 let start = i;
3628 while i < body.len() && body[i] != 0 { i += 1; }
3629 let text = String::from_utf8_lossy(&body[start..i]).into_owned();
3630 i += 1; match field {
3632 b'C' => code = Some(text),
3633 b'M' => msg = Some(text),
3634 _ => {}
3635 }
3636 }
3637 match (code, msg) {
3638 (Some(c), Some(m)) => format!("{} ({})", m, c),
3639 (None, Some(m)) => m,
3640 _ => format!(
3643 "the engine refused the statement and the error could not be decoded ({} bytes of wire response)", buf.len()
3644 ),
3645 }
3646}
3647
3648fn execute_stmt(
3649 stmt_sql: &str,
3650 db_name: &str,
3651 db: Option<&Arc<Db>>,
3652 read_only: bool,
3653) -> Result<Executed, Vec<u8>> {
3654 if let Some(inner) = strip_explain(stmt_sql) {
3663 if let Some((_, plan)) = try_catalog_select(inner, db)? {
3664 return Ok(plan_result(plan.render()));
3665 }
3666 let mut lines = vec![];
3667 match translate(inner) {
3668 Ok(_) => {
3669 lines.push(
3670 "NQL path — this statement is translated to NQL and \
3671 executed by the storage engine, not by the SQL evaluator."
3672 .to_string(),
3673 );
3674 lines.push(
3675 "No plan is reported, because the SQL evaluator is not \
3676 what runs it. Reporting one would describe a pipeline \
3677 that never executed."
3678 .to_string(),
3679 );
3680 lines.push(
3681 "The SQL evaluator (joins, CASE, scalar functions, a \
3682 hash-join planner) currently serves catalogue queries."
3683 .to_string(),
3684 );
3685 }
3686 Err(why) => lines.push(format!("cannot be executed: {why}")),
3687 }
3688 return Ok(plan_result(lines));
3689 }
3690
3691 if let Some((done, _plan)) = try_catalog_select(stmt_sql, db)? {
3692 return Ok(done);
3693 }
3694
3695 let stmt = translate(stmt_sql).map_err(|why| err_msg("0A000", &why))?;
3696
3697 macro_rules! need_db {
3700 () => {
3701 match db {
3702 Some(db) => db,
3703 None => return Err(no_db(db_name)),
3704 }
3705 };
3706 }
3707 macro_rules! need_write {
3708 () => {
3709 if read_only {
3710 return Err(err_msg("25006", READ_ONLY_MSG));
3711 }
3712 };
3713 }
3714
3715 match stmt {
3716 Stmt::Ok(tag) => Ok(Executed::nothing(if tag.is_empty() { "SELECT 0" } else { tag })),
3717
3718 Stmt::Canned { cols, row } => {
3719 let mut obj = serde_json::Map::new();
3722 for (c, v) in cols.iter().zip(row.iter()) {
3723 obj.insert(c.clone(), Value::String(v.clone()));
3724 }
3725 Ok(Executed {
3726 rows: vec![Value::Object(obj)],
3727 project: cols.iter().map(|c| Col::same(c)).collect(),
3728 has_rows: true,
3729 tag: "SELECT".into(),
3730 tag_counts_rows: true,
3731 })
3732 }
3733
3734 Stmt::Query { nql, project } => {
3735 if let Some(coll) = catalog_target(&nql) {
3745 let rows = crate::pgcatalog::rows(&coll, db)
3746 .expect("catalog_target only returns names pgcatalog serves");
3747 let rows = crate::nql::query_rows(rows, &nql)
3748 .map_err(|e| err_msg("42601", &e.to_string()))?;
3749 return Ok(Executed {
3750 rows, project, has_rows: true,
3751 tag: "SELECT".into(), tag_counts_rows: true,
3752 });
3753 }
3754 let db = need_db!();
3755 let (rows, _) = crate::nql::query(db, &nql).map_err(|e| {
3756 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3757 })?;
3758 Ok(Executed { rows, project, has_rows: true, tag: "SELECT".into(), tag_counts_rows: true })
3759 }
3760
3761 Stmt::Insert { coll, rows, returning } => {
3762 let db = need_db!();
3763 need_write!();
3764 let mut written: Vec<Value> = vec![];
3765 for (i, r) in rows.iter().enumerate() {
3766 let id = match &r.id {
3770 Some(id) => id.clone(),
3771 None => format!("{}-{}", next_row_id(), i),
3772 };
3773 let node = db
3774 .put(&coll, &id, Value::Object(r.doc.clone()),
3775 r.caused_by.clone(), r.valid_from.clone(), r.valid_to.clone())
3776 .map_err(|e| err_msg("XX000", &format!("INSERT failed: {}", e)))?;
3777 written.push(crate::nql::node_to_json(&node));
3778 }
3779 let n = written.len();
3780 let has_rows = wants_returning(stmt_sql);
3781 Ok(Executed {
3782 rows: if has_rows { written } else { vec![] },
3783 project: returning,
3784 has_rows,
3785 tag: format!("INSERT 0 {}", n),
3787 tag_counts_rows: false,
3788 })
3789 }
3790
3791 Stmt::Update { coll, set, nql, returning } => {
3792 let db = need_db!();
3793 need_write!();
3794 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
3797 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3798 })?;
3799 let mut written: Vec<Value> = vec![];
3800 for row in &matched {
3801 let id = match row.get("_id").and_then(|v| v.as_str()) {
3802 Some(id) => id.to_string(),
3803 None => continue,
3804 };
3805 let mut doc = match db.get(&coll, &id) {
3809 Some(n) => match n.data {
3810 Value::Object(m) => m,
3811 _ => serde_json::Map::new(),
3812 },
3813 None => continue,
3814 };
3815 for (k, v) in &set {
3816 doc.insert(k.clone(), v.clone());
3817 }
3818 let node = db
3821 .put(&coll, &id, Value::Object(doc), vec![], None, None)
3822 .map_err(|e| err_msg("XX000", &format!("UPDATE failed: {}", e)))?;
3823 written.push(crate::nql::node_to_json(&node));
3824 }
3825 let n = written.len();
3826 let has_rows = wants_returning(stmt_sql);
3827 Ok(Executed {
3828 rows: if has_rows { written } else { vec![] },
3829 project: returning,
3830 has_rows,
3831 tag: format!("UPDATE {}", n),
3832 tag_counts_rows: false,
3833 })
3834 }
3835
3836 Stmt::Delete { coll, nql, returning } => {
3837 let db = need_db!();
3838 need_write!();
3839 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
3840 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
3841 })?;
3842 let returned = matched.clone();
3845 let mut n = 0usize;
3846 for row in &matched {
3847 if let Some(id) = row.get("_id").and_then(|v| v.as_str()) {
3848 match db.delete(&coll, id) {
3849 Ok(true) => n += 1,
3850 Ok(false) => {}
3851 Err(e) => return Err(err_msg("XX000", &format!("DELETE failed: {}", e))),
3852 }
3853 }
3854 }
3855 let has_rows = wants_returning(stmt_sql);
3856 Ok(Executed {
3857 rows: if has_rows { returned } else { vec![] },
3858 project: returning,
3859 has_rows,
3860 tag: format!("DELETE {}", n),
3861 tag_counts_rows: false,
3862 })
3863 }
3864 }
3865}
3866
3867fn run_simple_query(sql: &str, db_name: &str, db: Option<&Arc<Db>>, read_only: bool) -> Vec<u8> {
3869 let mut out = vec![];
3870 let statements = split_statements(sql);
3871 if statements.is_empty() {
3872 return Out::msg(b'I').finish();
3874 }
3875 for stmt_sql in statements {
3876 match execute_stmt(&stmt_sql, db_name, db, read_only) {
3877 Err(encoded) => {
3879 out.extend_from_slice(&encoded);
3880 return out;
3881 }
3882 Ok(ex) => {
3883 if ex.has_rows {
3884 out.extend_from_slice(&encode_rows(&ex.rows, &ex.project));
3885 }
3886 out.extend_from_slice(&command_complete(&ex.tag_for(ex.rows.len())));
3887 }
3888 }
3889 }
3890 out
3891}
3892
3893fn split_statements(sql: &str) -> Vec<String> {
3895 let mut out = vec![];
3896 let mut cur = String::new();
3897 let mut in_s = false;
3898 for c in sql.chars() {
3899 match c {
3900 '\'' => { in_s = !in_s; cur.push(c); }
3901 ';' if !in_s => {
3902 if !cur.trim().is_empty() { out.push(cur.clone()); }
3903 cur.clear();
3904 }
3905 _ => cur.push(c),
3906 }
3907 }
3908 if !cur.trim().is_empty() {
3909 out.push(cur);
3910 }
3911 out
3912}
3913
3914pub async fn run(host: &str, port: u16, resolver: Arc<dyn DbResolver>) -> anyhow::Result<()> {
3916 let read_only = std::env::var("NEDBD_PG_READ_ONLY")
3920 .map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
3921 .unwrap_or(false);
3922 let listener = TcpListener::bind((host, port)).await?;
3923 println!(" pgwire postgres endpoint on {}:{} — psql / DBeaver / psycopg ({})",
3924 host, port,
3925 if read_only { "SELECT only — read-only mode" } else { "SELECT + INSERT/UPDATE/DELETE" });
3926 loop {
3927 let (sock, _peer) = match listener.accept().await {
3928 Ok(v) => v,
3929 Err(e) => {
3930 eprintln!(" [pgwire] accept failed: {}", e);
3931 continue;
3932 }
3933 };
3934 let r = Arc::clone(&resolver);
3935 tokio::spawn(async move {
3936 let _ = sock.set_nodelay(true);
3937 if let Err(e) = handle(sock, r, read_only).await {
3938 if e.kind() != std::io::ErrorKind::UnexpectedEof
3940 && e.kind() != std::io::ErrorKind::ConnectionReset
3941 {
3942 eprintln!(" [pgwire] connection error: {}", e);
3943 }
3944 }
3945 });
3946 }
3947}
3948
3949#[cfg(test)]
3952mod explain_tests {
3953 use super::*;
3954
3955 #[test]
3956 fn a_bare_explain_is_stripped() {
3957 assert_eq!(strip_explain("EXPLAIN SELECT 1"), Some("SELECT 1"));
3958 assert_eq!(strip_explain("explain select 1"), Some("select 1"));
3959 assert_eq!(strip_explain(" EXPLAIN SELECT 1 ; "), Some("SELECT 1"));
3960 }
3961
3962 #[test]
3963 fn analyze_and_verbose_are_accepted_and_ignored() {
3964 assert_eq!(strip_explain("EXPLAIN ANALYZE SELECT 1"), Some("SELECT 1"));
3968 assert_eq!(strip_explain("EXPLAIN ANALYSE SELECT 1"), Some("SELECT 1"));
3969 assert_eq!(strip_explain("EXPLAIN VERBOSE SELECT 1"), Some("SELECT 1"));
3970 assert_eq!(strip_explain("EXPLAIN ANALYZE VERBOSE SELECT 1"), Some("SELECT 1"));
3971 assert_eq!(strip_explain("explain analyze verbose select 1"), Some("select 1"));
3972 }
3973
3974 #[test]
3975 fn a_word_merely_starting_with_explain_is_not_a_keyword() {
3976 assert_eq!(strip_explain("EXPLAINED SELECT 1"), None);
3977 assert_eq!(strip_explain("SELECT 1"), None);
3978 assert_eq!(strip_explain("SELECT explain FROM t"), None);
3979 }
3980
3981 #[test]
3982 fn a_column_named_analyze_is_not_eaten() {
3983 assert_eq!(strip_explain("EXPLAIN analyzed_view"), Some("analyzed_view"));
3986 }
3987
3988 #[test]
3989 fn the_plan_result_has_postgres_shape() {
3990 let e = plan_result(vec!["Seq Scan on t".into(), "note".into()]);
3991 assert_eq!(e.project.len(), 1);
3992 assert_eq!(e.project[0].out, "QUERY PLAN");
3993 assert_eq!(e.rows.len(), 2);
3994 assert_eq!(e.rows[0]["QUERY PLAN"], "Seq Scan on t");
3995 assert_eq!(e.tag, "EXPLAIN");
3996 assert!(!e.tag_counts_rows);
3998 }
3999}
4000
4001#[cfg(test)]
4002mod tests {
4003 use super::*;
4004 use serde_json::json;
4005
4006 fn q(sql: &str) -> String {
4007 match translate(sql) {
4008 Ok(Stmt::Query { nql, .. }) => nql,
4009 other => panic!("expected a query for {:?}, got {:?}", sql, other),
4010 }
4011 }
4012 fn proj(sql: &str) -> Vec<String> {
4014 match translate(sql) {
4015 Ok(Stmt::Query { project, .. }) => project.iter().map(|c| c.out.clone()).collect(),
4016 other => panic!("expected a query for {:?}, got {:?}", sql, other),
4017 }
4018 }
4019 fn proj_pairs(sql: &str) -> Vec<(String, String)> {
4021 match translate(sql) {
4022 Ok(Stmt::Query { project, .. }) =>
4023 project.iter().map(|c| (c.src.clone(), c.out.clone())).collect(),
4024 other => panic!("expected a query for {:?}, got {:?}", sql, other),
4025 }
4026 }
4027 fn names(cols: &[Col]) -> Vec<String> { cols.iter().map(|c| c.out.clone()).collect() }
4028
4029 fn cols_of(sql: &str) -> Vec<Col> {
4033 match translate(sql).unwrap() {
4034 Stmt::Query { project, .. } => project,
4035 other => panic!("{:?}", other),
4036 }
4037 }
4038
4039 #[test]
4040 fn select_star_becomes_bare_from() {
4041 assert_eq!(q("SELECT * FROM orders"), "FROM orders");
4042 assert_eq!(q("select * from orders;"), "FROM orders");
4043 assert_eq!(proj("SELECT * FROM orders"), Vec::<String>::new());
4044 }
4045
4046 #[test]
4047 fn a_column_list_becomes_a_projection_not_a_clause() {
4048 assert_eq!(q("SELECT status, total FROM orders"), "FROM orders");
4051 assert_eq!(proj("SELECT status, total FROM orders"), vec!["status", "total"]);
4052 }
4053
4054 #[test]
4055 fn a_qualifier_reduces_to_the_field_while_an_ALIAS_is_the_name_the_client_sees() {
4056 let cols = cols_of("SELECT o.status AS s, o.total total, o.region FROM orders o");
4064 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
4065 vec!["status", "total", "region"]);
4066 assert_eq!(cols.iter().map(|c| c.out.clone()).collect::<Vec<_>>(),
4067 vec!["s", "total", "region"]);
4068 assert_eq!(q("SELECT * FROM public.orders"), "FROM orders");
4069 assert_eq!(q("SELECT * FROM \"orders\""), "FROM orders");
4070 }
4071
4072 #[test]
4073 fn a_select_list_may_MIX_columns_with_an_aggregate() {
4074 assert_eq!(q("SELECT status, count(*) AS count_1 FROM orders GROUP BY status"),
4082 "FROM orders GROUP BY status COUNT");
4083 assert_eq!(q("SELECT status, count(*) FROM orders WHERE total > 1 GROUP BY status ORDER BY status LIMIT 5"),
4086 "FROM orders WHERE total > 1 GROUP BY status COUNT ORDER BY status LIMIT 5");
4087 assert_eq!(q("SELECT count(*) FROM orders"), "FROM orders COUNT");
4089 assert_eq!(q("SELECT sum(total) FROM orders"), "FROM orders SUM total");
4090 let e = translate("SELECT status, count(*) FROM orders GROUP BY status, region").unwrap_err();
4094 assert!(e.contains("GROUP BY takes one key"), "{}", e);
4095 let cols = cols_of("SELECT status, count(*) AS count_1 FROM orders GROUP BY status");
4096 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
4097 vec!["status", "count"]);
4098 assert_eq!(cols.iter().map(|c| c.out.clone()).collect::<Vec<_>>(),
4099 vec!["status", "count_1"]);
4100
4101 let cols = cols_of("SELECT status, count(*), sum(total) FROM orders GROUP BY status");
4104 assert_eq!(cols.iter().map(|c| c.src.clone()).collect::<Vec<_>>(),
4105 vec!["status", "count", "sum_total"]);
4106 assert_eq!(q("SELECT status, count(*), sum(total) FROM orders GROUP BY status"),
4107 "FROM orders GROUP BY status SUM total");
4108
4109 assert_eq!(q("SELECT o.status, sum(o.total) FROM orders o GROUP BY o.status"),
4111 "FROM orders GROUP BY status SUM total");
4112
4113 let e = translate("SELECT status, sum(total), avg(total) FROM orders GROUP BY status")
4116 .unwrap_err();
4117 assert!(e.contains("only one of SUM/AVG/MIN/MAX"), "{}", e);
4118
4119 let e = translate("SELECT status, total, count(*) FROM orders GROUP BY status")
4121 .unwrap_err();
4122 assert!(e.contains("must appear in the GROUP BY clause"), "{}", e);
4123 }
4124
4125 #[test]
4126 fn ORDER_BY_an_ordinal_resolves_to_that_select_list_column() {
4127 assert_eq!(q("SELECT status, total FROM orders ORDER BY 1"),
4132 "FROM orders ORDER BY status");
4133 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 DESC"),
4134 "FROM orders ORDER BY total DESC");
4135 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 DESC, 1"),
4137 "FROM orders ORDER BY total DESC, status");
4138 assert_eq!(q("SELECT status, total FROM orders ORDER BY 1, total DESC"),
4139 "FROM orders ORDER BY status, total DESC");
4140 assert_eq!(q("SELECT status, count(*) AS n FROM orders GROUP BY status ORDER BY 1"),
4144 "FROM orders GROUP BY status COUNT ORDER BY status");
4145 assert_eq!(q("SELECT status, count(*) AS n FROM orders GROUP BY status ORDER BY 2 DESC"),
4147 "FROM orders GROUP BY status COUNT ORDER BY count DESC");
4148 assert_eq!(q("SELECT status, total FROM orders ORDER BY 2 LIMIT 1"),
4151 "FROM orders ORDER BY total LIMIT 1");
4152 assert_eq!(q("SELECT status FROM orders WHERE total > 1 ORDER BY 1"),
4154 "FROM orders WHERE total > 1 ORDER BY status");
4155
4156 let e = translate("SELECT status FROM orders ORDER BY 4").unwrap_err();
4160 assert!(e.contains("out of range") && e.contains("1 column"), "{}", e);
4161 let e = translate("SELECT * FROM orders ORDER BY 1").unwrap_err();
4162 assert!(e.contains("no list to index"), "{}", e);
4163 }
4164
4165 #[test]
4166 fn count_of_a_subquery_flattens_only_when_the_two_counts_MUST_agree() {
4167 assert_eq!(
4171 q("SELECT count(*) AS count_1 FROM (SELECT orders._id AS a, orders.status AS b \
4172 FROM orders WHERE orders.status = 'paid') AS anon_1"),
4173 r#"FROM orders COUNT WHERE status = "paid""#);
4176 assert_eq!(q("SELECT count(*) FROM (SELECT orders._id FROM orders) AS anon_1"),
4178 "FROM orders COUNT");
4179 assert_eq!(q("SELECT count(*) FROM (SELECT _id FROM orders ORDER BY total DESC) AS a"),
4181 "FROM orders COUNT");
4182 let cols = cols_of("SELECT count(*) AS count_1 FROM (SELECT _id FROM orders) AS a");
4184 assert_eq!(cols[0].src, "count");
4185 assert_eq!(cols[0].out, "count_1");
4186
4187 for sql in [
4190 "SELECT count(*) FROM (SELECT _id FROM orders LIMIT 1) AS a",
4192 "SELECT count(*) FROM (SELECT _id FROM orders OFFSET 1) AS a",
4193 "SELECT count(*) FROM (SELECT status FROM orders GROUP BY status) AS a",
4195 "SELECT count(*) FROM (SELECT count(*) FROM orders) AS a",
4197 "SELECT count(*) FROM (SELECT sum(total) FROM orders) AS a",
4198 "SELECT count(*), status FROM (SELECT status FROM orders) AS a",
4200 "SELECT status FROM (SELECT status FROM orders) AS a",
4201 "SELECT count(*) FROM (SELECT x FROM (SELECT _id AS x FROM orders) AS b) AS a",
4203 ] {
4204 let e = translate(sql).unwrap_err();
4205 assert!(e.contains("subqueries in FROM"), "{} -> {}", sql, e);
4206 }
4207
4208 for (sql, needle) in [
4213 ("SELECT count(*) FROM (SELECT DISTINCT status FROM orders) AS a", "DISTINCT"),
4214 ("SELECT count(*) FROM (SELECT a FROM t UNION SELECT b FROM u) AS x", "UNION"),
4215 ] {
4216 let e = translate(sql).unwrap_err();
4217 assert!(e.contains(needle), "{} -> {}", sql, e);
4218 }
4219 }
4220
4221 #[test]
4222 fn a_QUALIFIED_column_in_WHERE_finds_its_field_instead_of_ZERO_ROWS() {
4223 assert_eq!(q("SELECT _id FROM orders WHERE orders.status = 'paid'"),
4231 r#"FROM orders WHERE status = "paid""#);
4232 assert_eq!(q("SELECT _id FROM orders WHERE orders.total > 50"),
4233 "FROM orders WHERE total > 50");
4234 assert_eq!(q("SELECT _id FROM orders ORDER BY orders.total DESC LIMIT 2"),
4236 "FROM orders ORDER BY total DESC LIMIT 2");
4237 assert_eq!(q("SELECT status, count(*) FROM orders GROUP BY orders.status"),
4238 "FROM orders GROUP BY status COUNT");
4239
4240 assert_eq!(q("SELECT o.status FROM orders o WHERE o.status = 'paid'"),
4244 r#"FROM orders WHERE status = "paid""#);
4245 assert_eq!(q("SELECT o.status FROM orders AS o WHERE o.total > 1"),
4246 "FROM orders WHERE total > 1");
4247
4248 let e = translate("SELECT _id FROM orders WHERE nosuch.status = 'paid'").unwrap_err();
4253 assert!(e.contains("no table or alias named \"nosuch\""), "{}", e);
4254 let e = translate("SELECT _id FROM orders o WHERE p.status = 'paid'").unwrap_err();
4255 assert!(e.contains("aliased \"o\""), "the message names the alias in scope: {}", e);
4256
4257 assert_eq!(q("SELECT _id FROM orders WHERE status = 'pa.id'"),
4259 r#"FROM orders WHERE status = "pa.id""#);
4260 assert_eq!(q("SELECT _id FROM orders WHERE total > 1.5"),
4262 "FROM orders WHERE total > 1.5");
4263
4264 match translate("UPDATE orders o SET status = 'x' WHERE o.total > 5").unwrap() {
4266 Stmt::Update { coll, nql, .. } => {
4267 assert_eq!(coll, "orders", "the alias is not part of the collection name");
4268 assert_eq!(nql, "FROM orders WHERE total > 5");
4269 }
4270 other => panic!("{:?}", other),
4271 }
4272 match translate("DELETE FROM orders o WHERE o.status = 'paid'").unwrap() {
4273 Stmt::Delete { coll, nql, .. } => {
4274 assert_eq!(coll, "orders");
4275 assert_eq!(nql, r#"FROM orders WHERE status = "paid""#);
4276 }
4277 other => panic!("{:?}", other),
4278 }
4279
4280 assert_eq!(q("SELECT _id FROM orders AS OF SYSTEM TIME 3 WHERE orders.total > 1"),
4282 "FROM orders AS OF 3 WHERE total > 1");
4283 }
4284
4285 #[test]
4286 fn where_clauses_pass_through_with_sql_literals_rewritten() {
4287 assert_eq!(q("SELECT * FROM orders WHERE status = 'paid'"),
4288 r#"FROM orders WHERE status = "paid""#);
4289 assert_eq!(q("SELECT * FROM orders WHERE status <> 'paid'"),
4290 r#"FROM orders WHERE status != "paid""#);
4291 assert_eq!(q("SELECT * FROM orders WHERE status IN ('paid','open')"),
4292 r#"FROM orders WHERE status IN ("paid","open")"#);
4293 }
4294
4295 #[test]
4298 fn a_doubled_sql_quote_is_one_literal_character() {
4299 assert_eq!(q("SELECT * FROM t WHERE name = 'it''s'"),
4300 r#"FROM t WHERE name = "it's""#);
4301 }
4302
4303 #[test]
4306 fn a_double_quote_inside_a_sql_literal_is_escaped_for_nql() {
4307 assert_eq!(q(r#"SELECT * FROM t WHERE name = 'say "hi"'"#),
4308 r#"FROM t WHERE name = "say \"hi\"""#);
4309 }
4310
4311 #[test]
4312 fn the_shared_clauses_are_handed_to_nql_unchanged() {
4313 assert_eq!(q("SELECT * FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5"),
4314 "FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5");
4315 assert_eq!(q("SELECT * FROM orders GROUP BY region"), "FROM orders GROUP BY region");
4316 assert_eq!(q("SELECT * FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC"),
4317 "FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC");
4318 }
4319
4320 #[test]
4327 fn an_aggregate_is_one_column_named_as_sql_names_it() {
4328 assert_eq!(proj_pairs("SELECT COUNT(*) FROM orders"),
4329 vec![("count".to_string(), "count".to_string())]);
4330 assert_eq!(proj_pairs("SELECT SUM(total) FROM orders"),
4331 vec![("sum_total".to_string(), "sum".to_string())]);
4332 assert_eq!(proj_pairs("SELECT avg(total) FROM orders"),
4333 vec![("avg_total".to_string(), "avg".to_string())]);
4334 assert_eq!(proj_pairs("SELECT MIN(total) FROM orders"),
4335 vec![("min_total".to_string(), "min".to_string())]);
4336 let rows = vec![json!({"count": 4, "sum_total": 420, "value": 420})];
4338 let p = vec![Col::renamed("sum_total", "sum")];
4339 let cols = columns_for(&rows, &p);
4340 assert_eq!(names(&cols), vec!["sum"], "one column, SQL's name");
4341 assert_eq!(cell(rows[0].get(&cols[0].src)), Some("420".to_string()));
4342 }
4343
4344 #[test]
4349 fn a_bare_column_with_group_by_is_refused_not_nulled() {
4350 let e = translate("SELECT region, total FROM orders GROUP BY region").unwrap_err();
4351 assert!(e.contains("must appear in the GROUP BY clause"), "{}", e);
4352 assert!(e.contains("total"), "the message names the offending column: {}", e);
4353
4354 assert!(translate("SELECT region FROM orders GROUP BY region").is_ok());
4356 assert!(translate("SELECT region, count FROM orders GROUP BY region").is_ok());
4357 assert!(translate("SELECT SUM(total) FROM orders GROUP BY region").is_ok());
4359 assert!(translate("SELECT * FROM orders GROUP BY region").is_ok());
4361 }
4362
4363 #[test]
4364 fn count_star_becomes_nql_count() {
4365 assert_eq!(q("SELECT COUNT(*) FROM orders"), "FROM orders COUNT");
4366 assert_eq!(q("SELECT count(*) FROM orders WHERE total > 5"),
4367 "FROM orders COUNT WHERE total > 5");
4368 }
4369
4370 #[test]
4371 fn aggregates_carry_their_target_column() {
4372 assert_eq!(q("SELECT SUM(total) FROM orders"), "FROM orders SUM total");
4373 assert_eq!(q("SELECT avg(total) FROM orders WHERE region = 'eu'"),
4374 r#"FROM orders AVG total WHERE region = "eu""#);
4375 assert!(translate("SELECT SUM(*) FROM orders").is_err());
4376 }
4377
4378 #[test]
4381 fn as_of_system_time_bridges_to_nql_as_of() {
4382 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42"),
4383 "FROM orders AS OF 42");
4384 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42 WHERE total > 1"),
4385 "FROM orders AS OF 42 WHERE total > 1");
4386 let e = translate("SELECT * FROM orders AS OF SYSTEM TIME '2026-01-01'").unwrap_err();
4388 assert!(e.contains("sequence number"), "{}", e);
4389 }
4390
4391 #[test]
4400 fn a_select_list_expression_is_refused_rather_than_answered_blank() {
4401 for sql in [
4402 "SELECT total * 2 FROM orders",
4403 "SELECT total, total*2 AS doubled FROM orders",
4404 "SELECT total + 1 FROM orders",
4405 "SELECT status || 'x' FROM orders",
4406 "SELECT -total FROM orders",
4407 "SELECT lower(status) FROM orders",
4408 ] {
4409 let e = translate(sql).unwrap_err();
4410 assert!(e.contains("expressions in the select list"), "{} -> {}", sql, e);
4411 }
4412 assert_eq!(q("SELECT _id, status FROM orders"), "FROM orders");
4415 assert_eq!(q("SELECT \"status\" FROM orders"), "FROM orders");
4416 assert_eq!(q("SELECT orders.status FROM orders"), "FROM orders");
4417 assert_eq!(q("SELECT o.status FROM orders o"), "FROM orders");
4418 assert_eq!(q("SELECT total AS t FROM orders"), "FROM orders");
4419 assert!(translate("SELECT count(*) FROM orders").is_ok());
4420 assert!(translate("SELECT sum(total) FROM orders").is_ok());
4421 }
4422
4423 #[test]
4431 fn having_is_translated_to_nqls_spelling_and_refuses_an_unknown_key() {
4432 for sql in [
4434 "SELECT status, count(*) AS n FROM orders GROUP BY status HAVING count(*) > 1",
4435 "SELECT status, count(*) AS n FROM orders GROUP BY status HAVING n > 1",
4436 "SELECT status, count(*) FROM orders GROUP BY status HAVING COUNT > 1",
4437 "SELECT status, count(*) FROM orders GROUP BY status HAVING count > 1",
4438 ] {
4439 let got = q(sql);
4440 assert_eq!(got, "FROM orders GROUP BY status COUNT HAVING count > 1",
4441 "{} -> {}", sql, got);
4442 }
4443 assert_eq!(q("SELECT status, sum(total) AS s FROM orders GROUP BY status HAVING s > 100"),
4445 "FROM orders GROUP BY status SUM total HAVING sum_total > 100");
4446 assert_eq!(q("SELECT status, sum(total) FROM orders GROUP BY status HAVING sum_total > 100"),
4448 "FROM orders GROUP BY status SUM total HAVING sum_total > 100");
4449 assert_eq!(q("SELECT status, count(*) FROM orders GROUP BY status HAVING status > 'a'"),
4452 "FROM orders GROUP BY status COUNT HAVING status > \"a\"");
4453 let e = translate(
4455 "SELECT status, count(*) FROM orders GROUP BY status HAVING nosuch > 1").unwrap_err();
4456 assert!(e.contains("HAVING names") && e.contains("nosuch"), "{}", e);
4457 assert!(e.contains("zero rows"), "the message must say what it prevented: {}", e);
4458 }
4459
4460 #[test]
4461 fn handshake_queries_are_answered_so_clients_can_connect() {
4462 assert!(matches!(translate("SELECT version()"), Ok(Stmt::Canned { .. })));
4463 assert!(matches!(translate("SHOW transaction_isolation"), Ok(Stmt::Canned { .. })));
4464 assert!(matches!(translate("SELECT current_schema()"), Ok(Stmt::Canned { .. })));
4465 assert!(matches!(translate("SET extra_float_digits = 3"), Ok(Stmt::Ok(_))));
4466 assert!(matches!(translate("BEGIN"), Ok(Stmt::Ok(_))));
4467 assert!(matches!(translate(""), Ok(Stmt::Ok(_))));
4468 }
4469
4470 #[test]
4473 fn unsupported_sql_is_refused_with_a_reason() {
4474 for (sql, expect) in [
4475 ("INSERT INTO t VALUES (1)", "explicit column list"),
4476 ("CREATE TABLE t (a int)", "DDL"),
4477 ("TRUNCATE t", "append-only"),
4478 ("GRANT ALL ON t TO x", "privilege system"),
4479 ("SELECT * FROM a JOIN b ON a.x = b.x", "JOIN is not supported"),
4480 ("SELECT * FROM a UNION SELECT * FROM b", "UNION"),
4481 ("SELECT DISTINCT region FROM orders", "GROUP BY"),
4482 ("SELECT * FROM (SELECT 1) x", "subqueries in FROM"),
4483 ("SELECT * FROM a, b", "more than one collection"),
4484 ("SELECT lower(status) FROM orders", "expressions in the select list"),
4485 ("VACUUM", "only SELECT"),
4486 ] {
4487 let e = translate(sql).unwrap_err();
4488 assert!(e.contains(expect), "for {:?} expected {:?} in {:?}", sql, expect, e);
4489 }
4490 }
4491
4492 fn ins(sql: &str) -> (String, Vec<InsertRow>, Vec<Col>) {
4500 match translate(sql) {
4501 Ok(Stmt::Insert { coll, rows, returning }) => (coll, rows, returning),
4502 other => panic!("expected INSERT for {:?}, got {:?}", sql, other),
4503 }
4504 }
4505
4506 #[test]
4507 fn insert_becomes_a_put_per_row() {
4508 let (coll, rows, ret) = ins("INSERT INTO orders (_id, status, total) VALUES ('o1', 'paid', 120)");
4509 assert_eq!(coll, "orders");
4510 assert_eq!(rows.len(), 1);
4511 assert_eq!(rows[0].id.as_deref(), Some("o1"));
4512 assert_eq!(rows[0].doc.get("status"), Some(&json!("paid")));
4513 assert_eq!(rows[0].doc.get("total"), Some(&json!(120)));
4514 assert!(!rows[0].doc.contains_key("_id"));
4516 assert!(ret.is_empty());
4517 }
4518
4519 #[test]
4520 fn a_multi_row_insert_yields_one_row_each() {
4521 let (_, rows, _) = ins(
4522 "INSERT INTO t (id, n) VALUES ('a', 1), ('b', 2), ('c', 3)");
4523 assert_eq!(rows.len(), 3);
4524 assert_eq!(rows[1].id.as_deref(), Some("b"));
4525 assert_eq!(rows[2].doc.get("n"), Some(&json!(3)));
4526 }
4527
4528 #[test]
4529 fn an_insert_without_an_id_column_lets_the_server_assign_one() {
4530 let (_, rows, _) = ins("INSERT INTO t (n) VALUES (1)");
4531 assert_eq!(rows[0].id, None, "the executor mints a unique key");
4532 assert_eq!(rows[0].doc.get("n"), Some(&json!(1)));
4533 }
4534
4535 #[test]
4538 fn insert_lifts_provenance_out_of_reserved_columns() {
4539 let (_, rows, _) = ins(
4540 "INSERT INTO audit (_id, _caused_by, _valid_from, kind) \
4541 VALUES ('e1', 'abc123', '2026-01-01', 'reprice')");
4542 assert_eq!(rows[0].caused_by, vec!["abc123".to_string()]);
4543 assert_eq!(rows[0].valid_from.as_deref(), Some("2026-01-01"));
4544 assert_eq!(rows[0].doc.get("kind"), Some(&json!("reprice")));
4545 for k in ["_id", "_caused_by", "_valid_from"] {
4547 assert!(!rows[0].doc.contains_key(k), "{} leaked into the doc", k);
4548 }
4549 }
4550
4551 #[test]
4552 fn insert_values_cover_the_scalar_types() {
4553 let (_, rows, _) = ins(
4554 "INSERT INTO t (s, i, f, b, n) VALUES ('x', 42, 1.5, TRUE, NULL)");
4555 assert_eq!(rows[0].doc.get("s"), Some(&json!("x")));
4556 assert_eq!(rows[0].doc.get("i"), Some(&json!(42)));
4557 assert_eq!(rows[0].doc.get("f"), Some(&json!(1.5)));
4558 assert_eq!(rows[0].doc.get("b"), Some(&json!(true)));
4559 assert_eq!(rows[0].doc.get("n"), Some(&Value::Null));
4560 }
4561
4562 #[test]
4565 fn insert_literals_survive_quotes_and_commas() {
4566 let (_, rows, _) = ins("INSERT INTO t (a, b) VALUES ('it''s', 'x,y')");
4567 assert_eq!(rows[0].doc.get("a"), Some(&json!("it's")));
4568 assert_eq!(rows[0].doc.get("b"), Some(&json!("x,y")));
4569 }
4570
4571 #[test]
4572 fn insert_refuses_what_it_cannot_store_faithfully() {
4573 assert!(translate("INSERT INTO t (a) VALUES (1 + 1)").is_err());
4575 assert!(translate("INSERT INTO t (a) VALUES (now())").is_err());
4576 let e = translate("INSERT INTO t (a, b) VALUES (1)").unwrap_err();
4578 assert!(e.contains("values for"), "{}", e);
4579 let e2 = translate("INSERT INTO t VALUES (1)").unwrap_err();
4581 assert!(e2.contains("explicit column list"), "{}", e2);
4582 }
4583
4584 #[test]
4585 fn update_finds_rows_with_the_full_predicate_surface() {
4586 match translate("UPDATE orders SET status = 'void' WHERE total < 50 AND region IN ('eu')") {
4587 Ok(Stmt::Update { coll, set, nql, .. }) => {
4588 assert_eq!(coll, "orders");
4589 assert_eq!(set, vec![("status".to_string(), json!("void"))]);
4590 assert_eq!(nql, r#"FROM orders WHERE total < 50 AND region IN ("eu")"#);
4592 }
4593 other => panic!("expected UPDATE, got {:?}", other),
4594 }
4595 }
4596
4597 #[test]
4598 fn update_without_where_targets_the_whole_collection() {
4599 match translate("UPDATE t SET a = 1") {
4601 Ok(Stmt::Update { nql, .. }) => assert_eq!(nql, "FROM t"),
4602 other => panic!("expected UPDATE, got {:?}", other),
4603 }
4604 }
4605
4606 #[test]
4607 fn update_handles_several_assignments() {
4608 match translate("UPDATE t SET a = 1, b = 'x,y', c = NULL WHERE id = 'k'") {
4609 Ok(Stmt::Update { set, .. }) => {
4610 assert_eq!(set.len(), 3);
4611 assert_eq!(set[1], ("b".to_string(), json!("x,y")));
4612 assert_eq!(set[2], ("c".to_string(), Value::Null));
4613 }
4614 other => panic!("expected UPDATE, got {:?}", other),
4615 }
4616 assert!(translate("UPDATE t SET").is_err());
4617 assert!(translate("UPDATE t SET a").is_err());
4618 }
4619
4620 #[test]
4621 fn delete_becomes_a_predicate_over_the_collection() {
4622 match translate("DELETE FROM orders WHERE status = 'void'") {
4623 Ok(Stmt::Delete { coll, nql, .. }) => {
4624 assert_eq!(coll, "orders");
4625 assert_eq!(nql, r#"FROM orders WHERE status = "void""#);
4626 }
4627 other => panic!("expected DELETE, got {:?}", other),
4628 }
4629 match translate("DELETE FROM t") {
4630 Ok(Stmt::Delete { nql, .. }) => assert_eq!(nql, "FROM t"),
4631 other => panic!("expected DELETE, got {:?}", other),
4632 }
4633 }
4634
4635 #[test]
4636 fn returning_is_parsed_off_every_write() {
4637 let (_, _, ret) = ins("INSERT INTO t (a) VALUES (1) RETURNING a, _id");
4638 assert_eq!(ret.iter().map(|c| c.out.clone()).collect::<Vec<_>>(), vec!["a", "_id"]);
4639 let (_, _, star) = ins("INSERT INTO t (a) VALUES (1) RETURNING *");
4642 assert!(star.is_empty());
4643 assert!(wants_returning("INSERT INTO t (a) VALUES (1) RETURNING *"));
4644 assert!(!wants_returning("INSERT INTO t (a) VALUES (1)"));
4645
4646 match translate("UPDATE t SET a = 1 WHERE id = 'k' RETURNING a") {
4647 Ok(Stmt::Update { nql, returning, .. }) => {
4648 assert_eq!(returning.len(), 1);
4649 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql);
4651 }
4652 other => panic!("expected UPDATE, got {:?}", other),
4653 }
4654 match translate("DELETE FROM t WHERE id = 'k' RETURNING *") {
4655 Ok(Stmt::Delete { nql, .. }) =>
4656 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql),
4657 other => panic!("expected DELETE, got {:?}", other),
4658 }
4659 }
4660
4661 #[test]
4662 fn a_keyword_inside_a_value_is_not_a_clause() {
4663 match translate("UPDATE t SET note = 'where returning from' WHERE id = 'k'") {
4664 Ok(Stmt::Update { set, nql, .. }) => {
4665 assert_eq!(set[0].1, json!("where returning from"));
4666 assert_eq!(nql, r#"FROM t WHERE id = "k""#);
4667 }
4668 other => panic!("expected UPDATE, got {:?}", other),
4669 }
4670 }
4671
4672 #[test]
4673 fn split_top_respects_quotes_and_nesting() {
4674 assert_eq!(split_top("a, b, c", ',').len(), 3);
4675 assert_eq!(split_top("(1, 2), (3, 4)", ',').len(), 2);
4676 assert_eq!(split_top("'a,b', c", ',').len(), 2);
4677 assert_eq!(split_top("'it''s, fine', c", ',').len(), 2);
4678 }
4679
4680 #[test]
4681 fn comments_and_whitespace_do_not_confuse_the_translator() {
4682 assert_eq!(q("SELECT *\n FROM orders -- trailing note\n"), "FROM orders");
4683 assert_eq!(q("SELECT /* inline */ * FROM orders"), "FROM orders");
4684 assert_eq!(q("SELECT * FROM t WHERE note = 'from here to JOIN'"),
4686 r#"FROM t WHERE note = "from here to JOIN""#);
4687 }
4688
4689 #[test]
4690 fn find_kw_ignores_quotes_parens_and_substrings() {
4691 assert_eq!(find_kw("SELECT A FROM B", "FROM"), Some(9));
4692 assert_eq!(find_kw("SELECT 'FROM' FROM B", "FROM"), Some(14));
4693 assert_eq!(find_kw("SELECT F(x FROM y) FROM B", "FROM"), Some(19));
4694 assert_eq!(find_kw("SELECT FROMAGE", "FROM"), None);
4695 assert_eq!(find_kw("SELECT X_FROM", "FROM"), None);
4696 }
4697
4698 #[test]
4701 fn provenance_columns_sort_after_the_users_own_fields() {
4702 let rows = vec![json!({"_id":"1","_hash":"ab","status":"paid","total":9})];
4703 assert_eq!(names(&columns_for(&rows, &[])),
4704 vec!["status", "total", "_hash", "_id"]);
4705 }
4706
4707 #[test]
4708 fn an_explicit_projection_sets_the_column_order() {
4709 let rows = vec![json!({"a":1,"b":2})];
4710 let p = vec![Col::same("b"), Col::same("a")];
4711 assert_eq!(names(&columns_for(&rows, &p)), vec!["b", "a"]);
4712 }
4713
4714 #[test]
4715 fn columns_are_the_union_across_sparse_rows() {
4716 let rows = vec![json!({"a":1}), json!({"b":2})];
4718 assert_eq!(names(&columns_for(&rows, &[])), vec!["a", "b"]);
4719 }
4720
4721 #[test]
4722 fn type_oids_follow_the_first_non_null_value() {
4723 let rows = vec![json!({"i":1,"f":1.5,"b":true,"s":"x","n":null})];
4724 assert_eq!(oid_for(&rows, "i"), OID_INT8);
4725 assert_eq!(oid_for(&rows, "f"), OID_FLOAT8);
4726 assert_eq!(oid_for(&rows, "b"), OID_BOOL);
4727 assert_eq!(oid_for(&rows, "s"), OID_TEXT);
4728 assert_eq!(oid_for(&rows, "n"), OID_TEXT);
4730 assert_eq!(oid_for(&rows, "absent"), OID_TEXT);
4731 }
4732
4733 #[test]
4734 fn a_column_that_is_null_in_the_first_row_still_gets_its_type() {
4735 let rows = vec![json!({"v": null}), json!({"v": 7})];
4736 assert_eq!(oid_for(&rows, "v"), OID_INT8);
4737 }
4738
4739 #[test]
4740 fn cells_render_in_postgres_text_format() {
4741 assert_eq!(cell(Some(&json!("x"))), Some("x".to_string()));
4742 assert_eq!(cell(Some(&json!(true))), Some("t".to_string()));
4743 assert_eq!(cell(Some(&json!(false))), Some("f".to_string()));
4744 assert_eq!(cell(Some(&json!(42))), Some("42".to_string()));
4745 assert_eq!(cell(Some(&json!(null))), None);
4746 assert_eq!(cell(None), None);
4747 assert_eq!(cell(Some(&json!({"a":1}))), Some("{\"a\":1}".to_string()));
4749 }
4750
4751 #[test]
4754 fn message_framing_length_excludes_the_tag() {
4755 let mut m = Out::msg(b'Z');
4756 m.bytes(b"I");
4757 let bytes = m.finish();
4758 assert_eq!(bytes[0], b'Z');
4759 assert_eq!(i32::from_be_bytes([bytes[1], bytes[2], bytes[3], bytes[4]]), 5);
4760 assert_eq!(bytes.len(), 6);
4761 }
4762
4763 #[test]
4764 fn a_result_set_encodes_as_description_then_rows_then_complete() {
4765 let rows = vec![json!({"a": 1}), json!({"a": 2})];
4766 let out = encode_result(&rows, &[]);
4767 assert_eq!(out[0], b'T');
4768 let tags: Vec<u8> = {
4769 let mut t = vec![];
4771 let mut i = 0usize;
4772 while i < out.len() {
4773 t.push(out[i]);
4774 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
4775 i += 1 + len;
4776 }
4777 t
4778 };
4779 assert_eq!(tags, vec![b'T', b'D', b'D', b'C'],
4780 "one description, one row each, one completion");
4781 }
4782
4783 #[test]
4788 fn a_write_with_returning_emits_exactly_one_command_complete() {
4789 let rows = vec![json!({"_id": "o1", "total": 9})];
4790 let mut out = encode_rows(&rows, &[Col::same("_id")]);
4791 out.extend_from_slice(&command_complete("INSERT 0 1"));
4792 let mut tags = vec![];
4793 let mut i = 0usize;
4794 while i < out.len() {
4795 tags.push(out[i]);
4796 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
4797 i += 1 + len;
4798 }
4799 assert_eq!(tags, vec![b'T', b'D', b'C'], "one description, one row, ONE tag");
4800 assert_eq!(tags.iter().filter(|t| **t == b'C').count(), 1);
4801 assert!(!encode_rows(&rows, &[]).contains(&b'C')
4803 || encode_rows(&rows, &[]).iter().filter(|b| **b == b'C').count() > 0);
4804 let bare = encode_rows(&rows, &[Col::same("_id")]);
4805 let mut bare_tags = vec![];
4806 let mut j = 0usize;
4807 while j < bare.len() {
4808 bare_tags.push(bare[j]);
4809 let len = i32::from_be_bytes([bare[j+1], bare[j+2], bare[j+3], bare[j+4]]) as usize;
4810 j += 1 + len;
4811 }
4812 assert_eq!(bare_tags, vec![b'T', b'D'], "encode_rows never appends a tag");
4813 }
4814
4815 #[test]
4816 fn an_empty_result_still_sends_a_description() {
4817 let out = encode_result(&[], &[Col::same("a")]);
4818 assert_eq!(out[0], b'T', "clients need the shape even with no rows");
4819 }
4820
4821 #[test]
4822 fn statements_split_on_top_level_semicolons_only() {
4823 assert_eq!(split_statements("SELECT 1; SELECT 2").len(), 2);
4824 assert_eq!(split_statements("SELECT ';'").len(), 1);
4825 assert_eq!(split_statements("SELECT 1;").len(), 1);
4826 assert_eq!(split_statements(" ").len(), 0);
4827 }
4828
4829 #[test]
4830 fn an_error_names_its_sqlstate() {
4831 let e = String::from_utf8_lossy(&err_msg("0A000", "x")).to_string();
4832 assert!(e.contains("ERROR"));
4833 assert!(e.contains("0A000"));
4834 }
4835
4836 #[test]
4839 fn placeholders_are_counted_outside_string_literals() {
4840 assert_eq!(param_count("SELECT a FROM t WHERE b = $1 AND c = $2"), 2);
4841 assert_eq!(param_count("SELECT a FROM t"), 0);
4842 assert_eq!(param_count("WHERE a = $2 OR b = $2 OR c = $1"), 2);
4844 assert_eq!(param_count("SELECT a FROM t WHERE b = '$1'"), 0,
4845 "a placeholder inside a literal is data, not a parameter");
4846 assert_eq!(param_count("WHERE a = $10 AND b = $1"), 10,
4847 "two-digit indexes must not be read as $1 followed by 0");
4848 }
4849
4850 #[test]
4851 fn parameters_are_spliced_as_literals() {
4852 let out = substitute_params("WHERE a = $1 AND b = $2 AND c = $3",
4853 &[Some("'x'".into()), Some("42".into()), None]).unwrap();
4854 assert_eq!(out, "WHERE a = 'x' AND b = 42 AND c = NULL");
4855 }
4856
4857 #[test]
4858 fn substitution_leaves_string_literals_alone() {
4859 let out = substitute_params("WHERE a = '$1' AND b = $1", &[Some("9".into())]).unwrap();
4860 assert_eq!(out, "WHERE a = '$1' AND b = 9");
4861 }
4862
4863 #[test]
4864 fn too_few_parameters_is_an_error_not_a_silent_null() {
4865 let e = substitute_params("WHERE a = $2", &[Some("1".into())]).unwrap_err();
4868 assert!(e.contains("$2"), "{}", e);
4869 }
4870
4871 #[test]
4872 fn a_quote_in_a_parameter_cannot_escape_its_literal() {
4873 let lit = decode_param(Some(b"it's"), OID_TEXT, 0).unwrap().unwrap();
4874 assert_eq!(lit, "'it''s'");
4875 let out = substitute_params("WHERE a = $1", &[Some(lit)]).unwrap();
4877 assert_eq!(out, "WHERE a = 'it''s'");
4878 }
4879
4880 #[test]
4881 fn binary_parameters_decode_in_every_width_psycopg_sends() {
4882 assert_eq!(decode_param(Some(&[0x00, 0x2a]), OID_INT2, 1).unwrap().unwrap(), "42");
4885 assert_eq!(decode_param(Some(&[0, 0, 0, 7]), OID_INT4, 1).unwrap().unwrap(), "7");
4886 assert_eq!(
4887 decode_param(Some(&[0, 0, 0, 0, 0, 0, 0, 9]), OID_INT8, 1).unwrap().unwrap(), "9");
4888 assert_eq!(
4889 decode_param(Some(&0x400c_0000_0000_0000u64.to_be_bytes()), OID_FLOAT8, 1)
4890 .unwrap().unwrap(), "3.5");
4891 assert_eq!(decode_param(Some(&[1]), OID_BOOL, 1).unwrap().unwrap(), "TRUE");
4892 assert_eq!(decode_param(Some(&[0]), OID_BOOL, 1).unwrap().unwrap(), "FALSE");
4893 }
4894
4895 #[test]
4896 fn a_negative_binary_integer_keeps_its_sign() {
4897 assert_eq!(decode_param(Some(&(-5i32).to_be_bytes()), OID_INT4, 1).unwrap().unwrap(), "-5");
4898 assert_eq!(decode_param(Some(&(-5i16).to_be_bytes()), OID_INT2, 1).unwrap().unwrap(), "-5");
4899 }
4900
4901 #[test]
4902 fn a_binary_parameter_of_the_wrong_width_is_refused() {
4903 let e = decode_param(Some(&[0x2a]), OID_INT4, 1).unwrap_err();
4906 assert!(e.contains("4 bytes"), "{}", e);
4907 }
4908
4909 #[test]
4910 fn an_unspecified_text_parameter_is_treated_as_a_string() {
4911 assert_eq!(decode_param(Some(b"hello"), 0, 0).unwrap().unwrap(), "'hello'");
4914 }
4915
4916 #[test]
4917 fn a_null_parameter_decodes_to_none_in_every_format() {
4918 assert_eq!(decode_param(None, OID_TEXT, 0).unwrap(), None);
4919 assert_eq!(decode_param(None, OID_INT8, 1).unwrap(), None);
4920 }
4921
4922 #[test]
4923 fn an_unsupported_binary_type_says_so_by_name() {
4924 let e = decode_param(Some(&[0u8; 8]), 1114, 1).unwrap_err();
4925 assert!(e.contains("1114"), "{}", e);
4926 assert!(e.contains("text"), "the error should point at the way out: {}", e);
4927 }
4928
4929 #[test]
4930 fn a_text_number_that_is_not_a_number_gets_quoted() {
4931 assert_eq!(decode_param(Some(b"oops"), OID_INT8, 0).unwrap().unwrap(), "'oops'");
4934 }
4935
4936 #[test]
4937 fn a_client_declared_type_is_believed_over_inference() {
4938 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[OID_INT4, 0], None);
4941 assert_eq!(oids, vec![OID_INT4, OID_TEXT]);
4942 }
4943
4944 #[test]
4945 fn parameter_arity_is_taken_from_the_sql_when_the_client_declares_none() {
4946 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[], None);
4949 assert_eq!(oids.len(), 2);
4950 }
4951
4952 #[test]
4953 fn the_field_behind_each_placeholder_is_identified() {
4954 assert_eq!(
4955 param_fields("SELECT a FROM t WHERE qty > $1 AND status = $2", 2),
4956 vec![Some("qty".to_string()), Some("status".to_string())]);
4957 }
4958
4959 #[test]
4960 fn word_operators_do_not_hide_the_field() {
4961 assert_eq!(param_fields("SELECT a FROM t WHERE name LIKE $1", 1),
4962 vec![Some("name".to_string())]);
4963 assert_eq!(param_fields("SELECT a FROM t WHERE qty BETWEEN $1 AND $2", 2),
4964 vec![Some("qty".to_string()), Some("qty".to_string())]);
4965 assert_eq!(param_fields("SELECT a FROM t WHERE region IN ($1, $2)", 2),
4966 vec![Some("region".to_string()), Some("region".to_string())]);
4967 }
4968
4969 #[test]
4970 fn a_clause_position_types_from_the_grammar_not_from_a_column() {
4971 assert_eq!(
4975 infer_param_oids("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", &[], None),
4976 vec![OID_INT8, OID_TEXT]);
4977 assert_eq!(infer_param_oids("SELECT a FROM t AS OF $1", &[], None), vec![OID_INT8]);
4978 assert_eq!(
4981 infer_param_oids("SELECT a FROM t VALID AS OF $1", &[], None), vec![OID_TEXT]);
4982 assert_eq!(
4983 infer_param_oids("SELECT a FROM t LIMIT $1 OFFSET $2", &[], None),
4984 vec![OID_INT8, OID_INT8]);
4985 }
4986
4987 #[test]
4988 fn an_aggregate_column_types_from_what_the_aggregate_means() {
4989 assert_eq!(aggregate_oid("count", None, "t"), Some(OID_INT8));
4993 assert_eq!(aggregate_oid("avg_fee", None, "t"), Some(OID_FLOAT8),
4994 "an average is fractional even over integers");
4995 assert_eq!(aggregate_oid("max__seq", None, "t"), Some(OID_INT8));
4998 assert_eq!(aggregate_oid("total", None, "t"), None, "not an aggregate");
4999 }
5000
5001 #[test]
5002 fn the_parse_probe_uses_a_literal_that_every_clause_accepts() {
5003 let probe = probe_sql("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", 2);
5007 assert!(!probe.contains("NULL"), "{}", probe);
5008 assert!(translate(&probe).is_ok(), "the probe must parse: {}", probe);
5009 }
5010
5011 #[test]
5012 fn a_column_with_mixed_types_across_documents_is_advertised_as_text() {
5013 let rows = vec![json!({"x": 3}), json!({"x": "n/a"})];
5017 assert_eq!(oid_for(&rows, "x"), OID_TEXT);
5018 let rows = vec![json!({"x": 3}), json!({"x": 1.5})];
5020 assert_eq!(oid_for(&rows, "x"), OID_FLOAT8);
5021 let rows = vec![json!({"x": Value::Null}), json!({"x": 7})];
5023 assert_eq!(oid_for(&rows, "x"), OID_INT8);
5024 }
5025
5026 #[test]
5027 fn binary_output_encodes_each_advertised_type() {
5028 assert_eq!(cell_binary(Some(&json!(true)), OID_BOOL).unwrap().unwrap(), vec![1]);
5029 assert_eq!(cell_binary(Some(&json!(42)), OID_INT8).unwrap().unwrap(),
5030 42i64.to_be_bytes().to_vec());
5031 assert_eq!(cell_binary(Some(&json!(3.5)), OID_FLOAT8).unwrap().unwrap(),
5032 3.5f64.to_be_bytes().to_vec());
5033 assert_eq!(cell_binary(Some(&json!("hi")), OID_TEXT).unwrap().unwrap(), b"hi".to_vec());
5035 assert_eq!(cell_binary(Some(&Value::Null), OID_INT8).unwrap(), None);
5036 assert_eq!(cell(Some(&json!(true))).unwrap(), "t");
5038 }
5039
5040 #[test]
5041 fn a_value_that_does_not_fit_its_advertised_binary_type_is_refused() {
5042 let e = cell_binary(Some(&json!("nope")), OID_INT8).unwrap_err();
5047 assert!(e.contains("a string"), "{}", e);
5048 assert!(e.contains("more than one type"), "the error should explain WHY: {}", e);
5049 }
5050
5051 #[test]
5052 fn a_row_description_carries_the_requested_format_per_column() {
5053 let cols = [Col::same("a"), Col::same("b")];
5054 let m = row_description_fmt(&cols, &[OID_INT8, OID_TEXT], &[1, 0]);
5055 assert_eq!(m[0], b'T');
5056 assert_eq!(m[m.len() - 1], 0, "the last column was requested as text");
5058 }
5059
5060 #[test]
5061 fn a_qualified_column_resolves_to_its_bare_name() {
5062 assert_eq!(param_fields("SELECT a FROM t WHERE t.qty = $1", 1),
5063 vec![Some("qty".to_string())]);
5064 }
5065
5066 #[test]
5067 fn insert_placeholders_map_positionally_to_the_column_list() {
5068 assert_eq!(
5069 param_fields("INSERT INTO t (_id, qty, status) VALUES ($1, $2, $3)", 3),
5070 vec![Some("_id".to_string()), Some("qty".to_string()), Some("status".to_string())]);
5071 }
5072
5073 #[test]
5074 fn a_set_clause_placeholder_finds_its_column() {
5075 assert_eq!(param_fields("UPDATE t SET status = $1 WHERE _id = $2", 2),
5076 vec![Some("status".to_string()), Some("_id".to_string())]);
5077 }
5078
5079 #[test]
5080 fn the_target_collection_is_found_for_every_statement_kind() {
5081 assert_eq!(stmt_collection("SELECT a FROM inv WHERE b = $1"), "inv");
5082 assert_eq!(stmt_collection("UPDATE inv SET a = $1"), "inv");
5083 assert_eq!(stmt_collection("DELETE FROM inv WHERE a = $1"), "inv");
5084 assert_eq!(stmt_collection("INSERT INTO inv (a) VALUES ($1)"), "inv");
5085 assert_eq!(stmt_collection("SELECT a FROM public.inv"), "inv");
5087 assert_eq!(stmt_collection("INSERT INTO inv(a) VALUES ($1)"), "inv");
5088 }
5089
5090 #[test]
5091 fn engine_metadata_fields_type_without_touching_storage() {
5092 assert_eq!(infer_field_oid(None, "t", "_seq"), OID_INT8);
5093 assert_eq!(infer_field_oid(None, "t", "_id"), OID_TEXT);
5094 }
5095
5096 #[test]
5097 fn the_protocol_acknowledgements_are_single_empty_messages() {
5098 for (m, tag) in [
5100 (parse_complete(), b'1'), (bind_complete(), b'2'),
5101 (close_complete(), b'3'), (no_data(), b'n'), (portal_suspended(), b's'),
5102 ] {
5103 assert_eq!(m.len(), 5, "{:?}", tag as char);
5104 assert_eq!(m[0], tag);
5105 assert_eq!(i32::from_be_bytes([m[1], m[2], m[3], m[4]]), 4);
5106 }
5107 }
5108
5109 #[test]
5110 fn parameter_description_reports_its_arity_and_types() {
5111 let m = parameter_description(&[OID_TEXT, OID_INT8]);
5112 assert_eq!(m[0], b't');
5113 assert_eq!(i16::from_be_bytes([m[5], m[6]]), 2);
5114 assert_eq!(i32::from_be_bytes([m[7], m[8], m[9], m[10]]), OID_TEXT);
5115 assert_eq!(i32::from_be_bytes([m[11], m[12], m[13], m[14]]), OID_INT8);
5116 }
5117
5118 #[test]
5119 fn a_cstring_is_taken_without_its_terminator() {
5120 let body = b"one\0two\0".to_vec();
5121 let mut at = 0usize;
5122 assert_eq!(take_cstr(&body, &mut at), "one");
5123 assert_eq!(take_cstr(&body, &mut at), "two");
5124 assert_eq!(at, body.len());
5125 }
5126
5127 #[test]
5128 fn truncated_integers_are_reported_rather_than_read_past_the_end() {
5129 let body = vec![0u8, 1];
5130 let mut at = 0usize;
5131 assert!(take_i32(&body, &mut at).is_err());
5132 let mut at = 0usize;
5133 assert!(take_i16(&body, &mut at).is_ok());
5134 }
5135
5136 #[test]
5137 fn a_binary_result_format_request_is_refused_rather_than_faked() {
5138 let out = encode_rows(&[], &[Col::same("a")]);
5141 let desc_format = &out[out.len() - 2..];
5142 assert_eq!(i16::from_be_bytes([desc_format[0], desc_format[1]]), 0,
5143 "every column is advertised as text format");
5144 }
5145
5146 #[test]
5147 fn a_float_parameter_does_not_render_as_rust_infinity() {
5148 assert_eq!(fmt_float(f64::INFINITY), "'Infinity'");
5149 assert_eq!(fmt_float(f64::NEG_INFINITY), "'-Infinity'");
5150 assert_eq!(fmt_float(f64::NAN), "'NaN'");
5151 assert_eq!(fmt_float(3.0), "3", "a whole float should not gain a .0 tail");
5152 assert_eq!(fmt_float(3.5), "3.5");
5153 }
5154}