1use std::collections::HashMap;
128use std::sync::Arc;
129
130use serde_json::Value;
131use tokio::io::{AsyncReadExt, AsyncWriteExt};
132use tokio::net::{TcpListener, TcpStream};
133
134use crate::db::Db;
135
136const OID_BOOL: i32 = 16;
138const OID_INT8: i32 = 20;
139const OID_FLOAT8: i32 = 701;
140const OID_TEXT: i32 = 25;
141
142const PROTO_V3: i32 = 196_608; const SSL_REQUEST: i32 = 80_877_103;
144const GSS_REQUEST: i32 = 80_877_104;
145const CANCEL_REQUEST: i32 = 80_877_102;
146
147pub trait DbResolver: Send + Sync + 'static {
153 fn resolve(&self, name: &str) -> Option<Arc<Db>>;
161 fn token(&self) -> Option<String> {
163 None
164 }
165}
166
167struct Out(Vec<u8>);
170
171impl Out {
172 fn msg(tag: u8) -> Self {
173 Out(vec![tag, 0, 0, 0, 0])
175 }
176 fn i16(&mut self, v: i16) { self.0.extend_from_slice(&v.to_be_bytes()); }
177 fn i32(&mut self, v: i32) { self.0.extend_from_slice(&v.to_be_bytes()); }
178 fn cstr(&mut self, s: &str) {
179 self.0.extend_from_slice(s.replace('\0', "").as_bytes());
182 self.0.push(0);
183 }
184 fn bytes(&mut self, b: &[u8]) { self.0.extend_from_slice(b); }
185 fn finish(mut self) -> Vec<u8> {
187 let len = (self.0.len() - 1) as i32;
188 self.0[1..5].copy_from_slice(&len.to_be_bytes());
189 self.0
190 }
191}
192
193fn err_msg(code: &str, message: &str) -> Vec<u8> {
194 let mut m = Out::msg(b'E');
195 m.bytes(b"S"); m.cstr("ERROR");
196 m.bytes(b"C"); m.cstr(code);
197 m.bytes(b"M"); m.cstr(message);
198 m.0.push(0);
199 m.finish()
200}
201
202fn ready() -> Vec<u8> {
203 let mut m = Out::msg(b'Z');
204 m.bytes(b"I"); m.finish()
206}
207
208fn command_complete(tag: &str) -> Vec<u8> {
209 let mut m = Out::msg(b'C');
210 m.cstr(tag);
211 m.finish()
212}
213
214#[derive(Debug, PartialEq, Clone)]
224pub struct Col {
225 pub src: String,
226 pub out: String,
227}
228
229impl Col {
230 fn same(name: &str) -> Self {
231 Col { src: name.to_string(), out: name.to_string() }
232 }
233 fn renamed(src: &str, out: &str) -> Self {
234 Col { src: src.to_string(), out: out.to_string() }
235 }
236}
237
238#[derive(Debug, PartialEq)]
256pub enum Stmt {
257 Query { nql: String, project: Vec<Col> },
259 Insert { coll: String, rows: Vec<InsertRow>, returning: Vec<Col> },
261 Update { coll: String, set: Vec<(String, Value)>, nql: String, returning: Vec<Col> },
263 Delete { coll: String, nql: String, returning: Vec<Col> },
265 Canned { cols: Vec<String>, row: Vec<String> },
267 Ok(&'static str),
269}
270
271#[derive(Debug, PartialEq, Clone)]
274pub struct InsertRow {
275 pub id: Option<String>,
277 pub doc: serde_json::Map<String, Value>,
278 pub caused_by: Vec<String>,
281 pub valid_from: Option<String>,
282 pub valid_to: Option<String>,
283}
284
285fn normalise(sql: &str) -> String {
288 let mut out = String::with_capacity(sql.len());
289 let mut chars = sql.chars().peekable();
290 let mut in_s = false;
291 while let Some(c) = chars.next() {
292 if in_s {
293 out.push(c);
294 if c == '\'' { in_s = false; }
295 continue;
296 }
297 match c {
298 '\'' => { in_s = true; out.push(c); }
299 '-' if chars.peek() == Some(&'-') => {
300 for n in chars.by_ref() { if n == '\n' { break; } }
302 out.push(' ');
303 }
304 '/' if chars.peek() == Some(&'*') => {
305 chars.next();
306 let mut prev = ' ';
307 while let Some(n) = chars.next() {
308 if prev == '*' && n == '/' { break; }
309 prev = n;
310 }
311 out.push(' ');
312 }
313 _ => out.push(c),
314 }
315 }
316 out.split_whitespace().collect::<Vec<_>>().join(" ")
317}
318
319fn sql_literals_to_nql(s: &str) -> String {
326 let mut out = String::with_capacity(s.len());
327 let mut it = s.chars().peekable();
328 while let Some(c) = it.next() {
329 match c {
330 '\'' => {
331 out.push('"');
332 while let Some(ch) = it.next() {
333 if ch == '\'' {
334 if it.peek() == Some(&'\'') {
335 it.next();
336 out.push('\''); } else {
338 break;
339 }
340 } else if ch == '"' {
341 out.push('\\');
344 out.push('"');
345 } else {
346 out.push(ch);
347 }
348 }
349 out.push('"');
350 }
351 '<' if it.peek() == Some(&'>') => { it.next(); out.push_str("!="); }
352 _ => out.push(c),
353 }
354 }
355 out
356}
357
358fn strip_prefix_ci(s: &str, prefix: &str) -> Option<String> {
359 if s.len() >= prefix.len() && s[..prefix.len()].eq_ignore_ascii_case(prefix) {
360 Some(s[prefix.len()..].trim_start().to_string())
361 } else {
362 None
363 }
364}
365
366fn find_kw(s: &str, kw: &str) -> Option<usize> {
369 let bytes = s.as_bytes();
370 let k = kw.as_bytes();
371 let mut depth = 0i32;
372 let mut in_s = false;
373 let mut in_d = false;
374 let mut i = 0usize;
375 while i < bytes.len() {
376 let c = bytes[i];
377 if in_s { if c == b'\'' { in_s = false; } i += 1; continue; }
378 if in_d { if c == b'"' { in_d = false; } i += 1; continue; }
379 match c {
380 b'\'' => { in_s = true; i += 1; continue; }
381 b'"' => { in_d = true; i += 1; continue; }
382 b'(' => { depth += 1; i += 1; continue; }
383 b')' => { depth -= 1; i += 1; continue; }
384 _ => {}
385 }
386 if depth == 0 && i + k.len() <= bytes.len()
387 && bytes[i..i + k.len()].eq_ignore_ascii_case(k)
388 {
389 let before_ok = i == 0 || !(bytes[i - 1] as char).is_alphanumeric() && bytes[i - 1] != b'_';
390 let after = i + k.len();
391 let after_ok = after >= bytes.len()
392 || !(bytes[after] as char).is_alphanumeric() && bytes[after] != b'_';
393 if before_ok && after_ok {
394 return Some(i);
395 }
396 }
397 i += 1;
398 }
399 None
400}
401
402fn split_top(s: &str, sep: char) -> Vec<String> {
406 let mut out = vec![];
407 let mut cur = String::new();
408 let mut depth = 0i32;
409 let mut in_s = false;
410 let mut it = s.chars().peekable();
411 while let Some(c) = it.next() {
412 if in_s {
413 cur.push(c);
414 if c == '\'' {
415 if it.peek() == Some(&'\'') { cur.push(it.next().unwrap()); } else { in_s = false; }
417 }
418 continue;
419 }
420 match c {
421 '\'' => { in_s = true; cur.push(c); }
422 '(' => { depth += 1; cur.push(c); }
423 ')' => { depth -= 1; cur.push(c); }
424 x if x == sep && depth == 0 => { out.push(cur.trim().to_string()); cur.clear(); }
425 _ => cur.push(c),
426 }
427 }
428 if !cur.trim().is_empty() { out.push(cur.trim().to_string()); }
429 out
430}
431
432fn sql_value(raw: &str) -> Result<Value, String> {
438 let t = raw.trim();
439 if t.is_empty() {
440 return Err("empty value".into());
441 }
442 let up = t.to_uppercase();
443 if up == "NULL" { return Ok(Value::Null); }
444 if up == "TRUE" { return Ok(Value::Bool(true)); }
445 if up == "FALSE" { return Ok(Value::Bool(false)); }
446 if t.starts_with('\'') && t.ends_with('\'') && t.len() >= 2 {
447 let inner = &t[1..t.len() - 1];
449 return Ok(Value::String(inner.replace("''", "'")));
450 }
451 if let Ok(i) = t.parse::<i64>() { return Ok(Value::from(i)); }
452 if let Ok(f) = t.parse::<f64>() { return Ok(Value::from(f)); }
453 Err(format!(
454 "cannot use {:?} as a value — this endpoint accepts string literals, \
455 numbers, TRUE/FALSE and NULL. Expressions, casts and function calls \
456 are not evaluated, because storing an unevaluated expression as text \
457 would be worse than refusing it", t))
458}
459
460fn split_returning(tail: &str) -> (String, Vec<Col>) {
462 let tu = tail.to_uppercase();
463 match find_kw(&tu, "RETURNING") {
464 None => (tail.to_string(), vec![]),
465 Some(at) => {
466 let head = tail[..at].trim().to_string();
467 let list = tail[at + "RETURNING".len()..].trim();
468 if list == "*" {
469 return (head, vec![]); }
471 let cols = split_top(list, ',')
472 .into_iter()
473 .map(|p| {
474 let raw = p.split_whitespace().next().unwrap_or(&p).to_string();
475 let name = raw.rsplit('.').next().unwrap_or(&raw).trim_matches('"').to_string();
476 Col::same(&name)
477 })
478 .collect();
479 (head, cols)
480 }
481 }
482}
483
484fn take_reserved(doc: &mut serde_json::Map<String, Value>) -> (Option<String>, Vec<String>, Option<String>, Option<String>) {
486 let id = doc.remove("_id").or_else(|| doc.remove("id"))
487 .and_then(|v| match v {
488 Value::String(s) => Some(s),
489 Value::Null => None,
490 other => Some(other.to_string()), });
492 let caused_by = match doc.remove("_caused_by") {
493 Some(Value::String(s)) => vec![s],
494 Some(Value::Array(a)) => a.into_iter()
495 .filter_map(|v| v.as_str().map(str::to_string)).collect(),
496 _ => vec![],
497 };
498 let vf = doc.remove("_valid_from").and_then(|v| v.as_str().map(str::to_string));
499 let vt = doc.remove("_valid_to").and_then(|v| v.as_str().map(str::to_string));
500 (id, caused_by, vf, vt)
501}
502
503fn translate_insert(sql: &str) -> Result<Stmt, String> {
505 let rest = strip_prefix_ci(sql, "INSERT")
506 .and_then(|r| strip_prefix_ci(&r, "INTO"))
507 .ok_or("expected INSERT INTO")?;
508 let ru = rest.to_uppercase();
512 let values_at = find_kw(&ru, "VALUES").ok_or(
513 "expected VALUES — `INSERT … SELECT` is not supported on this endpoint")?;
514 let head = rest[..values_at].trim().to_string();
515 let open = head.find('(').ok_or(
516 "INSERT needs an explicit column list — `INSERT INTO t (a, b) VALUES (…)`. \
517 NEDB is schemaless, so there is no declared column order to infer from")?;
518 let coll = head[..open].trim().trim_matches('"');
519 let coll = coll.rsplit('.').next().unwrap_or(coll).to_string();
520 if coll.is_empty() {
521 return Err("expected a collection name after INSERT INTO".into());
522 }
523 let close = head.rfind(')').ok_or("unterminated column list")?;
524 if close < open {
525 return Err("malformed column list".into());
526 }
527 let tail_from_values = rest[values_at..].to_string();
528 let cols: Vec<String> = split_top(&head[open + 1..close], ',')
529 .into_iter()
530 .map(|c| c.trim().trim_matches('"').to_string())
531 .collect();
532 if cols.is_empty() {
533 return Err("the column list is empty".into());
534 }
535
536 let after = strip_prefix_ci(&tail_from_values, "VALUES")
537 .ok_or("expected VALUES after the column list")?;
538 let (values_part, returning) = split_returning(&after);
539
540 let mut rows = vec![];
541 for group in split_top(&values_part, ',') {
542 let g = group.trim();
543 if !(g.starts_with('(') && g.ends_with(')')) {
544 return Err(format!("expected a parenthesised row of values, got {:?}", g));
545 }
546 let vals = split_top(&g[1..g.len() - 1], ',');
547 if vals.len() != cols.len() {
548 return Err(format!(
549 "{} values for {} columns — every row must match the column list",
550 vals.len(), cols.len()));
551 }
552 let mut doc = serde_json::Map::new();
553 for (c, v) in cols.iter().zip(vals.iter()) {
554 doc.insert(c.clone(), sql_value(v)?);
555 }
556 let (id, caused_by, valid_from, valid_to) = take_reserved(&mut doc);
557 rows.push(InsertRow { id, doc, caused_by, valid_from, valid_to });
558 }
559 if rows.is_empty() {
560 return Err("INSERT with no rows".into());
561 }
562 Ok(Stmt::Insert { coll, rows, returning })
563}
564
565fn translate_update(sql: &str) -> Result<Stmt, String> {
567 let rest = strip_prefix_ci(sql, "UPDATE").ok_or("expected UPDATE")?;
568 let ru = rest.to_uppercase();
569 let set_at = find_kw(&ru, "SET").ok_or("expected SET in UPDATE")?;
570 let coll = rest[..set_at].trim().trim_matches('"');
571 let coll = coll.rsplit('.').next().unwrap_or(coll).to_string();
572 if coll.is_empty() {
573 return Err("expected a collection name after UPDATE".into());
574 }
575 let after_set = rest[set_at + 3..].trim().to_string();
576 let (after_set, returning) = split_returning(&after_set);
577
578 let au = after_set.to_uppercase();
580 let (assigns_raw, where_raw) = match find_kw(&au, "WHERE") {
581 Some(at) => (after_set[..at].to_string(), after_set[at..].to_string()),
582 None => (after_set.clone(), String::new()),
583 };
584
585 let mut set = vec![];
586 for a in split_top(&assigns_raw, ',') {
587 let eq = a.find('=').ok_or(format!("expected `col = value` in SET, got {:?}", a))?;
588 let col = a[..eq].trim().trim_matches('"').to_string();
589 if col.is_empty() {
590 return Err("empty column name in SET".into());
591 }
592 set.push((col, sql_value(&a[eq + 1..])?));
593 }
594 if set.is_empty() {
595 return Err("UPDATE with no assignments".into());
596 }
597 let nql = format!("FROM {} {}", coll, sql_literals_to_nql(where_raw.trim()))
600 .trim().to_string();
601 Ok(Stmt::Update { coll, set, nql, returning })
602}
603
604fn translate_delete(sql: &str) -> Result<Stmt, String> {
606 let rest = strip_prefix_ci(sql, "DELETE")
607 .and_then(|r| strip_prefix_ci(&r, "FROM"))
608 .ok_or("expected DELETE FROM")?;
609 let (rest, returning) = split_returning(&rest);
610 let end = rest.find(' ').unwrap_or(rest.len());
611 let coll = rest[..end].trim().trim_matches('"');
612 let coll = coll.rsplit('.').next().unwrap_or(coll).to_string();
613 if coll.is_empty() {
614 return Err("expected a collection name after DELETE FROM".into());
615 }
616 let where_raw = rest[end..].trim();
617 let nql = format!("FROM {} {}", coll, sql_literals_to_nql(where_raw))
618 .trim().to_string();
619 Ok(Stmt::Delete { coll, nql, returning })
620}
621
622pub fn translate(sql_raw: &str) -> Result<Stmt, String> {
624 let sql = normalise(sql_raw);
625 let sql = sql.trim().trim_end_matches(';').trim();
626 if sql.is_empty() {
627 return Ok(Stmt::Ok(""));
628 }
629 let upper = sql.to_uppercase();
630
631 if upper.starts_with("SET ") || upper.starts_with("BEGIN") || upper.starts_with("COMMIT")
636 || upper.starts_with("ROLLBACK") || upper.starts_with("DISCARD")
637 || upper.starts_with("LISTEN ") || upper.starts_with("UNLISTEN ")
638 {
639 return Ok(Stmt::Ok(if upper.starts_with("SET") { "SET" } else { "OK" }));
641 }
642 if upper.starts_with("SHOW ") {
643 let name = sql[5..].trim().to_lowercase();
644 let val = match name.as_str() {
645 "transaction_isolation" | "default_transaction_isolation" => "read committed",
646 "server_version" => SERVER_VERSION,
647 "server_encoding" | "client_encoding" => "UTF8",
648 "standard_conforming_strings" => "on",
649 "is_superuser" => "off",
650 _ => "",
651 };
652 return Ok(Stmt::Canned { cols: vec![name], row: vec![val.to_string()] });
653 }
654 if upper == "SELECT VERSION()" {
655 return Ok(Stmt::Canned {
656 cols: vec!["version".into()],
657 row: vec![full_version_string()],
658 });
659 }
660 if upper == "SELECT 1" || upper == "SELECT 1;" {
661 return Ok(Stmt::Canned { cols: vec!["?column?".into()], row: vec!["1".into()] });
662 }
663 if upper.starts_with("SELECT CURRENT_SCHEMA") {
664 return Ok(Stmt::Canned { cols: vec!["current_schema".into()], row: vec!["public".into()] });
665 }
666 if upper.starts_with("SELECT CURRENT_DATABASE") {
667 return Ok(Stmt::Canned { cols: vec!["current_database".into()], row: vec!["nedb".into()] });
668 }
669 if upper.starts_with("SELECT CURRENT_USER") || upper.starts_with("SELECT USER") {
670 return Ok(Stmt::Canned { cols: vec!["current_user".into()], row: vec!["nedb".into()] });
671 }
672
673 if upper.starts_with("INSERT") { return translate_insert(sql); }
677 if upper.starts_with("UPDATE") { return translate_update(sql); }
678 if upper.starts_with("DELETE") { return translate_delete(sql); }
679
680 for (kw, why) in [
682 ("CREATE", "DDL is not supported — collections are created implicitly by the first write to them, because NEDB is schemaless"),
683 ("ALTER", "DDL is not supported — there is no schema to alter"),
684 ("DROP", "DDL is not supported; drop a database with DELETE /v1/databases/<db>"),
685 ("TRUNCATE", "not supported, and not an oversight: NEDB is append-only so that history cannot be discarded. That is the product"),
686 ("COPY", "not supported; use GET /v1/databases/<db>/since for bulk export"),
687 ("GRANT", "there is no SQL-level privilege system; auth is the bearer token"),
688 ("REVOKE", "there is no SQL-level privilege system; auth is the bearer token"),
689 ] {
690 if upper.starts_with(kw) {
691 return Err(format!("{} is not supported — {}", kw, why));
692 }
693 }
694 if !upper.starts_with("SELECT") {
695 return Err(format!(
696 "only SELECT, INSERT, UPDATE and DELETE are supported on the Postgres \
697 endpoint (got {:?})",
698 sql.split_whitespace().next().unwrap_or("")
699 ));
700 }
701 for (kw, why) in [
702 (" JOIN ", "JOIN is not supported — NQL is single-collection; join in your client or model the relation with LINK/TRAVERSE"),
703 (" UNION ", "UNION is not supported"),
704 (" INTERSECT ", "INTERSECT is not supported"),
705 (" EXCEPT ", "EXCEPT is not supported"),
706 (" OVER (", "window functions are not supported"),
707 ("DISTINCT ", "DISTINCT is not supported — GROUP BY <col> gives the distinct values with counts"),
708 ] {
709 if upper.contains(kw) {
710 return Err(why.to_string());
711 }
712 }
713 if find_kw(&upper, "FROM").is_none() {
714 return Err("SELECT without FROM is not supported on this endpoint".into());
715 }
716
717 let after_select = strip_prefix_ci(sql, "SELECT").ok_or("expected SELECT")?;
719 let from_at = find_kw(&after_select.to_uppercase(), "FROM")
720 .ok_or("expected FROM after the select list")?;
721 let projection = after_select[..from_at].trim().to_string();
722 let rest = after_select[from_at + 4..].trim().to_string();
723 if rest.is_empty() {
724 return Err("expected a collection name after FROM".into());
725 }
726 if rest.starts_with('(') {
729 return Err("subqueries in FROM are not supported".into());
730 }
731 let coll_end = rest.find(' ').unwrap_or(rest.len());
732 let coll = &rest[..coll_end];
733 if coll.contains(',') {
734 return Err("selecting from more than one collection is not supported (no JOIN)".into());
735 }
736 let bare = coll.rsplit('.').next().unwrap_or(coll).trim_matches('"');
743 let qualified = coll
744 .split('.')
745 .map(|p| p.trim_matches('"'))
746 .collect::<Vec<_>>()
747 .join(".");
748 let coll = if qualified.starts_with("information_schema.") {
749 qualified.as_str()
750 } else {
751 bare
752 };
753 let tail = rest[coll_end..].trim();
754
755 let pu = projection.to_uppercase();
757 let mut agg_clause = String::new();
758 let mut project: Vec<Col> = vec![];
759
760 if projection == "*" {
761 } else if pu.starts_with("COUNT(") {
763 agg_clause = " COUNT".to_string();
766 project.push(Col::same("count"));
767 } else if let Some(agg) = ["SUM", "AVG", "MIN", "MAX"]
768 .iter()
769 .find(|a| pu.starts_with(&format!("{}(", a)))
770 {
771 let inner = projection[agg.len() + 1..]
772 .trim_end_matches(')')
773 .trim()
774 .to_string();
775 if inner.is_empty() || inner == "*" {
776 return Err(format!("{}() needs a column", agg));
777 }
778 agg_clause = format!(" {} {}", agg, inner);
779 project.push(Col::renamed(
781 &format!("{}_{}", agg.to_lowercase(), inner),
782 &agg.to_lowercase(),
783 ));
784 } else {
785 for part in projection.split(',') {
786 let p = part.trim();
787 if p.is_empty() {
788 return Err("empty column in the select list".into());
789 }
790 if p.contains('(') {
791 return Err(format!(
792 "expressions in the select list are not supported ({:?}) — \
793 supported: *, a column list, COUNT(*), or SUM/AVG/MIN/MAX(col)", p));
794 }
795 let raw = p.split_whitespace().next().unwrap_or(p);
797 let name = raw.rsplit('.').next().unwrap_or(raw).trim_matches('"');
798 project.push(Col::same(name));
799 }
800 }
801
802 let mut tail = tail.to_string();
809 let tu = tail.to_uppercase();
810 if let Some(at) = find_kw(&tu, "AS OF SYSTEM TIME") {
811 let before = tail[..at].to_string();
812 let after = tail[at + "AS OF SYSTEM TIME".len()..].trim_start().to_string();
813 let end = after.find(' ').unwrap_or(after.len());
815 let seq = after[..end].trim().trim_matches('\'').trim_matches('"').to_string();
816 if seq.parse::<u64>().is_err() {
817 return Err(format!(
818 "AS OF SYSTEM TIME takes a NEDB sequence number here, not a timestamp (got {:?}). \
819 NEDB's history is sequence-addressed and never garbage-collected, so a seq is \
820 exact where a wall-clock time would be approximate", seq));
821 }
822 tail = format!("{} AS OF {} {}", before.trim(), seq, after[end..].trim())
823 .trim()
824 .to_string();
825 }
826
827 let tu_all = tail.to_uppercase();
835 if let Some(gb_at) = find_kw(&tu_all, "GROUP BY") {
836 let after = tail[gb_at + "GROUP BY".len()..].trim_start();
837 let key_end = after.find(|c: char| c == ' ' || c == ',').unwrap_or(after.len());
838 let group_key = after[..key_end].trim().trim_matches('"').to_string();
839 let is_agg = !agg_clause.is_empty();
840 for c in &project {
841 let ok = c.src == group_key
842 || c.src == "count"
843 || (is_agg && c.out == agg_clause.trim().split(' ').next()
844 .unwrap_or("").to_lowercase());
845 if !ok {
846 return Err(format!(
847 "column {:?} must appear in the GROUP BY clause or be used in an \
848 aggregate function — a grouped row carries the group key, `count`, \
849 and the aggregate, nothing else",
850 c.src));
851 }
852 }
853 }
854
855 let tail = sql_literals_to_nql(&tail);
856 let nql = format!("FROM {}{}{}", coll,
857 if agg_clause.is_empty() { String::new() } else { agg_clause },
858 if tail.is_empty() { String::new() } else { format!(" {}", tail) });
859
860 Ok(Stmt::Query { nql: nql.trim().to_string(), project })
861}
862
863const SERVER_VERSION: &str = "15.0";
864
865pub fn version_string() -> String {
867 full_version_string()
868}
869
870fn full_version_string() -> String {
871 format!(
872 "PostgreSQL {} (NEDB {}) — tamper-evident, append-only, permanent \
873 history. SELECT + INSERT/UPDATE/DELETE; an UPDATE is a new version, \
874 so prior values stay readable with AS OF SYSTEM TIME.",
875 SERVER_VERSION,
876 env!("CARGO_PKG_VERSION")
877 )
878}
879
880fn columns_for(rows: &[Value], project: &[Col]) -> Vec<Col> {
888 if !project.is_empty() {
889 return project.to_vec();
890 }
891 let mut plain: Vec<String> = vec![];
892 let mut meta: Vec<String> = vec![];
893 for r in rows {
894 if let Value::Object(m) = r {
895 for k in m.keys() {
896 let target = if k.starts_with('_') { &mut meta } else { &mut plain };
897 if !target.contains(k) {
898 target.push(k.clone());
899 }
900 }
901 }
902 }
903 plain.sort();
904 meta.sort();
905 plain.extend(meta);
906 plain.into_iter().map(|k| Col::same(&k)).collect()
907}
908
909fn oid_of_value(v: &Value) -> Option<i32> {
911 match v {
912 Value::Null => None,
913 Value::Bool(_) => Some(OID_BOOL),
914 Value::Number(n) => Some(if n.is_i64() || n.is_u64() { OID_INT8 } else { OID_FLOAT8 }),
915 Value::String(_) => Some(OID_TEXT),
916 _ => Some(OID_TEXT),
918 }
919}
920
921fn unify_oid(a: i32, b: i32) -> i32 {
928 if a == b {
929 return a;
930 }
931 match (a, b) {
932 (OID_INT8, OID_FLOAT8) | (OID_FLOAT8, OID_INT8) => OID_FLOAT8,
933 _ => OID_TEXT,
934 }
935}
936
937pub fn oid_for_column(rows: &[Value], col: &str) -> i32 {
949 oid_for(rows, col)
950}
951
952fn oid_for(rows: &[Value], col: &str) -> i32 {
953 let mut acc: Option<i32> = None;
954 for r in rows {
955 if let Some(o) = r.get(col).and_then(oid_of_value) {
956 acc = Some(match acc {
957 None => o,
958 Some(prev) => unify_oid(prev, o),
959 });
960 if acc == Some(OID_TEXT) {
961 break; }
963 }
964 }
965 acc.unwrap_or(OID_TEXT)
966}
967
968fn cell(v: Option<&Value>) -> Option<String> {
970 match v {
971 None | Some(Value::Null) => None, Some(Value::String(s)) => Some(s.clone()),
973 Some(Value::Bool(b)) => Some(if *b { "t".into() } else { "f".into() }),
974 Some(other) => Some(other.to_string()),
975 }
976}
977
978fn cell_binary(v: Option<&Value>, oid: i32) -> Result<Option<Vec<u8>>, String> {
990 let v = match v {
991 None | Some(Value::Null) => return Ok(None),
992 Some(v) => v,
993 };
994 let as_f64 = |n: &serde_json::Number| n.as_f64()
995 .ok_or_else(|| "a number too large to send as float8".to_string());
996 Ok(Some(match (oid, v) {
997 (OID_BOOL, Value::Bool(b)) => vec![u8::from(*b)],
998 (OID_INT2, Value::Number(n)) => {
999 let i = n.as_i64().ok_or("not an integer")?;
1000 i16::try_from(i).map_err(|_| format!("{} does not fit in int2", i))?
1001 .to_be_bytes().to_vec()
1002 }
1003 (OID_INT4, Value::Number(n)) => {
1004 let i = n.as_i64().ok_or("not an integer")?;
1005 i32::try_from(i).map_err(|_| format!("{} does not fit in int4", i))?
1006 .to_be_bytes().to_vec()
1007 }
1008 (OID_INT8, Value::Number(n)) => {
1009 n.as_i64().ok_or("not an integer")?.to_be_bytes().to_vec()
1010 }
1011 (OID_FLOAT4, Value::Number(n)) => (as_f64(n)? as f32).to_be_bytes().to_vec(),
1012 (OID_FLOAT8, Value::Number(n)) => as_f64(n)?.to_be_bytes().to_vec(),
1013 (OID_TEXT | OID_VARCHAR | OID_NAME | OID_UNKNOWN | OID_JSON, _) => {
1015 cell(Some(v)).unwrap_or_default().into_bytes()
1016 }
1017 (OID_JSONB, _) => {
1019 let mut b = vec![1u8];
1020 b.extend_from_slice(cell(Some(v)).unwrap_or_default().as_bytes());
1021 b
1022 }
1023 (oid, val) => {
1024 let kind = match val {
1025 Value::Bool(_) => "a boolean",
1026 Value::Number(_) => "a number",
1027 Value::String(_) => "a string",
1028 Value::Array(_) => "an array",
1029 _ => "an object",
1030 };
1031 return Err(format!(
1032 "cannot send {} in binary format as type OID {} — the field holds \
1033 more than one type across documents, so it cannot be described \
1034 by a single Postgres type. Select it with a text cast, or use a \
1035 text-format client",
1036 kind, oid
1037 ));
1038 }
1039 }))
1040}
1041
1042fn row_description_fmt(cols: &[Col], oids: &[i32], fmts: &[i16]) -> Vec<u8> {
1044 let mut m = Out::msg(b'T');
1045 m.i16(cols.len() as i16);
1046 for (i, c) in cols.iter().enumerate() {
1047 m.cstr(&c.out);
1048 m.i32(0); m.i16((i + 1) as i16); m.i32(oids.get(i).copied().unwrap_or(OID_TEXT));
1051 m.i16(-1); m.i32(-1); m.i16(fmts.get(i).copied().unwrap_or(0));
1054 }
1055 m.finish()
1056}
1057
1058fn row_description(cols: &[Col], oids: &[i32]) -> Vec<u8> {
1059 row_description_fmt(cols, oids, &[])
1060}
1061
1062fn data_row_bytes(vals: &[Option<Vec<u8>>]) -> Vec<u8> {
1063 let mut m = Out::msg(b'D');
1064 m.i16(vals.len() as i16);
1065 for v in vals {
1066 match v {
1067 None => m.i32(-1),
1068 Some(b) => {
1069 m.i32(b.len() as i32);
1070 m.bytes(b);
1071 }
1072 }
1073 }
1074 m.finish()
1075}
1076
1077fn data_row(vals: &[Option<String>]) -> Vec<u8> {
1078 let owned: Vec<Option<Vec<u8>>> =
1079 vals.iter().map(|v| v.as_ref().map(|s| s.as_bytes().to_vec())).collect();
1080 data_row_bytes(&owned)
1081}
1082
1083pub fn encode_rows(rows: &[Value], project: &[Col]) -> Vec<u8> {
1093 let cols = columns_for(rows, project);
1094 let oids: Vec<i32> = cols.iter().map(|c| oid_for(rows, &c.src)).collect();
1095 let mut out = row_description(&cols, &oids);
1096 for r in rows {
1097 let vals: Vec<Option<String>> = cols.iter().map(|c| cell(r.get(&c.src))).collect();
1098 out.extend_from_slice(&data_row(&vals));
1099 }
1100 out
1101}
1102
1103pub fn encode_result(rows: &[Value], project: &[Col]) -> Vec<u8> {
1105 let mut out = encode_rows(rows, project);
1106 out.extend_from_slice(&command_complete(&format!("SELECT {}", rows.len())));
1107 out
1108}
1109
1110const OID_INT2: i32 = 21;
1137const OID_INT4: i32 = 23;
1138const OID_OID: i32 = 26;
1139const OID_FLOAT4: i32 = 700;
1140const OID_VARCHAR: i32 = 1043;
1141const OID_NAME: i32 = 19;
1142const OID_UNKNOWN: i32 = 705;
1143const OID_JSON: i32 = 114;
1144const OID_JSONB: i32 = 3802;
1145
1146fn param_count(sql: &str) -> usize {
1152 let b = sql.as_bytes();
1153 let mut i = 0usize;
1154 let mut in_s = false;
1155 let mut max = 0usize;
1156 while i < b.len() {
1157 let c = b[i];
1158 if in_s {
1159 if c == b'\'' {
1160 in_s = false;
1161 }
1162 i += 1;
1163 continue;
1164 }
1165 if c == b'\'' {
1166 in_s = true;
1167 i += 1;
1168 continue;
1169 }
1170 if c == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1171 let mut j = i + 1;
1172 let mut n = 0usize;
1173 while j < b.len() && b[j].is_ascii_digit() {
1174 n = n * 10 + (b[j] - b'0') as usize;
1175 j += 1;
1176 }
1177 max = max.max(n);
1178 i = j;
1179 continue;
1180 }
1181 i += 1;
1182 }
1183 max
1184}
1185
1186fn infer_field_oid(db: Option<&Arc<Db>>, coll: &str, field: &str) -> i32 {
1195 match field {
1199 "_seq" => return OID_INT8,
1200 "_id" | "_hash" | "_prev" | "_collection" | "_valid_from" | "_valid_to" => return OID_TEXT,
1201 _ => {}
1202 }
1203 let db = match db {
1204 Some(db) => db,
1205 None => return OID_TEXT,
1206 };
1207 if coll.is_empty() || field.is_empty() {
1208 return OID_TEXT;
1209 }
1210 let rows = match crate::nql::query(db, &format!("FROM {} LIMIT {}", coll, TYPE_SAMPLE)) {
1211 Ok((rows, _)) => rows,
1212 Err(_) => return OID_TEXT,
1213 };
1214 oid_for(&rows, field)
1218}
1219
1220fn aggregate_oid(src: &str, db: Option<&Arc<Db>>, coll: &str) -> Option<i32> {
1232 if src == "count" {
1233 return Some(OID_INT8);
1234 }
1235 for (prefix, fixed) in [
1236 ("count_", Some(OID_INT8)),
1237 ("avg_", Some(OID_FLOAT8)),
1238 ("sum_", None),
1239 ("min_", None),
1240 ("max_", None),
1241 ] {
1242 if let Some(field) = src.strip_prefix(prefix) {
1243 return Some(match fixed {
1244 Some(oid) => oid,
1245 None => match infer_field_oid(db, coll, field) {
1248 OID_INT8 => OID_INT8,
1249 OID_FLOAT8 => OID_FLOAT8,
1250 other => other,
1253 },
1254 });
1255 }
1256 }
1257 None
1258}
1259
1260const TYPE_SAMPLE: usize = 200;
1266
1267fn stmt_collection(sql: &str) -> String {
1269 let s = normalise(sql);
1270 let up = s.to_uppercase();
1271 let after = if let Some(at) = find_kw(&up, "FROM") {
1272 &s[at + 4..]
1273 } else if let Some(rest) = strip_prefix_ci(&s, "UPDATE") {
1274 return rest
1275 .split_whitespace()
1276 .next()
1277 .unwrap_or("")
1278 .rsplit('.')
1279 .next()
1280 .unwrap_or("")
1281 .trim_matches('"')
1282 .to_string();
1283 } else if let Some(rest) = strip_prefix_ci(&s, "INSERT INTO") {
1284 return rest
1285 .split(|c: char| c.is_whitespace() || c == '(')
1286 .find(|t| !t.is_empty())
1287 .unwrap_or("")
1288 .rsplit('.')
1289 .next()
1290 .unwrap_or("")
1291 .trim_matches('"')
1292 .to_string();
1293 } else {
1294 return String::new();
1295 };
1296 after
1297 .trim()
1298 .split(|c: char| c.is_whitespace())
1299 .find(|t| !t.is_empty())
1300 .unwrap_or("")
1301 .rsplit('.')
1302 .next()
1303 .unwrap_or("")
1304 .trim_matches('"')
1305 .to_string()
1306}
1307
1308fn param_fields(sql: &str, n_params: usize) -> Vec<Option<String>> {
1319 let s = normalise(sql);
1320 let mut out = vec![None; n_params];
1321
1322 let up = s.to_uppercase();
1325 if up.starts_with("INSERT") {
1326 if let (Some(open), Some(vals_at)) = (s.find('('), find_kw(&up, "VALUES")) {
1327 if open < vals_at {
1328 if let Some(close) = s[open..vals_at].rfind(')') {
1329 let cols: Vec<String> = split_top(&s[open + 1..open + close], ',')
1330 .into_iter()
1331 .map(|c| c.trim().trim_matches('"').to_string())
1332 .collect();
1333 let tail = &s[vals_at..];
1335 let mut seen = 0usize;
1336 let b = tail.as_bytes();
1337 let mut i = 0usize;
1338 let mut in_s = false;
1339 while i < b.len() {
1340 if in_s {
1341 if b[i] == b'\'' { in_s = false; }
1342 i += 1;
1343 continue;
1344 }
1345 if b[i] == b'\'' { in_s = true; i += 1; continue; }
1346 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1347 let mut j = i + 1;
1348 let mut num = 0usize;
1349 while j < b.len() && b[j].is_ascii_digit() {
1350 num = num * 10 + (b[j] - b'0') as usize;
1351 j += 1;
1352 }
1353 if num >= 1 && num <= n_params {
1354 if let Some(c) = cols.get(seen % cols.len().max(1)) {
1355 out[num - 1] = Some(c.clone());
1356 }
1357 }
1358 seen += 1;
1359 i = j;
1360 continue;
1361 }
1362 i += 1;
1363 }
1364 return out;
1365 }
1366 }
1367 }
1368 }
1369
1370 let b = s.as_bytes();
1372 let mut i = 0usize;
1373 let mut in_s = false;
1374 while i < b.len() {
1375 if in_s {
1376 if b[i] == b'\'' { in_s = false; }
1377 i += 1;
1378 continue;
1379 }
1380 if b[i] == b'\'' { in_s = true; i += 1; continue; }
1381 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1382 let mut j = i + 1;
1383 let mut num = 0usize;
1384 while j < b.len() && b[j].is_ascii_digit() {
1385 num = num * 10 + (b[j] - b'0') as usize;
1386 j += 1;
1387 }
1388 if num >= 1 && num <= n_params {
1389 let left = &s[..i];
1390 let trimmed = left.trim_end_matches(|c: char| {
1393 c.is_whitespace() || "=<>!+-*/%(,".contains(c)
1394 });
1395 let mut tok = trimmed
1398 .rsplit(|c: char| c.is_whitespace() || c == '(' || c == ',')
1399 .find(|t| !t.is_empty())
1400 .unwrap_or("")
1401 .trim_matches('"');
1402 let mut before = trimmed;
1403 for _ in 0..4 {
1404 let upper_tok = tok.to_uppercase();
1405 if upper_tok.starts_with('$')
1411 || matches!(upper_tok.as_str(),
1412 "LIKE" | "ILIKE" | "IN" | "BETWEEN" | "AND" | "OR" | "NOT" | "IS") {
1413 before = before[..before.len() - tok.len()].trim_end_matches(|c: char| {
1414 c.is_whitespace() || "=<>!(,".contains(c)
1415 });
1416 tok = before
1417 .rsplit(|c: char| c.is_whitespace() || c == '(' || c == ',')
1418 .find(|t| !t.is_empty())
1419 .unwrap_or("")
1420 .trim_matches('"');
1421 } else {
1422 break;
1423 }
1424 }
1425 if !tok.is_empty()
1426 && tok.chars().all(|c| c.is_alphanumeric() || c == '_' || c == '.')
1427 && !tok.chars().next().map(|c| c.is_ascii_digit()).unwrap_or(true)
1428 {
1429 out[num - 1] = Some(tok.rsplit('.').next().unwrap_or(tok).to_string());
1430 }
1431 }
1432 i = j;
1433 continue;
1434 }
1435 i += 1;
1436 }
1437 out
1438}
1439
1440fn clause_param_oids(sql: &str, n_params: usize) -> Vec<Option<i32>> {
1450 let s = normalise(sql);
1451 let mut out = vec![None; n_params];
1452 let b = s.as_bytes();
1453 let mut i = 0usize;
1454 let mut in_s = false;
1455 while i < b.len() {
1456 if in_s {
1457 if b[i] == b'\'' { in_s = false; }
1458 i += 1;
1459 continue;
1460 }
1461 if b[i] == b'\'' { in_s = true; i += 1; continue; }
1462 if b[i] == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1463 let mut j = i + 1;
1464 let mut num = 0usize;
1465 while j < b.len() && b[j].is_ascii_digit() {
1466 num = num * 10 + (b[j] - b'0') as usize;
1467 j += 1;
1468 }
1469 if num >= 1 && num <= n_params {
1470 let left = s[..i].trim_end().to_uppercase();
1471 out[num - 1] = if left.ends_with("VALID AS OF") {
1474 Some(OID_TEXT)
1475 } else if left.ends_with("AS OF SYSTEM TIME")
1476 || left.ends_with("FOR SYSTEM_TIME AS OF")
1477 || left.ends_with("AS OF")
1478 || left.ends_with("LIMIT")
1479 || left.ends_with("OFFSET")
1480 {
1481 Some(OID_INT8)
1482 } else {
1483 None
1484 };
1485 }
1486 i = j;
1487 continue;
1488 }
1489 i += 1;
1490 }
1491 out
1492}
1493
1494fn infer_param_oids(sql: &str, declared: &[i32], db: Option<&Arc<Db>>) -> Vec<i32> {
1500 let n = param_count(sql).max(declared.len());
1501 if n == 0 {
1502 return vec![];
1503 }
1504 let coll = stmt_collection(sql);
1505 let fields = param_fields(sql, n);
1506 let clauses = clause_param_oids(sql, n);
1507 (0..n)
1508 .map(|i| match declared.get(i) {
1509 Some(&oid) if oid != 0 => oid,
1510 _ => match clauses[i] {
1513 Some(oid) => oid,
1514 None => match &fields[i] {
1515 Some(f) => infer_field_oid(db, &coll, f),
1516 None => OID_TEXT,
1517 },
1518 },
1519 })
1520 .collect()
1521}
1522
1523fn decode_param(raw: Option<&[u8]>, oid: i32, format: i16) -> Result<Option<String>, String> {
1529 let bytes = match raw {
1530 None => return Ok(None),
1531 Some(b) => b,
1532 };
1533 let quote = |s: &str| format!("'{}'", s.replace('\'', "''"));
1534
1535 if format == 0 {
1536 let s = String::from_utf8_lossy(bytes).to_string();
1537 return Ok(Some(match oid {
1538 OID_BOOL => {
1539 let t = matches!(s.as_str(), "t" | "true" | "TRUE" | "1" | "yes" | "on");
1540 if t { "TRUE".into() } else { "FALSE".into() }
1541 }
1542 OID_INT2 | OID_INT4 | OID_INT8 | OID_OID | OID_FLOAT4 | OID_FLOAT8 => {
1543 if s.parse::<f64>().is_ok() { s } else { quote(&s) }
1547 }
1548 _ => quote(&s),
1553 }));
1554 }
1555 if format != 1 {
1556 return Err(format!("unsupported parameter format code {}", format));
1557 }
1558
1559 let need = |n: usize| -> Result<(), String> {
1561 if bytes.len() == n {
1562 Ok(())
1563 } else {
1564 Err(format!(
1565 "binary parameter of type OID {} should be {} bytes, got {}",
1566 oid, n, bytes.len()
1567 ))
1568 }
1569 };
1570 Ok(Some(match oid {
1571 OID_BOOL => {
1572 need(1)?;
1573 if bytes[0] != 0 { "TRUE".into() } else { "FALSE".into() }
1574 }
1575 OID_INT2 => {
1576 need(2)?;
1577 i16::from_be_bytes([bytes[0], bytes[1]]).to_string()
1578 }
1579 OID_INT4 => {
1580 need(4)?;
1581 i32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]).to_string()
1582 }
1583 OID_OID => {
1584 need(4)?;
1585 u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]).to_string()
1586 }
1587 OID_INT8 => {
1588 need(8)?;
1589 i64::from_be_bytes(bytes[..8].try_into().unwrap()).to_string()
1590 }
1591 OID_FLOAT4 => {
1592 need(4)?;
1593 let f = f32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]);
1594 fmt_float(f as f64)
1595 }
1596 OID_FLOAT8 => {
1597 need(8)?;
1598 fmt_float(f64::from_be_bytes(bytes[..8].try_into().unwrap()))
1599 }
1600 OID_TEXT | OID_VARCHAR | OID_NAME | OID_UNKNOWN | OID_JSON | 0 => {
1601 quote(&String::from_utf8_lossy(bytes))
1602 }
1603 OID_JSONB => {
1604 let body = if bytes.first() == Some(&1) { &bytes[1..] } else { bytes };
1606 quote(&String::from_utf8_lossy(body))
1607 }
1608 other => {
1609 return Err(format!(
1610 "parameter type OID {} is not supported in binary format — \
1611 the supported set is bool, int2/int4/int8, float4/float8, \
1612 text/varchar/json/jsonb. Send it as text, or cast it in the \
1613 statement",
1614 other
1615 ))
1616 }
1617 }))
1618}
1619
1620fn fmt_float(f: f64) -> String {
1622 if f.is_nan() {
1623 "'NaN'".into()
1624 } else if f.is_infinite() {
1625 if f > 0.0 { "'Infinity'".into() } else { "'-Infinity'".into() }
1626 } else if f.fract() == 0.0 && f.abs() < 1e15 {
1627 format!("{:.0}", f)
1628 } else {
1629 f.to_string()
1630 }
1631}
1632
1633fn substitute_params(sql: &str, params: &[Option<String>]) -> Result<String, String> {
1641 let b = sql.as_bytes();
1642 let mut out = String::with_capacity(sql.len() + 16);
1643 let mut i = 0usize;
1644 let mut in_s = false;
1645 while i < b.len() {
1646 let c = b[i];
1647 if in_s {
1648 out.push(c as char);
1649 if c == b'\'' { in_s = false; }
1650 i += 1;
1651 continue;
1652 }
1653 if c == b'\'' {
1654 in_s = true;
1655 out.push('\'');
1656 i += 1;
1657 continue;
1658 }
1659 if c == b'$' && i + 1 < b.len() && b[i + 1].is_ascii_digit() {
1660 let mut j = i + 1;
1661 let mut n = 0usize;
1662 while j < b.len() && b[j].is_ascii_digit() {
1663 n = n * 10 + (b[j] - b'0') as usize;
1664 j += 1;
1665 }
1666 match params.get(n.wrapping_sub(1)) {
1667 Some(Some(lit)) => out.push_str(lit),
1668 Some(None) => out.push_str("NULL"),
1669 None => {
1670 return Err(format!(
1671 "bind message supplies {} parameter(s) but the statement uses ${}",
1672 params.len(), n
1673 ))
1674 }
1675 }
1676 i = j;
1677 continue;
1678 }
1679 out.push(c as char);
1680 i += 1;
1681 }
1682 Ok(out)
1683}
1684
1685struct Prepared {
1687 sql: String,
1688 param_oids: Vec<i32>,
1691 out_shape: Option<Option<(Vec<Col>, Vec<i32>)>>,
1699}
1700
1701fn prepared_shape<'a>(
1703 p: &'a mut Prepared,
1704 db: Option<&Arc<Db>>,
1705) -> &'a Option<(Vec<Col>, Vec<i32>)> {
1706 if p.out_shape.is_none() {
1707 p.out_shape = Some(describe_shape(&p.sql, db, p.param_oids.len()));
1708 }
1709 p.out_shape.as_ref().expect("just filled")
1710}
1711
1712struct Portal {
1714 sql: String,
1715 result: Option<PortalResult>,
1721 frozen: Option<Vec<Col>>,
1731 formats: Vec<i16>,
1733 declared: Option<(Vec<Col>, Vec<i32>)>,
1740}
1741
1742impl Portal {
1743 fn format_of(&self, i: usize) -> i16 {
1746 match self.formats.len() {
1747 0 => 0,
1748 1 => self.formats[0],
1749 _ => self.formats.get(i).copied().unwrap_or(0),
1750 }
1751 }
1752 fn shape(&self, r: &PortalResult) -> (Vec<Col>, Vec<i32>) {
1754 match &self.declared {
1755 Some((cols, oids)) if self.formats.iter().any(|f| *f == 1) => {
1756 (cols.clone(), oids.clone())
1757 }
1758 _ => {
1759 let cols = columns_for(&r.rows, &r.project);
1760 let oids = cols.iter().map(|c| oid_for(&r.rows, &c.src)).collect();
1761 (cols, oids)
1762 }
1763 }
1764 }
1765}
1766
1767struct PortalResult {
1768 rows: Vec<Value>,
1769 project: Vec<Col>,
1770 has_rows: bool,
1771 tag: String,
1772 tag_counts_rows: bool,
1773 sent: usize,
1775}
1776
1777fn parse_complete() -> Vec<u8> { Out::msg(b'1').finish() }
1778fn bind_complete() -> Vec<u8> { Out::msg(b'2').finish() }
1779fn close_complete() -> Vec<u8> { Out::msg(b'3').finish() }
1780fn no_data() -> Vec<u8> { Out::msg(b'n').finish() }
1781fn portal_suspended() -> Vec<u8> { Out::msg(b's').finish() }
1782
1783fn parameter_description(oids: &[i32]) -> Vec<u8> {
1784 let mut m = Out::msg(b't');
1785 m.i16(oids.len() as i16);
1786 for o in oids {
1787 m.i32(*o);
1788 }
1789 m.finish()
1790}
1791
1792fn take_cstr(body: &[u8], at: &mut usize) -> String {
1794 let start = *at;
1795 while *at < body.len() && body[*at] != 0 {
1796 *at += 1;
1797 }
1798 let s = String::from_utf8_lossy(&body[start..*at]).to_string();
1799 if *at < body.len() {
1800 *at += 1; }
1802 s
1803}
1804
1805fn take_i16(body: &[u8], at: &mut usize) -> Result<i16, String> {
1806 if *at + 2 > body.len() {
1807 return Err("truncated message".into());
1808 }
1809 let v = i16::from_be_bytes([body[*at], body[*at + 1]]);
1810 *at += 2;
1811 Ok(v)
1812}
1813
1814fn take_i32(body: &[u8], at: &mut usize) -> Result<i32, String> {
1815 if *at + 4 > body.len() {
1816 return Err("truncated message".into());
1817 }
1818 let v = i32::from_be_bytes([body[*at], body[*at + 1], body[*at + 2], body[*at + 3]]);
1819 *at += 4;
1820 Ok(v)
1821}
1822
1823fn sample_columns(db: Option<&Arc<Db>>, coll: &str) -> Vec<Col> {
1829 let db = match db {
1830 Some(db) => db,
1831 None => return vec![],
1832 };
1833 let rows = match crate::nql::query(db, &format!("FROM {} LIMIT 25", coll)) {
1834 Ok((rows, _)) => rows,
1835 Err(_) => return vec![],
1836 };
1837 let mut names: Vec<String> = vec![];
1838 for r in &rows {
1839 if let Value::Object(m) = r {
1840 for k in m.keys() {
1841 if !names.iter().any(|n| n == k) {
1842 names.push(k.clone());
1843 }
1844 }
1845 }
1846 }
1847 names.sort();
1848 names.iter().map(|n| Col::same(n)).collect()
1849}
1850
1851fn describe_shape(
1859 sql: &str,
1860 db: Option<&Arc<Db>>,
1861 n_params: usize,
1862) -> Option<(Vec<Col>, Vec<i32>)> {
1863 let probe = probe_sql(sql, n_params);
1864 let stmt = translate(&probe).ok()?;
1865 let coll = stmt_collection(sql);
1866
1867 let cols = match stmt {
1868 Stmt::Ok(_) => return None,
1869 Stmt::Canned { cols, .. } => cols.iter().map(|c| Col::same(c)).collect(),
1870 Stmt::Query { project, .. } => {
1871 if project.is_empty() { sample_columns(db, &coll) } else { project }
1872 }
1873 Stmt::Insert { returning, .. } | Stmt::Update { returning, .. } | Stmt::Delete { returning, .. } => {
1874 if !wants_returning(sql) {
1875 return None;
1876 }
1877 if returning.is_empty() { sample_columns(db, &coll) } else { returning }
1878 }
1879 };
1880 if cols.is_empty() {
1881 return None;
1885 }
1886 let oids = cols
1887 .iter()
1888 .map(|c| {
1889 aggregate_oid(&c.src, db, &coll)
1890 .unwrap_or_else(|| infer_field_oid(db, &coll, &c.src))
1891 })
1892 .collect();
1893 Some((cols, oids))
1894}
1895
1896fn probe_sql(sql: &str, n_params: usize) -> String {
1904 let stub: Vec<Option<String>> = vec![Some("0".to_string()); n_params];
1905 substitute_params(sql, &stub).unwrap_or_else(|_| sql.to_string())
1906}
1907
1908fn ensure_executed(
1910 portal: &mut Portal,
1911 db_name: &str,
1912 db: Option<&Arc<Db>>,
1913 read_only: bool,
1914) -> Result<(), Vec<u8>> {
1915 if portal.result.is_some() {
1916 return Ok(());
1917 }
1918 let ex = execute_stmt(&portal.sql, db_name, db, read_only)?;
1919 let project = if let Some(f) = &portal.frozen {
1922 f.clone()
1923 } else {
1924 let p = if ex.project.is_empty() {
1925 columns_for(&ex.rows, &[])
1926 } else {
1927 ex.project.clone()
1928 };
1929 portal.frozen = Some(p.clone());
1930 p
1931 };
1932 portal.result = Some(PortalResult {
1933 rows: ex.rows,
1934 project,
1935 has_rows: ex.has_rows,
1936 tag: ex.tag,
1937 tag_counts_rows: ex.tag_counts_rows,
1938 sent: 0,
1939 });
1940 Ok(())
1941}
1942
1943async fn read_exact(sock: &mut TcpStream, n: usize) -> std::io::Result<Vec<u8>> {
1946 let mut buf = vec![0u8; n];
1947 sock.read_exact(&mut buf).await?;
1948 Ok(buf)
1949}
1950
1951async fn read_i32(sock: &mut TcpStream) -> std::io::Result<i32> {
1952 let b = read_exact(sock, 4).await?;
1953 Ok(i32::from_be_bytes([b[0], b[1], b[2], b[3]]))
1954}
1955
1956fn parse_startup_params(body: &[u8]) -> HashMap<String, String> {
1957 let mut out = HashMap::new();
1958 let mut parts = body.split(|b| *b == 0).map(|s| String::from_utf8_lossy(s).to_string());
1959 while let (Some(k), Some(v)) = (parts.next(), parts.next()) {
1960 if k.is_empty() {
1961 break;
1962 }
1963 out.insert(k, v);
1964 }
1965 out
1966}
1967
1968async fn handle(mut sock: TcpStream, resolver: Arc<dyn DbResolver>, read_only: bool) -> std::io::Result<()> {
1970 let params = loop {
1972 let len = read_i32(&mut sock).await?;
1973 if len < 8 || len > 1 << 20 {
1974 return Ok(()); }
1976 let code = read_i32(&mut sock).await?;
1977 let body = read_exact(&mut sock, (len - 8) as usize).await?;
1978 match code {
1979 SSL_REQUEST | GSS_REQUEST => {
1980 sock.write_all(b"N").await?;
1982 continue;
1983 }
1984 CANCEL_REQUEST => return Ok(()), PROTO_V3 => break parse_startup_params(&body),
1986 other => {
1987 let major = other >> 16;
1988 sock.write_all(&err_msg(
1989 "0A000",
1990 &format!("unsupported frontend protocol {}.{} — this endpoint speaks 3.0",
1991 major, other & 0xffff),
1992 )).await?;
1993 return Ok(());
1994 }
1995 }
1996 };
1997
1998 let db_name = params.get("database").cloned().unwrap_or_default();
1999
2000 let resolved: Option<Arc<Db>> = {
2006 let r = Arc::clone(&resolver);
2007 let name = db_name.clone();
2008 tokio::task::spawn_blocking(move || r.resolve(&name))
2009 .await
2010 .unwrap_or(None)
2011 };
2012
2013 if let Some(expected) = resolver.token() {
2015 let mut m = Out::msg(b'R');
2017 m.i32(3);
2018 sock.write_all(&m.finish()).await?;
2019
2020 let tag = read_exact(&mut sock, 1).await?;
2021 if tag[0] != b'p' {
2022 sock.write_all(&err_msg("28000", "expected a password message")).await?;
2023 return Ok(());
2024 }
2025 let len = read_i32(&mut sock).await?;
2026 if len < 4 || len > 1 << 16 {
2027 return Ok(());
2028 }
2029 let body = read_exact(&mut sock, (len - 4) as usize).await?;
2030 let supplied = String::from_utf8_lossy(&body).trim_end_matches('\0').to_string();
2031 let ok = supplied.len() == expected.len()
2033 && supplied.bytes().zip(expected.bytes()).fold(0u8, |a, (x, y)| a | (x ^ y)) == 0;
2034 if !ok {
2035 sock.write_all(&err_msg("28P01", "password authentication failed")).await?;
2036 return Ok(());
2037 }
2038 }
2039
2040 let mut m = Out::msg(b'R');
2041 m.i32(0); sock.write_all(&m.finish()).await?;
2043
2044 for (k, v) in [
2045 ("server_version", SERVER_VERSION),
2046 ("server_encoding", "UTF8"),
2047 ("client_encoding", "UTF8"),
2048 ("DateStyle", "ISO, MDY"),
2049 ("integer_datetimes", "on"),
2050 ("standard_conforming_strings", "on"),
2051 ("application_name", "nedbd"),
2052 ] {
2053 let mut p = Out::msg(b'S');
2054 p.cstr(k);
2055 p.cstr(v);
2056 sock.write_all(&p.finish()).await?;
2057 }
2058 let mut k = Out::msg(b'K');
2059 k.i32(std::process::id() as i32);
2060 k.i32(0);
2061 sock.write_all(&k.finish()).await?;
2062 sock.write_all(&ready()).await?;
2063
2064 let mut prepared: HashMap<String, Prepared> = HashMap::new();
2070 let mut portals: HashMap<String, Portal> = HashMap::new();
2071 let mut failed = false;
2075
2076 loop {
2077 let mut tag = [0u8; 1];
2078 if sock.read_exact(&mut tag).await.is_err() {
2079 return Ok(()); }
2081 let len = read_i32(&mut sock).await?;
2082 if len < 4 || len > 64 << 20 {
2083 return Ok(());
2084 }
2085 let body = read_exact(&mut sock, (len - 4) as usize).await?;
2086
2087 if failed && tag[0] != b'S' && tag[0] != b'X' {
2089 continue;
2090 }
2091
2092 match tag[0] {
2093 b'X' => return Ok(()), b'Q' => {
2096 let sql = String::from_utf8_lossy(&body).trim_end_matches('\0').to_string();
2097 let out = run_simple_query(&sql, &db_name, resolved.as_ref(), read_only);
2098 sock.write_all(&out).await?;
2099 sock.write_all(&ready()).await?;
2100 portals.remove("");
2102 }
2103
2104 b'P' => {
2106 let mut at = 0usize;
2107 let name = take_cstr(&body, &mut at);
2108 let sql = take_cstr(&body, &mut at);
2109 let n = take_i16(&body, &mut at).unwrap_or(0).max(0) as usize;
2110 let mut declared = Vec::with_capacity(n);
2111 let mut bad = false;
2112 for _ in 0..n {
2113 match take_i32(&body, &mut at) {
2114 Ok(o) => declared.push(o),
2115 Err(_) => { bad = true; break; }
2116 }
2117 }
2118 if bad {
2119 sock.write_all(&err_msg("08P01", "malformed Parse message")).await?;
2120 failed = true;
2121 continue;
2122 }
2123 if let Err(why) = translate(&probe_sql(&sql, param_count(&sql))) {
2127 sock.write_all(&err_msg("0A000", &why)).await?;
2128 failed = true;
2129 continue;
2130 }
2131 let param_oids = infer_param_oids(&sql, &declared, resolved.as_ref());
2132 prepared.insert(name, Prepared { sql, param_oids, out_shape: None });
2133 sock.write_all(&parse_complete()).await?;
2134 }
2135
2136 b'B' => {
2138 let mut at = 0usize;
2139 let portal_name = take_cstr(&body, &mut at);
2140 let stmt_name = take_cstr(&body, &mut at);
2141 if !prepared.contains_key(&stmt_name) {
2142 sock.write_all(&err_msg("26000", &format!(
2143 "prepared statement {:?} does not exist", stmt_name))).await?;
2144 failed = true;
2145 continue;
2146 }
2147 let p = &prepared[&stmt_name];
2148 let mut want_formats: Vec<i16> = vec![];
2149 let res: Result<String, String> = (|| {
2150 let nfmt = take_i16(&body, &mut at)? .max(0) as usize;
2151 let mut fmts = Vec::with_capacity(nfmt);
2152 for _ in 0..nfmt {
2153 fmts.push(take_i16(&body, &mut at)?);
2154 }
2155 let nparam = take_i16(&body, &mut at)?.max(0) as usize;
2156 let mut vals: Vec<Option<String>> = Vec::with_capacity(nparam);
2157 for i in 0..nparam {
2158 let l = take_i32(&body, &mut at)?;
2159 let raw: Option<Vec<u8>> = if l < 0 {
2160 None
2161 } else {
2162 let l = l as usize;
2163 if at + l > body.len() {
2164 return Err("truncated Bind parameter".into());
2165 }
2166 let v = body[at..at + l].to_vec();
2167 at += l;
2168 Some(v)
2169 };
2170 let f = match fmts.len() {
2173 0 => 0,
2174 1 => fmts[0],
2175 _ => *fmts.get(i).unwrap_or(&0),
2176 };
2177 let oid = *p.param_oids.get(i).unwrap_or(&OID_TEXT);
2178 vals.push(decode_param(raw.as_deref(), oid, f)?);
2179 }
2180 let nres = take_i16(&body, &mut at)?.max(0) as usize;
2185 for _ in 0..nres {
2186 let f = take_i16(&body, &mut at)?;
2187 if f != 0 && f != 1 {
2188 return Err(format!("unknown result format code {}", f));
2189 }
2190 want_formats.push(f);
2191 }
2192 substitute_params(&p.sql, &vals)
2193 })();
2194 match res {
2195 Ok(sql) => {
2196 let declared = if want_formats.iter().any(|f| *f == 1) {
2199 let p = prepared.get_mut(&stmt_name).expect("checked above");
2200 prepared_shape(p, resolved.as_ref()).clone()
2201 } else {
2202 None
2203 };
2204 portals.insert(portal_name, Portal {
2205 sql, result: None, frozen: None,
2206 formats: want_formats, declared,
2207 });
2208 sock.write_all(&bind_complete()).await?;
2209 }
2210 Err(why) => {
2211 sock.write_all(&err_msg("08P01", &why)).await?;
2212 failed = true;
2213 }
2214 }
2215 }
2216
2217 b'D' => {
2219 let kind = body.first().copied().unwrap_or(b'S');
2220 let mut at = 1usize;
2221 let name = take_cstr(&body, &mut at);
2222 if kind == b'S' {
2223 if !prepared.contains_key(&name) {
2224 sock.write_all(&err_msg("26000", &format!(
2225 "prepared statement {:?} does not exist", name))).await?;
2226 failed = true;
2227 continue;
2228 }
2229 let p = prepared.get_mut(&name).expect("checked above");
2230 let oids = p.param_oids.clone();
2231 sock.write_all(¶meter_description(&oids)).await?;
2234 let out = match prepared_shape(p, resolved.as_ref()) {
2238 Some((cols, col_oids)) => row_description(cols, col_oids),
2239 None => no_data(),
2240 };
2241 sock.write_all(&out).await?;
2242 } else {
2243 let portal = match portals.get_mut(&name) {
2244 Some(p) => p,
2245 None => {
2246 sock.write_all(&err_msg("34000", &format!(
2247 "portal {:?} does not exist", name))).await?;
2248 failed = true;
2249 continue;
2250 }
2251 };
2252 match ensure_executed(portal, &db_name, resolved.as_ref(), read_only) {
2257 Err(encoded) => {
2258 sock.write_all(&encoded).await?;
2259 failed = true;
2260 }
2261 Ok(()) => {
2262 let r = portal.result.as_ref().expect("just executed");
2263 if !r.has_rows {
2264 sock.write_all(&no_data()).await?;
2265 } else {
2266 let (cols, oids) = portal.shape(r);
2267 let fmts: Vec<i16> =
2268 (0..cols.len()).map(|i| portal.format_of(i)).collect();
2269 sock.write_all(&row_description_fmt(&cols, &oids, &fmts)).await?;
2270 }
2271 }
2272 }
2273 }
2274 }
2275
2276 b'E' => {
2278 let mut at = 0usize;
2279 let name = take_cstr(&body, &mut at);
2280 let max_rows = take_i32(&body, &mut at).unwrap_or(0);
2281 let portal = match portals.get_mut(&name) {
2282 Some(p) => p,
2283 None => {
2284 sock.write_all(&err_msg("34000", &format!(
2285 "portal {:?} does not exist", name))).await?;
2286 failed = true;
2287 continue;
2288 }
2289 };
2290 if let Err(encoded) = ensure_executed(portal, &db_name, resolved.as_ref(), read_only) {
2291 sock.write_all(&encoded).await?;
2292 failed = true;
2293 continue;
2294 }
2295 let r = portal.result.as_ref().expect("just executed");
2296 if !r.has_rows {
2297 let tag = r.tag.clone();
2298 sock.write_all(&command_complete(&tag)).await?;
2299 continue;
2300 }
2301 let (cols, oids) = portal.shape(r);
2302 let limit = if max_rows > 0 {
2303 (r.sent + max_rows as usize).min(r.rows.len())
2304 } else {
2305 r.rows.len()
2306 };
2307 let mut encoded: Vec<Vec<u8>> = Vec::with_capacity(limit - r.sent);
2312 let mut fail: Option<String> = None;
2313 for row in &r.rows[r.sent..limit] {
2314 let mut vals: Vec<Option<Vec<u8>>> = Vec::with_capacity(cols.len());
2315 for (i, c) in cols.iter().enumerate() {
2316 let v = row.get(&c.src);
2317 let got = if portal.format_of(i) == 1 {
2318 cell_binary(v, oids.get(i).copied().unwrap_or(OID_TEXT))
2319 .map_err(|e| format!("column {:?}: {}", c.out, e))
2320 } else {
2321 Ok(cell(v).map(|s| s.into_bytes()))
2322 };
2323 match got {
2324 Ok(b) => vals.push(b),
2325 Err(e) => { fail = Some(e); break; }
2326 }
2327 }
2328 if fail.is_some() {
2329 break;
2330 }
2331 encoded.push(data_row_bytes(&vals));
2332 }
2333 if let Some(why) = fail {
2334 sock.write_all(&err_msg("22P03", &why)).await?;
2335 failed = true;
2336 continue;
2337 }
2338 let mut out = vec![];
2339 for e in &encoded {
2340 out.extend_from_slice(e);
2341 }
2342 let r = portal.result.as_mut().expect("just executed");
2343 r.sent = limit;
2344 if max_rows > 0 && r.sent < r.rows.len() {
2348 out.extend_from_slice(&portal_suspended());
2349 } else {
2350 let tag = if r.tag_counts_rows {
2351 format!("{} {}", r.tag, r.sent)
2352 } else {
2353 r.tag.clone()
2354 };
2355 out.extend_from_slice(&command_complete(&tag));
2356 }
2357 sock.write_all(&out).await?;
2358 }
2359
2360 b'C' => {
2362 let kind = body.first().copied().unwrap_or(b'S');
2363 let mut at = 1usize;
2364 let name = take_cstr(&body, &mut at);
2365 if kind == b'S' {
2366 prepared.remove(&name);
2367 } else {
2368 portals.remove(&name);
2369 }
2370 sock.write_all(&close_complete()).await?;
2373 }
2374
2375 b'H' => {}
2379
2380 b'S' => {
2381 failed = false;
2382 sock.write_all(&ready()).await?;
2383 }
2384
2385 other => {
2386 sock.write_all(&err_msg(
2387 "08P01",
2388 &format!("unexpected frontend message {:?}", other as char),
2389 )).await?;
2390 failed = true;
2391 }
2392 }
2393 }
2394}
2395
2396const READ_ONLY_MSG: &str =
2397 "this endpoint is running read-only (NEDBD_PG_READ_ONLY=1). Writes are \
2398 implemented but disabled on this server — unset the flag to allow them.";
2399
2400fn no_db(db_name: &str) -> Vec<u8> {
2401 err_msg("3D000", &format!(
2402 "database {:?} is not open on this server — create it first \
2403 (POST /v1/databases), or connect with -d <name>", db_name))
2404}
2405
2406fn try_catalog_select(
2419 sql: &str,
2420 db: Option<&Arc<Db>>,
2421) -> Result<Option<(Executed, crate::sqlplan::Plan)>, Vec<u8>> {
2422 let sel = match crate::sqlselect::parse(sql) {
2423 Ok(sel) => sel,
2424 Err(why) => {
2425 if mentions_catalog(sql) {
2434 return Err(err_msg("0A000", &format!(
2435 "this catalogue query uses SQL this endpoint does not \
2436 implement: {}", why)));
2437 }
2438 return Ok(None);
2439 }
2440 };
2441
2442 let touched: Vec<String> = sel.base_relations();
2448 let catalog_name = |n: &str| -> String {
2449 let joined: Vec<&str> = n.split('.').collect();
2453 if joined.len() >= 2 && joined[joined.len() - 2] == "information_schema" {
2454 format!("information_schema.{}", joined[joined.len() - 1])
2455 } else {
2456 joined[joined.len() - 1].to_string()
2457 }
2458 };
2459 if !touched.iter().any(|t| crate::pgcatalog::is_catalog(&catalog_name(t))) {
2460 return Ok(None);
2461 }
2462
2463 let resolve = |name: &str| -> anyhow::Result<Option<Box<dyn crate::sqlselect::Relation>>> {
2464 let cname = catalog_name(name);
2465 if let Some(rows) = crate::pgcatalog::rows(&cname, db) {
2466 return Ok(Some(crate::sqlselect::from_vec(rows)));
2470 }
2471 match db {
2480 Some(db) => match crate::nql::query(db, &format!("FROM {}", cname)) {
2481 Ok((rows, _)) => Ok(Some(crate::sqlselect::from_vec(rows))),
2482 Err(_) => Ok(None),
2483 },
2484 None => Ok(None),
2485 }
2486 };
2487
2488 let (cols, rows, plan) = crate::sqlselect::execute_explain(
2489 &sel,
2490 &resolve,
2491 crate::sqljoin::JoinExec::Auto,
2492 )
2493 .map_err(|e| err_msg("42601", &e.to_string()))?;
2494
2495 Ok(Some((
2496 Executed {
2497 rows,
2498 project: cols
2502 .iter()
2503 .map(|c| Col::renamed(&c.key, &c.name))
2504 .collect(),
2505 has_rows: true,
2506 tag: "SELECT".into(),
2507 tag_counts_rows: true,
2508 },
2509 plan,
2510 )))
2511}
2512
2513fn strip_explain(sql: &str) -> Option<&str> {
2520 let t = sql.trim().trim_end_matches(';').trim();
2521 let mut rest = t.strip_prefix("EXPLAIN").or_else(|| t.strip_prefix("explain"))?;
2522 if !rest.starts_with(char::is_whitespace) {
2524 return None;
2525 }
2526 rest = rest.trim_start();
2527 loop {
2528 let low = rest.to_lowercase();
2529 if let Some(r) = low.strip_prefix("analyze").or_else(|| low.strip_prefix("analyse")) {
2530 if r.starts_with(char::is_whitespace) || r.is_empty() {
2531 rest = rest[rest.len() - r.len()..].trim_start();
2532 continue;
2533 }
2534 }
2535 if let Some(r) = low.strip_prefix("verbose") {
2536 if r.starts_with(char::is_whitespace) || r.is_empty() {
2537 rest = rest[rest.len() - r.len()..].trim_start();
2538 continue;
2539 }
2540 }
2541 break;
2542 }
2543 Some(rest)
2544}
2545
2546fn plan_result(lines: Vec<String>) -> Executed {
2549 Executed {
2550 rows: lines
2551 .into_iter()
2552 .map(|l| serde_json::json!({ "QUERY PLAN": l }))
2553 .collect(),
2554 project: vec![Col::same("QUERY PLAN")],
2555 has_rows: true,
2556 tag: "EXPLAIN".into(),
2557 tag_counts_rows: false,
2558 }
2559}
2560
2561fn mentions_catalog(sql: &str) -> bool {
2568 let low = sql.to_lowercase();
2569 low.contains("pg_catalog.")
2570 || low.contains("information_schema.")
2571 || low.contains("from pg_")
2572 || low.contains("join pg_")
2573}
2574
2575fn catalog_target(nql: &str) -> Option<String> {
2580 let coll = crate::nql::parse(nql).ok()?.coll;
2581 if crate::pgcatalog::is_catalog(&coll) {
2582 Some(coll)
2583 } else {
2584 None
2585 }
2586}
2587
2588fn wants_returning(sql: &str) -> bool {
2592 find_kw(&sql.to_uppercase(), "RETURNING").is_some()
2593}
2594
2595fn next_row_id() -> String {
2597 use std::sync::atomic::{AtomicU64, Ordering};
2598 static N: AtomicU64 = AtomicU64::new(0);
2599 let n = N.fetch_add(1, Ordering::Relaxed);
2600 let ts = std::time::SystemTime::now()
2601 .duration_since(std::time::UNIX_EPOCH)
2602 .map(|d| d.as_micros())
2603 .unwrap_or(0);
2604 format!("r{}{}", ts, n)
2605}
2606
2607pub struct Executed {
2615 pub rows: Vec<Value>,
2617 pub project: Vec<Col>,
2619 pub has_rows: bool,
2623 pub tag: String,
2626 pub tag_counts_rows: bool,
2628}
2629
2630impl Executed {
2631 fn nothing(tag: &str) -> Self {
2632 Executed { rows: vec![], project: vec![], has_rows: false, tag: tag.to_string(), tag_counts_rows: false }
2633 }
2634 fn tag_for(&self, sent: usize) -> String {
2636 if self.tag_counts_rows { format!("{} {}", self.tag, sent) } else { self.tag.clone() }
2637 }
2638}
2639
2640fn execute_stmt(
2646 stmt_sql: &str,
2647 db_name: &str,
2648 db: Option<&Arc<Db>>,
2649 read_only: bool,
2650) -> Result<Executed, Vec<u8>> {
2651 if let Some(inner) = strip_explain(stmt_sql) {
2660 if let Some((_, plan)) = try_catalog_select(inner, db)? {
2661 return Ok(plan_result(plan.render()));
2662 }
2663 let mut lines = vec![];
2664 match translate(inner) {
2665 Ok(_) => {
2666 lines.push(
2667 "NQL path — this statement is translated to NQL and \
2668 executed by the storage engine, not by the SQL evaluator."
2669 .to_string(),
2670 );
2671 lines.push(
2672 "No plan is reported, because the SQL evaluator is not \
2673 what runs it. Reporting one would describe a pipeline \
2674 that never executed."
2675 .to_string(),
2676 );
2677 lines.push(
2678 "The SQL evaluator (joins, CASE, scalar functions, a \
2679 hash-join planner) currently serves catalogue queries."
2680 .to_string(),
2681 );
2682 }
2683 Err(why) => lines.push(format!("cannot be executed: {why}")),
2684 }
2685 return Ok(plan_result(lines));
2686 }
2687
2688 if let Some((done, _plan)) = try_catalog_select(stmt_sql, db)? {
2689 return Ok(done);
2690 }
2691
2692 let stmt = translate(stmt_sql).map_err(|why| err_msg("0A000", &why))?;
2693
2694 macro_rules! need_db {
2697 () => {
2698 match db {
2699 Some(db) => db,
2700 None => return Err(no_db(db_name)),
2701 }
2702 };
2703 }
2704 macro_rules! need_write {
2705 () => {
2706 if read_only {
2707 return Err(err_msg("25006", READ_ONLY_MSG));
2708 }
2709 };
2710 }
2711
2712 match stmt {
2713 Stmt::Ok(tag) => Ok(Executed::nothing(if tag.is_empty() { "SELECT 0" } else { tag })),
2714
2715 Stmt::Canned { cols, row } => {
2716 let mut obj = serde_json::Map::new();
2719 for (c, v) in cols.iter().zip(row.iter()) {
2720 obj.insert(c.clone(), Value::String(v.clone()));
2721 }
2722 Ok(Executed {
2723 rows: vec![Value::Object(obj)],
2724 project: cols.iter().map(|c| Col::same(c)).collect(),
2725 has_rows: true,
2726 tag: "SELECT".into(),
2727 tag_counts_rows: true,
2728 })
2729 }
2730
2731 Stmt::Query { nql, project } => {
2732 if let Some(coll) = catalog_target(&nql) {
2742 let rows = crate::pgcatalog::rows(&coll, db)
2743 .expect("catalog_target only returns names pgcatalog serves");
2744 let rows = crate::nql::query_rows(rows, &nql)
2745 .map_err(|e| err_msg("42601", &e.to_string()))?;
2746 return Ok(Executed {
2747 rows, project, has_rows: true,
2748 tag: "SELECT".into(), tag_counts_rows: true,
2749 });
2750 }
2751 let db = need_db!();
2752 let (rows, _) = crate::nql::query(db, &nql).map_err(|e| {
2753 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
2754 })?;
2755 Ok(Executed { rows, project, has_rows: true, tag: "SELECT".into(), tag_counts_rows: true })
2756 }
2757
2758 Stmt::Insert { coll, rows, returning } => {
2759 let db = need_db!();
2760 need_write!();
2761 let mut written: Vec<Value> = vec![];
2762 for (i, r) in rows.iter().enumerate() {
2763 let id = match &r.id {
2767 Some(id) => id.clone(),
2768 None => format!("{}-{}", next_row_id(), i),
2769 };
2770 let node = db
2771 .put(&coll, &id, Value::Object(r.doc.clone()),
2772 r.caused_by.clone(), r.valid_from.clone(), r.valid_to.clone())
2773 .map_err(|e| err_msg("XX000", &format!("INSERT failed: {}", e)))?;
2774 written.push(crate::nql::node_to_json(&node));
2775 }
2776 let n = written.len();
2777 let has_rows = wants_returning(stmt_sql);
2778 Ok(Executed {
2779 rows: if has_rows { written } else { vec![] },
2780 project: returning,
2781 has_rows,
2782 tag: format!("INSERT 0 {}", n),
2784 tag_counts_rows: false,
2785 })
2786 }
2787
2788 Stmt::Update { coll, set, nql, returning } => {
2789 let db = need_db!();
2790 need_write!();
2791 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
2794 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
2795 })?;
2796 let mut written: Vec<Value> = vec![];
2797 for row in &matched {
2798 let id = match row.get("_id").and_then(|v| v.as_str()) {
2799 Some(id) => id.to_string(),
2800 None => continue,
2801 };
2802 let mut doc = match db.get(&coll, &id) {
2806 Some(n) => match n.data {
2807 Value::Object(m) => m,
2808 _ => serde_json::Map::new(),
2809 },
2810 None => continue,
2811 };
2812 for (k, v) in &set {
2813 doc.insert(k.clone(), v.clone());
2814 }
2815 let node = db
2818 .put(&coll, &id, Value::Object(doc), vec![], None, None)
2819 .map_err(|e| err_msg("XX000", &format!("UPDATE failed: {}", e)))?;
2820 written.push(crate::nql::node_to_json(&node));
2821 }
2822 let n = written.len();
2823 let has_rows = wants_returning(stmt_sql);
2824 Ok(Executed {
2825 rows: if has_rows { written } else { vec![] },
2826 project: returning,
2827 has_rows,
2828 tag: format!("UPDATE {}", n),
2829 tag_counts_rows: false,
2830 })
2831 }
2832
2833 Stmt::Delete { coll, nql, returning } => {
2834 let db = need_db!();
2835 need_write!();
2836 let (matched, _) = crate::nql::query(db, &nql).map_err(|e| {
2837 err_msg("42601", &format!("{} (translated to NQL: {})", e, nql))
2838 })?;
2839 let returned = matched.clone();
2842 let mut n = 0usize;
2843 for row in &matched {
2844 if let Some(id) = row.get("_id").and_then(|v| v.as_str()) {
2845 match db.delete(&coll, id) {
2846 Ok(true) => n += 1,
2847 Ok(false) => {}
2848 Err(e) => return Err(err_msg("XX000", &format!("DELETE failed: {}", e))),
2849 }
2850 }
2851 }
2852 let has_rows = wants_returning(stmt_sql);
2853 Ok(Executed {
2854 rows: if has_rows { returned } else { vec![] },
2855 project: returning,
2856 has_rows,
2857 tag: format!("DELETE {}", n),
2858 tag_counts_rows: false,
2859 })
2860 }
2861 }
2862}
2863
2864fn run_simple_query(sql: &str, db_name: &str, db: Option<&Arc<Db>>, read_only: bool) -> Vec<u8> {
2866 let mut out = vec![];
2867 let statements = split_statements(sql);
2868 if statements.is_empty() {
2869 return Out::msg(b'I').finish();
2871 }
2872 for stmt_sql in statements {
2873 match execute_stmt(&stmt_sql, db_name, db, read_only) {
2874 Err(encoded) => {
2876 out.extend_from_slice(&encoded);
2877 return out;
2878 }
2879 Ok(ex) => {
2880 if ex.has_rows {
2881 out.extend_from_slice(&encode_rows(&ex.rows, &ex.project));
2882 }
2883 out.extend_from_slice(&command_complete(&ex.tag_for(ex.rows.len())));
2884 }
2885 }
2886 }
2887 out
2888}
2889
2890fn split_statements(sql: &str) -> Vec<String> {
2892 let mut out = vec![];
2893 let mut cur = String::new();
2894 let mut in_s = false;
2895 for c in sql.chars() {
2896 match c {
2897 '\'' => { in_s = !in_s; cur.push(c); }
2898 ';' if !in_s => {
2899 if !cur.trim().is_empty() { out.push(cur.clone()); }
2900 cur.clear();
2901 }
2902 _ => cur.push(c),
2903 }
2904 }
2905 if !cur.trim().is_empty() {
2906 out.push(cur);
2907 }
2908 out
2909}
2910
2911pub async fn run(host: &str, port: u16, resolver: Arc<dyn DbResolver>) -> anyhow::Result<()> {
2913 let read_only = std::env::var("NEDBD_PG_READ_ONLY")
2917 .map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
2918 .unwrap_or(false);
2919 let listener = TcpListener::bind((host, port)).await?;
2920 println!(" pgwire postgres endpoint on {}:{} — psql / DBeaver / psycopg ({})",
2921 host, port,
2922 if read_only { "SELECT only — read-only mode" } else { "SELECT + INSERT/UPDATE/DELETE" });
2923 loop {
2924 let (sock, _peer) = match listener.accept().await {
2925 Ok(v) => v,
2926 Err(e) => {
2927 eprintln!(" [pgwire] accept failed: {}", e);
2928 continue;
2929 }
2930 };
2931 let r = Arc::clone(&resolver);
2932 tokio::spawn(async move {
2933 let _ = sock.set_nodelay(true);
2934 if let Err(e) = handle(sock, r, read_only).await {
2935 if e.kind() != std::io::ErrorKind::UnexpectedEof
2937 && e.kind() != std::io::ErrorKind::ConnectionReset
2938 {
2939 eprintln!(" [pgwire] connection error: {}", e);
2940 }
2941 }
2942 });
2943 }
2944}
2945
2946#[cfg(test)]
2949mod explain_tests {
2950 use super::*;
2951
2952 #[test]
2953 fn a_bare_explain_is_stripped() {
2954 assert_eq!(strip_explain("EXPLAIN SELECT 1"), Some("SELECT 1"));
2955 assert_eq!(strip_explain("explain select 1"), Some("select 1"));
2956 assert_eq!(strip_explain(" EXPLAIN SELECT 1 ; "), Some("SELECT 1"));
2957 }
2958
2959 #[test]
2960 fn analyze_and_verbose_are_accepted_and_ignored() {
2961 assert_eq!(strip_explain("EXPLAIN ANALYZE SELECT 1"), Some("SELECT 1"));
2965 assert_eq!(strip_explain("EXPLAIN ANALYSE SELECT 1"), Some("SELECT 1"));
2966 assert_eq!(strip_explain("EXPLAIN VERBOSE SELECT 1"), Some("SELECT 1"));
2967 assert_eq!(strip_explain("EXPLAIN ANALYZE VERBOSE SELECT 1"), Some("SELECT 1"));
2968 assert_eq!(strip_explain("explain analyze verbose select 1"), Some("select 1"));
2969 }
2970
2971 #[test]
2972 fn a_word_merely_starting_with_explain_is_not_a_keyword() {
2973 assert_eq!(strip_explain("EXPLAINED SELECT 1"), None);
2974 assert_eq!(strip_explain("SELECT 1"), None);
2975 assert_eq!(strip_explain("SELECT explain FROM t"), None);
2976 }
2977
2978 #[test]
2979 fn a_column_named_analyze_is_not_eaten() {
2980 assert_eq!(strip_explain("EXPLAIN analyzed_view"), Some("analyzed_view"));
2983 }
2984
2985 #[test]
2986 fn the_plan_result_has_postgres_shape() {
2987 let e = plan_result(vec!["Seq Scan on t".into(), "note".into()]);
2988 assert_eq!(e.project.len(), 1);
2989 assert_eq!(e.project[0].out, "QUERY PLAN");
2990 assert_eq!(e.rows.len(), 2);
2991 assert_eq!(e.rows[0]["QUERY PLAN"], "Seq Scan on t");
2992 assert_eq!(e.tag, "EXPLAIN");
2993 assert!(!e.tag_counts_rows);
2995 }
2996}
2997
2998#[cfg(test)]
2999mod tests {
3000 use super::*;
3001 use serde_json::json;
3002
3003 fn q(sql: &str) -> String {
3004 match translate(sql) {
3005 Ok(Stmt::Query { nql, .. }) => nql,
3006 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3007 }
3008 }
3009 fn proj(sql: &str) -> Vec<String> {
3011 match translate(sql) {
3012 Ok(Stmt::Query { project, .. }) => project.iter().map(|c| c.out.clone()).collect(),
3013 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3014 }
3015 }
3016 fn proj_pairs(sql: &str) -> Vec<(String, String)> {
3018 match translate(sql) {
3019 Ok(Stmt::Query { project, .. }) =>
3020 project.iter().map(|c| (c.src.clone(), c.out.clone())).collect(),
3021 other => panic!("expected a query for {:?}, got {:?}", sql, other),
3022 }
3023 }
3024 fn names(cols: &[Col]) -> Vec<String> { cols.iter().map(|c| c.out.clone()).collect() }
3025
3026 #[test]
3027 fn select_star_becomes_bare_from() {
3028 assert_eq!(q("SELECT * FROM orders"), "FROM orders");
3029 assert_eq!(q("select * from orders;"), "FROM orders");
3030 assert_eq!(proj("SELECT * FROM orders"), Vec::<String>::new());
3031 }
3032
3033 #[test]
3034 fn a_column_list_becomes_a_projection_not_a_clause() {
3035 assert_eq!(q("SELECT status, total FROM orders"), "FROM orders");
3038 assert_eq!(proj("SELECT status, total FROM orders"), vec!["status", "total"]);
3039 }
3040
3041 #[test]
3042 fn aliases_and_qualified_names_reduce_to_the_field() {
3043 assert_eq!(proj("SELECT o.status AS s, o.total total FROM orders o"),
3044 vec!["status", "total"]);
3045 assert_eq!(q("SELECT * FROM public.orders"), "FROM orders");
3046 assert_eq!(q("SELECT * FROM \"orders\""), "FROM orders");
3047 }
3048
3049 #[test]
3050 fn where_clauses_pass_through_with_sql_literals_rewritten() {
3051 assert_eq!(q("SELECT * FROM orders WHERE status = 'paid'"),
3052 r#"FROM orders WHERE status = "paid""#);
3053 assert_eq!(q("SELECT * FROM orders WHERE status <> 'paid'"),
3054 r#"FROM orders WHERE status != "paid""#);
3055 assert_eq!(q("SELECT * FROM orders WHERE status IN ('paid','open')"),
3056 r#"FROM orders WHERE status IN ("paid","open")"#);
3057 }
3058
3059 #[test]
3062 fn a_doubled_sql_quote_is_one_literal_character() {
3063 assert_eq!(q("SELECT * FROM t WHERE name = 'it''s'"),
3064 r#"FROM t WHERE name = "it's""#);
3065 }
3066
3067 #[test]
3070 fn a_double_quote_inside_a_sql_literal_is_escaped_for_nql() {
3071 assert_eq!(q(r#"SELECT * FROM t WHERE name = 'say "hi"'"#),
3072 r#"FROM t WHERE name = "say \"hi\"""#);
3073 }
3074
3075 #[test]
3076 fn the_shared_clauses_are_handed_to_nql_unchanged() {
3077 assert_eq!(q("SELECT * FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5"),
3078 "FROM orders ORDER BY total DESC LIMIT 10 OFFSET 5");
3079 assert_eq!(q("SELECT * FROM orders GROUP BY region"), "FROM orders GROUP BY region");
3080 assert_eq!(q("SELECT * FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC"),
3081 "FROM o WHERE total BETWEEN 1 AND 9 ORDER BY a, b DESC");
3082 }
3083
3084 #[test]
3091 fn an_aggregate_is_one_column_named_as_sql_names_it() {
3092 assert_eq!(proj_pairs("SELECT COUNT(*) FROM orders"),
3093 vec![("count".to_string(), "count".to_string())]);
3094 assert_eq!(proj_pairs("SELECT SUM(total) FROM orders"),
3095 vec![("sum_total".to_string(), "sum".to_string())]);
3096 assert_eq!(proj_pairs("SELECT avg(total) FROM orders"),
3097 vec![("avg_total".to_string(), "avg".to_string())]);
3098 assert_eq!(proj_pairs("SELECT MIN(total) FROM orders"),
3099 vec![("min_total".to_string(), "min".to_string())]);
3100 let rows = vec![json!({"count": 4, "sum_total": 420, "value": 420})];
3102 let p = vec![Col::renamed("sum_total", "sum")];
3103 let cols = columns_for(&rows, &p);
3104 assert_eq!(names(&cols), vec!["sum"], "one column, SQL's name");
3105 assert_eq!(cell(rows[0].get(&cols[0].src)), Some("420".to_string()));
3106 }
3107
3108 #[test]
3113 fn a_bare_column_with_group_by_is_refused_not_nulled() {
3114 let e = translate("SELECT region, total FROM orders GROUP BY region").unwrap_err();
3115 assert!(e.contains("must appear in the GROUP BY clause"), "{}", e);
3116 assert!(e.contains("total"), "the message names the offending column: {}", e);
3117
3118 assert!(translate("SELECT region FROM orders GROUP BY region").is_ok());
3120 assert!(translate("SELECT region, count FROM orders GROUP BY region").is_ok());
3121 assert!(translate("SELECT SUM(total) FROM orders GROUP BY region").is_ok());
3123 assert!(translate("SELECT * FROM orders GROUP BY region").is_ok());
3125 }
3126
3127 #[test]
3128 fn count_star_becomes_nql_count() {
3129 assert_eq!(q("SELECT COUNT(*) FROM orders"), "FROM orders COUNT");
3130 assert_eq!(q("SELECT count(*) FROM orders WHERE total > 5"),
3131 "FROM orders COUNT WHERE total > 5");
3132 }
3133
3134 #[test]
3135 fn aggregates_carry_their_target_column() {
3136 assert_eq!(q("SELECT SUM(total) FROM orders"), "FROM orders SUM total");
3137 assert_eq!(q("SELECT avg(total) FROM orders WHERE region = 'eu'"),
3138 r#"FROM orders AVG total WHERE region = "eu""#);
3139 assert!(translate("SELECT SUM(*) FROM orders").is_err());
3140 }
3141
3142 #[test]
3145 fn as_of_system_time_bridges_to_nql_as_of() {
3146 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42"),
3147 "FROM orders AS OF 42");
3148 assert_eq!(q("SELECT * FROM orders AS OF SYSTEM TIME 42 WHERE total > 1"),
3149 "FROM orders AS OF 42 WHERE total > 1");
3150 let e = translate("SELECT * FROM orders AS OF SYSTEM TIME '2026-01-01'").unwrap_err();
3152 assert!(e.contains("sequence number"), "{}", e);
3153 }
3154
3155 #[test]
3156 fn handshake_queries_are_answered_so_clients_can_connect() {
3157 assert!(matches!(translate("SELECT version()"), Ok(Stmt::Canned { .. })));
3158 assert!(matches!(translate("SHOW transaction_isolation"), Ok(Stmt::Canned { .. })));
3159 assert!(matches!(translate("SELECT current_schema()"), Ok(Stmt::Canned { .. })));
3160 assert!(matches!(translate("SET extra_float_digits = 3"), Ok(Stmt::Ok(_))));
3161 assert!(matches!(translate("BEGIN"), Ok(Stmt::Ok(_))));
3162 assert!(matches!(translate(""), Ok(Stmt::Ok(_))));
3163 }
3164
3165 #[test]
3168 fn unsupported_sql_is_refused_with_a_reason() {
3169 for (sql, expect) in [
3170 ("INSERT INTO t VALUES (1)", "explicit column list"),
3171 ("CREATE TABLE t (a int)", "DDL"),
3172 ("TRUNCATE t", "append-only"),
3173 ("GRANT ALL ON t TO x", "privilege system"),
3174 ("SELECT * FROM a JOIN b ON a.x = b.x", "JOIN is not supported"),
3175 ("SELECT * FROM a UNION SELECT * FROM b", "UNION"),
3176 ("SELECT DISTINCT region FROM orders", "GROUP BY"),
3177 ("SELECT * FROM (SELECT 1) x", "subqueries in FROM"),
3178 ("SELECT * FROM a, b", "more than one collection"),
3179 ("SELECT lower(status) FROM orders", "expressions in the select list"),
3180 ("VACUUM", "only SELECT"),
3181 ] {
3182 let e = translate(sql).unwrap_err();
3183 assert!(e.contains(expect), "for {:?} expected {:?} in {:?}", sql, expect, e);
3184 }
3185 }
3186
3187 fn ins(sql: &str) -> (String, Vec<InsertRow>, Vec<Col>) {
3195 match translate(sql) {
3196 Ok(Stmt::Insert { coll, rows, returning }) => (coll, rows, returning),
3197 other => panic!("expected INSERT for {:?}, got {:?}", sql, other),
3198 }
3199 }
3200
3201 #[test]
3202 fn insert_becomes_a_put_per_row() {
3203 let (coll, rows, ret) = ins("INSERT INTO orders (_id, status, total) VALUES ('o1', 'paid', 120)");
3204 assert_eq!(coll, "orders");
3205 assert_eq!(rows.len(), 1);
3206 assert_eq!(rows[0].id.as_deref(), Some("o1"));
3207 assert_eq!(rows[0].doc.get("status"), Some(&json!("paid")));
3208 assert_eq!(rows[0].doc.get("total"), Some(&json!(120)));
3209 assert!(!rows[0].doc.contains_key("_id"));
3211 assert!(ret.is_empty());
3212 }
3213
3214 #[test]
3215 fn a_multi_row_insert_yields_one_row_each() {
3216 let (_, rows, _) = ins(
3217 "INSERT INTO t (id, n) VALUES ('a', 1), ('b', 2), ('c', 3)");
3218 assert_eq!(rows.len(), 3);
3219 assert_eq!(rows[1].id.as_deref(), Some("b"));
3220 assert_eq!(rows[2].doc.get("n"), Some(&json!(3)));
3221 }
3222
3223 #[test]
3224 fn an_insert_without_an_id_column_lets_the_server_assign_one() {
3225 let (_, rows, _) = ins("INSERT INTO t (n) VALUES (1)");
3226 assert_eq!(rows[0].id, None, "the executor mints a unique key");
3227 assert_eq!(rows[0].doc.get("n"), Some(&json!(1)));
3228 }
3229
3230 #[test]
3233 fn insert_lifts_provenance_out_of_reserved_columns() {
3234 let (_, rows, _) = ins(
3235 "INSERT INTO audit (_id, _caused_by, _valid_from, kind) \
3236 VALUES ('e1', 'abc123', '2026-01-01', 'reprice')");
3237 assert_eq!(rows[0].caused_by, vec!["abc123".to_string()]);
3238 assert_eq!(rows[0].valid_from.as_deref(), Some("2026-01-01"));
3239 assert_eq!(rows[0].doc.get("kind"), Some(&json!("reprice")));
3240 for k in ["_id", "_caused_by", "_valid_from"] {
3242 assert!(!rows[0].doc.contains_key(k), "{} leaked into the doc", k);
3243 }
3244 }
3245
3246 #[test]
3247 fn insert_values_cover_the_scalar_types() {
3248 let (_, rows, _) = ins(
3249 "INSERT INTO t (s, i, f, b, n) VALUES ('x', 42, 1.5, TRUE, NULL)");
3250 assert_eq!(rows[0].doc.get("s"), Some(&json!("x")));
3251 assert_eq!(rows[0].doc.get("i"), Some(&json!(42)));
3252 assert_eq!(rows[0].doc.get("f"), Some(&json!(1.5)));
3253 assert_eq!(rows[0].doc.get("b"), Some(&json!(true)));
3254 assert_eq!(rows[0].doc.get("n"), Some(&Value::Null));
3255 }
3256
3257 #[test]
3260 fn insert_literals_survive_quotes_and_commas() {
3261 let (_, rows, _) = ins("INSERT INTO t (a, b) VALUES ('it''s', 'x,y')");
3262 assert_eq!(rows[0].doc.get("a"), Some(&json!("it's")));
3263 assert_eq!(rows[0].doc.get("b"), Some(&json!("x,y")));
3264 }
3265
3266 #[test]
3267 fn insert_refuses_what_it_cannot_store_faithfully() {
3268 assert!(translate("INSERT INTO t (a) VALUES (1 + 1)").is_err());
3270 assert!(translate("INSERT INTO t (a) VALUES (now())").is_err());
3271 let e = translate("INSERT INTO t (a, b) VALUES (1)").unwrap_err();
3273 assert!(e.contains("values for"), "{}", e);
3274 let e2 = translate("INSERT INTO t VALUES (1)").unwrap_err();
3276 assert!(e2.contains("explicit column list"), "{}", e2);
3277 }
3278
3279 #[test]
3280 fn update_finds_rows_with_the_full_predicate_surface() {
3281 match translate("UPDATE orders SET status = 'void' WHERE total < 50 AND region IN ('eu')") {
3282 Ok(Stmt::Update { coll, set, nql, .. }) => {
3283 assert_eq!(coll, "orders");
3284 assert_eq!(set, vec![("status".to_string(), json!("void"))]);
3285 assert_eq!(nql, r#"FROM orders WHERE total < 50 AND region IN ("eu")"#);
3287 }
3288 other => panic!("expected UPDATE, got {:?}", other),
3289 }
3290 }
3291
3292 #[test]
3293 fn update_without_where_targets_the_whole_collection() {
3294 match translate("UPDATE t SET a = 1") {
3296 Ok(Stmt::Update { nql, .. }) => assert_eq!(nql, "FROM t"),
3297 other => panic!("expected UPDATE, got {:?}", other),
3298 }
3299 }
3300
3301 #[test]
3302 fn update_handles_several_assignments() {
3303 match translate("UPDATE t SET a = 1, b = 'x,y', c = NULL WHERE id = 'k'") {
3304 Ok(Stmt::Update { set, .. }) => {
3305 assert_eq!(set.len(), 3);
3306 assert_eq!(set[1], ("b".to_string(), json!("x,y")));
3307 assert_eq!(set[2], ("c".to_string(), Value::Null));
3308 }
3309 other => panic!("expected UPDATE, got {:?}", other),
3310 }
3311 assert!(translate("UPDATE t SET").is_err());
3312 assert!(translate("UPDATE t SET a").is_err());
3313 }
3314
3315 #[test]
3316 fn delete_becomes_a_predicate_over_the_collection() {
3317 match translate("DELETE FROM orders WHERE status = 'void'") {
3318 Ok(Stmt::Delete { coll, nql, .. }) => {
3319 assert_eq!(coll, "orders");
3320 assert_eq!(nql, r#"FROM orders WHERE status = "void""#);
3321 }
3322 other => panic!("expected DELETE, got {:?}", other),
3323 }
3324 match translate("DELETE FROM t") {
3325 Ok(Stmt::Delete { nql, .. }) => assert_eq!(nql, "FROM t"),
3326 other => panic!("expected DELETE, got {:?}", other),
3327 }
3328 }
3329
3330 #[test]
3331 fn returning_is_parsed_off_every_write() {
3332 let (_, _, ret) = ins("INSERT INTO t (a) VALUES (1) RETURNING a, _id");
3333 assert_eq!(ret.iter().map(|c| c.out.clone()).collect::<Vec<_>>(), vec!["a", "_id"]);
3334 let (_, _, star) = ins("INSERT INTO t (a) VALUES (1) RETURNING *");
3337 assert!(star.is_empty());
3338 assert!(wants_returning("INSERT INTO t (a) VALUES (1) RETURNING *"));
3339 assert!(!wants_returning("INSERT INTO t (a) VALUES (1)"));
3340
3341 match translate("UPDATE t SET a = 1 WHERE id = 'k' RETURNING a") {
3342 Ok(Stmt::Update { nql, returning, .. }) => {
3343 assert_eq!(returning.len(), 1);
3344 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql);
3346 }
3347 other => panic!("expected UPDATE, got {:?}", other),
3348 }
3349 match translate("DELETE FROM t WHERE id = 'k' RETURNING *") {
3350 Ok(Stmt::Delete { nql, .. }) =>
3351 assert!(!nql.to_uppercase().contains("RETURNING"), "{}", nql),
3352 other => panic!("expected DELETE, got {:?}", other),
3353 }
3354 }
3355
3356 #[test]
3357 fn a_keyword_inside_a_value_is_not_a_clause() {
3358 match translate("UPDATE t SET note = 'where returning from' WHERE id = 'k'") {
3359 Ok(Stmt::Update { set, nql, .. }) => {
3360 assert_eq!(set[0].1, json!("where returning from"));
3361 assert_eq!(nql, r#"FROM t WHERE id = "k""#);
3362 }
3363 other => panic!("expected UPDATE, got {:?}", other),
3364 }
3365 }
3366
3367 #[test]
3368 fn split_top_respects_quotes_and_nesting() {
3369 assert_eq!(split_top("a, b, c", ',').len(), 3);
3370 assert_eq!(split_top("(1, 2), (3, 4)", ',').len(), 2);
3371 assert_eq!(split_top("'a,b', c", ',').len(), 2);
3372 assert_eq!(split_top("'it''s, fine', c", ',').len(), 2);
3373 }
3374
3375 #[test]
3376 fn comments_and_whitespace_do_not_confuse_the_translator() {
3377 assert_eq!(q("SELECT *\n FROM orders -- trailing note\n"), "FROM orders");
3378 assert_eq!(q("SELECT /* inline */ * FROM orders"), "FROM orders");
3379 assert_eq!(q("SELECT * FROM t WHERE note = 'from here to JOIN'"),
3381 r#"FROM t WHERE note = "from here to JOIN""#);
3382 }
3383
3384 #[test]
3385 fn find_kw_ignores_quotes_parens_and_substrings() {
3386 assert_eq!(find_kw("SELECT A FROM B", "FROM"), Some(9));
3387 assert_eq!(find_kw("SELECT 'FROM' FROM B", "FROM"), Some(14));
3388 assert_eq!(find_kw("SELECT F(x FROM y) FROM B", "FROM"), Some(19));
3389 assert_eq!(find_kw("SELECT FROMAGE", "FROM"), None);
3390 assert_eq!(find_kw("SELECT X_FROM", "FROM"), None);
3391 }
3392
3393 #[test]
3396 fn provenance_columns_sort_after_the_users_own_fields() {
3397 let rows = vec![json!({"_id":"1","_hash":"ab","status":"paid","total":9})];
3398 assert_eq!(names(&columns_for(&rows, &[])),
3399 vec!["status", "total", "_hash", "_id"]);
3400 }
3401
3402 #[test]
3403 fn an_explicit_projection_sets_the_column_order() {
3404 let rows = vec![json!({"a":1,"b":2})];
3405 let p = vec![Col::same("b"), Col::same("a")];
3406 assert_eq!(names(&columns_for(&rows, &p)), vec!["b", "a"]);
3407 }
3408
3409 #[test]
3410 fn columns_are_the_union_across_sparse_rows() {
3411 let rows = vec![json!({"a":1}), json!({"b":2})];
3413 assert_eq!(names(&columns_for(&rows, &[])), vec!["a", "b"]);
3414 }
3415
3416 #[test]
3417 fn type_oids_follow_the_first_non_null_value() {
3418 let rows = vec![json!({"i":1,"f":1.5,"b":true,"s":"x","n":null})];
3419 assert_eq!(oid_for(&rows, "i"), OID_INT8);
3420 assert_eq!(oid_for(&rows, "f"), OID_FLOAT8);
3421 assert_eq!(oid_for(&rows, "b"), OID_BOOL);
3422 assert_eq!(oid_for(&rows, "s"), OID_TEXT);
3423 assert_eq!(oid_for(&rows, "n"), OID_TEXT);
3425 assert_eq!(oid_for(&rows, "absent"), OID_TEXT);
3426 }
3427
3428 #[test]
3429 fn a_column_that_is_null_in_the_first_row_still_gets_its_type() {
3430 let rows = vec![json!({"v": null}), json!({"v": 7})];
3431 assert_eq!(oid_for(&rows, "v"), OID_INT8);
3432 }
3433
3434 #[test]
3435 fn cells_render_in_postgres_text_format() {
3436 assert_eq!(cell(Some(&json!("x"))), Some("x".to_string()));
3437 assert_eq!(cell(Some(&json!(true))), Some("t".to_string()));
3438 assert_eq!(cell(Some(&json!(false))), Some("f".to_string()));
3439 assert_eq!(cell(Some(&json!(42))), Some("42".to_string()));
3440 assert_eq!(cell(Some(&json!(null))), None);
3441 assert_eq!(cell(None), None);
3442 assert_eq!(cell(Some(&json!({"a":1}))), Some("{\"a\":1}".to_string()));
3444 }
3445
3446 #[test]
3449 fn message_framing_length_excludes_the_tag() {
3450 let mut m = Out::msg(b'Z');
3451 m.bytes(b"I");
3452 let bytes = m.finish();
3453 assert_eq!(bytes[0], b'Z');
3454 assert_eq!(i32::from_be_bytes([bytes[1], bytes[2], bytes[3], bytes[4]]), 5);
3455 assert_eq!(bytes.len(), 6);
3456 }
3457
3458 #[test]
3459 fn a_result_set_encodes_as_description_then_rows_then_complete() {
3460 let rows = vec![json!({"a": 1}), json!({"a": 2})];
3461 let out = encode_result(&rows, &[]);
3462 assert_eq!(out[0], b'T');
3463 let tags: Vec<u8> = {
3464 let mut t = vec![];
3466 let mut i = 0usize;
3467 while i < out.len() {
3468 t.push(out[i]);
3469 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
3470 i += 1 + len;
3471 }
3472 t
3473 };
3474 assert_eq!(tags, vec![b'T', b'D', b'D', b'C'],
3475 "one description, one row each, one completion");
3476 }
3477
3478 #[test]
3483 fn a_write_with_returning_emits_exactly_one_command_complete() {
3484 let rows = vec![json!({"_id": "o1", "total": 9})];
3485 let mut out = encode_rows(&rows, &[Col::same("_id")]);
3486 out.extend_from_slice(&command_complete("INSERT 0 1"));
3487 let mut tags = vec![];
3488 let mut i = 0usize;
3489 while i < out.len() {
3490 tags.push(out[i]);
3491 let len = i32::from_be_bytes([out[i+1], out[i+2], out[i+3], out[i+4]]) as usize;
3492 i += 1 + len;
3493 }
3494 assert_eq!(tags, vec![b'T', b'D', b'C'], "one description, one row, ONE tag");
3495 assert_eq!(tags.iter().filter(|t| **t == b'C').count(), 1);
3496 assert!(!encode_rows(&rows, &[]).contains(&b'C')
3498 || encode_rows(&rows, &[]).iter().filter(|b| **b == b'C').count() > 0);
3499 let bare = encode_rows(&rows, &[Col::same("_id")]);
3500 let mut bare_tags = vec![];
3501 let mut j = 0usize;
3502 while j < bare.len() {
3503 bare_tags.push(bare[j]);
3504 let len = i32::from_be_bytes([bare[j+1], bare[j+2], bare[j+3], bare[j+4]]) as usize;
3505 j += 1 + len;
3506 }
3507 assert_eq!(bare_tags, vec![b'T', b'D'], "encode_rows never appends a tag");
3508 }
3509
3510 #[test]
3511 fn an_empty_result_still_sends_a_description() {
3512 let out = encode_result(&[], &[Col::same("a")]);
3513 assert_eq!(out[0], b'T', "clients need the shape even with no rows");
3514 }
3515
3516 #[test]
3517 fn statements_split_on_top_level_semicolons_only() {
3518 assert_eq!(split_statements("SELECT 1; SELECT 2").len(), 2);
3519 assert_eq!(split_statements("SELECT ';'").len(), 1);
3520 assert_eq!(split_statements("SELECT 1;").len(), 1);
3521 assert_eq!(split_statements(" ").len(), 0);
3522 }
3523
3524 #[test]
3525 fn an_error_names_its_sqlstate() {
3526 let e = String::from_utf8_lossy(&err_msg("0A000", "x")).to_string();
3527 assert!(e.contains("ERROR"));
3528 assert!(e.contains("0A000"));
3529 }
3530
3531 #[test]
3534 fn placeholders_are_counted_outside_string_literals() {
3535 assert_eq!(param_count("SELECT a FROM t WHERE b = $1 AND c = $2"), 2);
3536 assert_eq!(param_count("SELECT a FROM t"), 0);
3537 assert_eq!(param_count("WHERE a = $2 OR b = $2 OR c = $1"), 2);
3539 assert_eq!(param_count("SELECT a FROM t WHERE b = '$1'"), 0,
3540 "a placeholder inside a literal is data, not a parameter");
3541 assert_eq!(param_count("WHERE a = $10 AND b = $1"), 10,
3542 "two-digit indexes must not be read as $1 followed by 0");
3543 }
3544
3545 #[test]
3546 fn parameters_are_spliced_as_literals() {
3547 let out = substitute_params("WHERE a = $1 AND b = $2 AND c = $3",
3548 &[Some("'x'".into()), Some("42".into()), None]).unwrap();
3549 assert_eq!(out, "WHERE a = 'x' AND b = 42 AND c = NULL");
3550 }
3551
3552 #[test]
3553 fn substitution_leaves_string_literals_alone() {
3554 let out = substitute_params("WHERE a = '$1' AND b = $1", &[Some("9".into())]).unwrap();
3555 assert_eq!(out, "WHERE a = '$1' AND b = 9");
3556 }
3557
3558 #[test]
3559 fn too_few_parameters_is_an_error_not_a_silent_null() {
3560 let e = substitute_params("WHERE a = $2", &[Some("1".into())]).unwrap_err();
3563 assert!(e.contains("$2"), "{}", e);
3564 }
3565
3566 #[test]
3567 fn a_quote_in_a_parameter_cannot_escape_its_literal() {
3568 let lit = decode_param(Some(b"it's"), OID_TEXT, 0).unwrap().unwrap();
3569 assert_eq!(lit, "'it''s'");
3570 let out = substitute_params("WHERE a = $1", &[Some(lit)]).unwrap();
3572 assert_eq!(out, "WHERE a = 'it''s'");
3573 }
3574
3575 #[test]
3576 fn binary_parameters_decode_in_every_width_psycopg_sends() {
3577 assert_eq!(decode_param(Some(&[0x00, 0x2a]), OID_INT2, 1).unwrap().unwrap(), "42");
3580 assert_eq!(decode_param(Some(&[0, 0, 0, 7]), OID_INT4, 1).unwrap().unwrap(), "7");
3581 assert_eq!(
3582 decode_param(Some(&[0, 0, 0, 0, 0, 0, 0, 9]), OID_INT8, 1).unwrap().unwrap(), "9");
3583 assert_eq!(
3584 decode_param(Some(&0x400c_0000_0000_0000u64.to_be_bytes()), OID_FLOAT8, 1)
3585 .unwrap().unwrap(), "3.5");
3586 assert_eq!(decode_param(Some(&[1]), OID_BOOL, 1).unwrap().unwrap(), "TRUE");
3587 assert_eq!(decode_param(Some(&[0]), OID_BOOL, 1).unwrap().unwrap(), "FALSE");
3588 }
3589
3590 #[test]
3591 fn a_negative_binary_integer_keeps_its_sign() {
3592 assert_eq!(decode_param(Some(&(-5i32).to_be_bytes()), OID_INT4, 1).unwrap().unwrap(), "-5");
3593 assert_eq!(decode_param(Some(&(-5i16).to_be_bytes()), OID_INT2, 1).unwrap().unwrap(), "-5");
3594 }
3595
3596 #[test]
3597 fn a_binary_parameter_of_the_wrong_width_is_refused() {
3598 let e = decode_param(Some(&[0x2a]), OID_INT4, 1).unwrap_err();
3601 assert!(e.contains("4 bytes"), "{}", e);
3602 }
3603
3604 #[test]
3605 fn an_unspecified_text_parameter_is_treated_as_a_string() {
3606 assert_eq!(decode_param(Some(b"hello"), 0, 0).unwrap().unwrap(), "'hello'");
3609 }
3610
3611 #[test]
3612 fn a_null_parameter_decodes_to_none_in_every_format() {
3613 assert_eq!(decode_param(None, OID_TEXT, 0).unwrap(), None);
3614 assert_eq!(decode_param(None, OID_INT8, 1).unwrap(), None);
3615 }
3616
3617 #[test]
3618 fn an_unsupported_binary_type_says_so_by_name() {
3619 let e = decode_param(Some(&[0u8; 8]), 1114, 1).unwrap_err();
3620 assert!(e.contains("1114"), "{}", e);
3621 assert!(e.contains("text"), "the error should point at the way out: {}", e);
3622 }
3623
3624 #[test]
3625 fn a_text_number_that_is_not_a_number_gets_quoted() {
3626 assert_eq!(decode_param(Some(b"oops"), OID_INT8, 0).unwrap().unwrap(), "'oops'");
3629 }
3630
3631 #[test]
3632 fn a_client_declared_type_is_believed_over_inference() {
3633 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[OID_INT4, 0], None);
3636 assert_eq!(oids, vec![OID_INT4, OID_TEXT]);
3637 }
3638
3639 #[test]
3640 fn parameter_arity_is_taken_from_the_sql_when_the_client_declares_none() {
3641 let oids = infer_param_oids("SELECT a FROM t WHERE b = $1 AND c = $2", &[], None);
3644 assert_eq!(oids.len(), 2);
3645 }
3646
3647 #[test]
3648 fn the_field_behind_each_placeholder_is_identified() {
3649 assert_eq!(
3650 param_fields("SELECT a FROM t WHERE qty > $1 AND status = $2", 2),
3651 vec![Some("qty".to_string()), Some("status".to_string())]);
3652 }
3653
3654 #[test]
3655 fn word_operators_do_not_hide_the_field() {
3656 assert_eq!(param_fields("SELECT a FROM t WHERE name LIKE $1", 1),
3657 vec![Some("name".to_string())]);
3658 assert_eq!(param_fields("SELECT a FROM t WHERE qty BETWEEN $1 AND $2", 2),
3659 vec![Some("qty".to_string()), Some("qty".to_string())]);
3660 assert_eq!(param_fields("SELECT a FROM t WHERE region IN ($1, $2)", 2),
3661 vec![Some("region".to_string()), Some("region".to_string())]);
3662 }
3663
3664 #[test]
3665 fn a_clause_position_types_from_the_grammar_not_from_a_column() {
3666 assert_eq!(
3670 infer_param_oids("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", &[], None),
3671 vec![OID_INT8, OID_TEXT]);
3672 assert_eq!(infer_param_oids("SELECT a FROM t AS OF $1", &[], None), vec![OID_INT8]);
3673 assert_eq!(
3676 infer_param_oids("SELECT a FROM t VALID AS OF $1", &[], None), vec![OID_TEXT]);
3677 assert_eq!(
3678 infer_param_oids("SELECT a FROM t LIMIT $1 OFFSET $2", &[], None),
3679 vec![OID_INT8, OID_INT8]);
3680 }
3681
3682 #[test]
3683 fn an_aggregate_column_types_from_what_the_aggregate_means() {
3684 assert_eq!(aggregate_oid("count", None, "t"), Some(OID_INT8));
3688 assert_eq!(aggregate_oid("avg_fee", None, "t"), Some(OID_FLOAT8),
3689 "an average is fractional even over integers");
3690 assert_eq!(aggregate_oid("max__seq", None, "t"), Some(OID_INT8));
3693 assert_eq!(aggregate_oid("total", None, "t"), None, "not an aggregate");
3694 }
3695
3696 #[test]
3697 fn the_parse_probe_uses_a_literal_that_every_clause_accepts() {
3698 let probe = probe_sql("SELECT a FROM t AS OF SYSTEM TIME $1 WHERE b = $2", 2);
3702 assert!(!probe.contains("NULL"), "{}", probe);
3703 assert!(translate(&probe).is_ok(), "the probe must parse: {}", probe);
3704 }
3705
3706 #[test]
3707 fn a_column_with_mixed_types_across_documents_is_advertised_as_text() {
3708 let rows = vec![json!({"x": 3}), json!({"x": "n/a"})];
3712 assert_eq!(oid_for(&rows, "x"), OID_TEXT);
3713 let rows = vec![json!({"x": 3}), json!({"x": 1.5})];
3715 assert_eq!(oid_for(&rows, "x"), OID_FLOAT8);
3716 let rows = vec![json!({"x": Value::Null}), json!({"x": 7})];
3718 assert_eq!(oid_for(&rows, "x"), OID_INT8);
3719 }
3720
3721 #[test]
3722 fn binary_output_encodes_each_advertised_type() {
3723 assert_eq!(cell_binary(Some(&json!(true)), OID_BOOL).unwrap().unwrap(), vec![1]);
3724 assert_eq!(cell_binary(Some(&json!(42)), OID_INT8).unwrap().unwrap(),
3725 42i64.to_be_bytes().to_vec());
3726 assert_eq!(cell_binary(Some(&json!(3.5)), OID_FLOAT8).unwrap().unwrap(),
3727 3.5f64.to_be_bytes().to_vec());
3728 assert_eq!(cell_binary(Some(&json!("hi")), OID_TEXT).unwrap().unwrap(), b"hi".to_vec());
3730 assert_eq!(cell_binary(Some(&Value::Null), OID_INT8).unwrap(), None);
3731 assert_eq!(cell(Some(&json!(true))).unwrap(), "t");
3733 }
3734
3735 #[test]
3736 fn a_value_that_does_not_fit_its_advertised_binary_type_is_refused() {
3737 let e = cell_binary(Some(&json!("nope")), OID_INT8).unwrap_err();
3742 assert!(e.contains("a string"), "{}", e);
3743 assert!(e.contains("more than one type"), "the error should explain WHY: {}", e);
3744 }
3745
3746 #[test]
3747 fn a_row_description_carries_the_requested_format_per_column() {
3748 let cols = [Col::same("a"), Col::same("b")];
3749 let m = row_description_fmt(&cols, &[OID_INT8, OID_TEXT], &[1, 0]);
3750 assert_eq!(m[0], b'T');
3751 assert_eq!(m[m.len() - 1], 0, "the last column was requested as text");
3753 }
3754
3755 #[test]
3756 fn a_qualified_column_resolves_to_its_bare_name() {
3757 assert_eq!(param_fields("SELECT a FROM t WHERE t.qty = $1", 1),
3758 vec![Some("qty".to_string())]);
3759 }
3760
3761 #[test]
3762 fn insert_placeholders_map_positionally_to_the_column_list() {
3763 assert_eq!(
3764 param_fields("INSERT INTO t (_id, qty, status) VALUES ($1, $2, $3)", 3),
3765 vec![Some("_id".to_string()), Some("qty".to_string()), Some("status".to_string())]);
3766 }
3767
3768 #[test]
3769 fn a_set_clause_placeholder_finds_its_column() {
3770 assert_eq!(param_fields("UPDATE t SET status = $1 WHERE _id = $2", 2),
3771 vec![Some("status".to_string()), Some("_id".to_string())]);
3772 }
3773
3774 #[test]
3775 fn the_target_collection_is_found_for_every_statement_kind() {
3776 assert_eq!(stmt_collection("SELECT a FROM inv WHERE b = $1"), "inv");
3777 assert_eq!(stmt_collection("UPDATE inv SET a = $1"), "inv");
3778 assert_eq!(stmt_collection("DELETE FROM inv WHERE a = $1"), "inv");
3779 assert_eq!(stmt_collection("INSERT INTO inv (a) VALUES ($1)"), "inv");
3780 assert_eq!(stmt_collection("SELECT a FROM public.inv"), "inv");
3782 assert_eq!(stmt_collection("INSERT INTO inv(a) VALUES ($1)"), "inv");
3783 }
3784
3785 #[test]
3786 fn engine_metadata_fields_type_without_touching_storage() {
3787 assert_eq!(infer_field_oid(None, "t", "_seq"), OID_INT8);
3788 assert_eq!(infer_field_oid(None, "t", "_id"), OID_TEXT);
3789 }
3790
3791 #[test]
3792 fn the_protocol_acknowledgements_are_single_empty_messages() {
3793 for (m, tag) in [
3795 (parse_complete(), b'1'), (bind_complete(), b'2'),
3796 (close_complete(), b'3'), (no_data(), b'n'), (portal_suspended(), b's'),
3797 ] {
3798 assert_eq!(m.len(), 5, "{:?}", tag as char);
3799 assert_eq!(m[0], tag);
3800 assert_eq!(i32::from_be_bytes([m[1], m[2], m[3], m[4]]), 4);
3801 }
3802 }
3803
3804 #[test]
3805 fn parameter_description_reports_its_arity_and_types() {
3806 let m = parameter_description(&[OID_TEXT, OID_INT8]);
3807 assert_eq!(m[0], b't');
3808 assert_eq!(i16::from_be_bytes([m[5], m[6]]), 2);
3809 assert_eq!(i32::from_be_bytes([m[7], m[8], m[9], m[10]]), OID_TEXT);
3810 assert_eq!(i32::from_be_bytes([m[11], m[12], m[13], m[14]]), OID_INT8);
3811 }
3812
3813 #[test]
3814 fn a_cstring_is_taken_without_its_terminator() {
3815 let body = b"one\0two\0".to_vec();
3816 let mut at = 0usize;
3817 assert_eq!(take_cstr(&body, &mut at), "one");
3818 assert_eq!(take_cstr(&body, &mut at), "two");
3819 assert_eq!(at, body.len());
3820 }
3821
3822 #[test]
3823 fn truncated_integers_are_reported_rather_than_read_past_the_end() {
3824 let body = vec![0u8, 1];
3825 let mut at = 0usize;
3826 assert!(take_i32(&body, &mut at).is_err());
3827 let mut at = 0usize;
3828 assert!(take_i16(&body, &mut at).is_ok());
3829 }
3830
3831 #[test]
3832 fn a_binary_result_format_request_is_refused_rather_than_faked() {
3833 let out = encode_rows(&[], &[Col::same("a")]);
3836 let desc_format = &out[out.len() - 2..];
3837 assert_eq!(i16::from_be_bytes([desc_format[0], desc_format[1]]), 0,
3838 "every column is advertised as text format");
3839 }
3840
3841 #[test]
3842 fn a_float_parameter_does_not_render_as_rust_infinity() {
3843 assert_eq!(fmt_float(f64::INFINITY), "'Infinity'");
3844 assert_eq!(fmt_float(f64::NEG_INFINITY), "'-Infinity'");
3845 assert_eq!(fmt_float(f64::NAN), "'NaN'");
3846 assert_eq!(fmt_float(3.0), "3", "a whole float should not gain a .0 tail");
3847 assert_eq!(fmt_float(3.5), "3.5");
3848 }
3849}