Skip to main content

pylon_pgcon/
wire.rs

1//
2// This source file is part of the Pylon open source project.
3//
4// Copyright (c) 2026 Jaldis B.V.
5//
6// Licensed under the MIT OR Apache-2.0 license (the "License");
7// you may not use this file except in compliance with the License.
8// You may obtain a copy of the License at
9//
10//     https://opensource.org/licenses/MIT
11//     https://www.apache.org/licenses/LICENSE-2.0
12//
13// Unless required by applicable law or agreed to in writing, software
14// distributed under the License is distributed on an "AS IS" BASIS,
15// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
16// See the License for the specific language governing permissions and
17// limitations under the License.
18//
19
20//! Recursive decoder from PostgreSQL's binary wire format into
21//! `pylon_value::DecodedValue` — the shared decode target `pylon-cache` also
22//! stores, so a cache hit and a fresh row decode into the exact same shape.
23//!
24//! Every PyQL query result is emitted by `pylon-core` as a single
25//! `SELECT (...) AS result` — an anonymous composite (`record`, OID 2249).
26//! `tokio-postgres` has no generic "decode any composite into a value tree"
27//! API (its `FromSql` machinery targets known Rust types), so this module
28//! walks PostgreSQL's own documented binary wire format directly, the same
29//! way the outgoing Python implementation's `_pg_decode_record`/
30//! `_pg_decode_value` (`pylon/client.py`) did — except comprehensive,
31//! comprehensively, including the cases a generic composite decoder plus
32//! hand-registered codec overrides tends to miss (jsonb, `record[]`,
33//! `vector`).
34//!
35//! Composite field layout (used recursively for `record`/`record[]`):
36//! `i32 nfields`, then per field: `u32 type_oid`, `i32 field_len`
37//! (`-1` = NULL), `field_len` bytes of that field's own wire encoding.
38//! Array layout (used for every `T[]` OID below): `i32 ndim`, `i32
39//! has_null_flag`, `u32 element_oid`, then *one* `(i32 dim_size, i32
40//! lower_bound)` pair — Pylon's `pylon.Array[T]` is always 1-dimensional,
41//! so multi-dimensional arrays are out of scope, same as the Python
42//! implementation this replaces — then per element: `i32 len` (`-1` =
43//! NULL) + `len` bytes.
44
45use pylon_value::DecodedValue;
46
47use crate::numeric;
48
49pub use crate::error::Error;
50pub type Result<T> = crate::Result<T>;
51
52// Fixed, well-known OIDs (see `pg_type.h` / `SELECT oid, typname FROM
53// pg_type`) — stable across every Postgres install, unlike extension types
54// (`vector`, PostGIS geometry/geography), enums, and domains, whose OIDs are
55// assigned by the local database and are discovered per-connection instead
56// (see `ExtensionOids` and `TYPE_DISCOVERY_SQL`).
57const OID_BOOL: u32 = 16;
58const OID_BYTEA: u32 = 17;
59const OID_INT8: u32 = 20;
60const OID_INT2: u32 = 21;
61const OID_INT4: u32 = 23;
62const OID_TEXT: u32 = 25;
63const OID_JSONB: u32 = 3802;
64const OID_FLOAT4: u32 = 700;
65const OID_FLOAT8: u32 = 701;
66const OID_BPCHAR: u32 = 1042;
67const OID_VARCHAR: u32 = 1043;
68const OID_NUMERIC: u32 = 1700;
69const OID_DATE: u32 = 1082;
70const OID_TIME: u32 = 1083;
71const OID_TIMESTAMP: u32 = 1114;
72const OID_TIMESTAMPTZ: u32 = 1184;
73const OID_INTERVAL: u32 = 1186;
74const OID_UUID: u32 = 2950;
75const OID_RECORD: u32 = 2249;
76const OID_RECORD_ARRAY: u32 = 2287;
77/// The pseudo-type Postgres reports for a literal it never had to resolve to
78/// a concrete type — `SELECT ('doc', ...)` inside a row constructor, for
79/// instance. Its wire format is the value's text representation, so it
80/// decodes exactly like `text`.
81const OID_UNKNOWN: u32 = 705;
82/// `name`, used by the catalogs (`pg_type.typname` and friends). Text-shaped
83/// on the wire, and reachable through the introspection queries.
84const OID_NAME: u32 = 19;
85
86const OID_BOOL_ARRAY: u32 = 1000;
87const OID_BYTEA_ARRAY: u32 = 1001;
88const OID_INT2_ARRAY: u32 = 1005;
89const OID_INT4_ARRAY: u32 = 1007;
90const OID_TEXT_ARRAY: u32 = 1009;
91const OID_BPCHAR_ARRAY: u32 = 1014;
92const OID_VARCHAR_ARRAY: u32 = 1015;
93const OID_INT8_ARRAY: u32 = 1016;
94const OID_FLOAT4_ARRAY: u32 = 1021;
95const OID_FLOAT8_ARRAY: u32 = 1022;
96const OID_NUMERIC_ARRAY: u32 = 1231;
97const OID_UUID_ARRAY: u32 = 2951;
98const OID_JSONB_ARRAY: u32 = 3807;
99
100// Native PostgreSQL range/multirange type OIDs, paired with the element
101// type OID their bound values decode with (int4range's bounds are int4,
102// etc.) — mirrors `range_ctor_for_pg_type`/`multirange_ctor_for_range_ctor`
103// in `pylon-core`'s `ir/compiler.rs`, the other side of this same "which 5
104// PG range families does Pylon support" decision.
105const OID_INT4RANGE: u32 = 3904;
106const OID_INT8RANGE: u32 = 3926;
107const OID_NUMRANGE: u32 = 3906;
108const OID_TSRANGE: u32 = 3908;
109const OID_TSTZRANGE: u32 = 3910;
110const OID_DATERANGE: u32 = 3912;
111const OID_INT4MULTIRANGE: u32 = 4451;
112const OID_INT8MULTIRANGE: u32 = 4536;
113const OID_NUMMULTIRANGE: u32 = 4532;
114const OID_TSMULTIRANGE: u32 = 4533;
115const OID_TSTZMULTIRANGE: u32 = 4534;
116const OID_DATEMULTIRANGE: u32 = 4535;
117
118/// The element OID a range/multirange type's bound values decode with —
119/// `None` for anything that isn't one of the 6 native range/multirange
120/// families this module knows about.
121fn range_element_oid(oid: u32) -> Option<u32> {
122    match oid {
123        OID_INT4RANGE | OID_INT4MULTIRANGE => Some(OID_INT4),
124        OID_INT8RANGE | OID_INT8MULTIRANGE => Some(OID_INT8),
125        OID_NUMRANGE | OID_NUMMULTIRANGE => Some(OID_NUMERIC),
126        OID_TSRANGE | OID_TSMULTIRANGE => Some(OID_TIMESTAMP),
127        OID_TSTZRANGE | OID_TSTZMULTIRANGE => Some(OID_TIMESTAMPTZ),
128        OID_DATERANGE | OID_DATEMULTIRANGE => Some(OID_DATE),
129        _ => None,
130    }
131}
132
133/// Per-database type OIDs, discovered once at connect time and threaded
134/// through decode calls.
135///
136/// Everything in here is assigned by the local database rather than fixed by
137/// `pg_type.h`: extension types get their OID at `CREATE EXTENSION` time,
138/// and enums and domains get theirs when the migration that declares them
139/// runs. The constants at the top of this module cover the built-in types
140/// whose OIDs are stable everywhere; this covers the rest.
141///
142/// `Default` means "nothing discovered yet", under which any OID not in the
143/// built-in set is an `Error::UnknownTypeOid` rather than a guess — see
144/// `decode_value`.
145#[derive(Debug, Clone, Default)]
146pub struct ExtensionOids {
147    /// pgvector's `vector` type, if the extension is installed.
148    pub vector: Option<u32>,
149    /// Every enum type OID. Postgres sends an enum's binary value as its
150    /// label text, so these decode as `Str`.
151    pub enums: std::collections::HashSet<u32>,
152    /// Domain OID to the OID of the type it wraps. A domain's wire format is
153    /// its base type's, so these decode by recursing on the base.
154    pub domains: std::collections::HashMap<u32, u32>,
155    /// Array types whose element is one of the enums or domains above. The
156    /// element OID travels in the array's own binary header, so recognising
157    /// the array OID is all `decode_array` needs; without this an
158    /// `AuthenticationMethod[]` column is an `UnknownTypeOid` even though
159    /// every value in it would decode.
160    pub arrays: std::collections::HashSet<u32>,
161}
162
163/// One round trip that classifies every non-builtin type in the database:
164/// the `vector` extension type, all enums, all domains with the base type
165/// each resolves to, and the array types over any of those.
166///
167/// Arrays are reported with a `'A'` in the typtype column. Postgres never
168/// uses that letter itself (`b`, `c`, `d`, `e`, `m`, `p`, `r` are the real
169/// ones) -- it is `typcategory`'s letter for an array, borrowed here so the
170/// four-column shape holds for every row.
171pub(crate) const TYPE_DISCOVERY_SQL: &str = "\
172SELECT t.oid::int8, t.typtype::text, COALESCE(b.oid, 0)::int8, t.typname::text \
173FROM pg_type t \
174LEFT JOIN pg_type b ON b.oid = t.typbasetype \
175WHERE t.typtype IN ('e', 'd') OR t.typname = 'vector' \
176UNION ALL \
177SELECT a.oid::int8, 'A', e.oid::int8, a.typname::text \
178FROM pg_type a \
179JOIN pg_type e ON e.oid = a.typelem \
180WHERE a.typcategory = 'A' AND e.typtype IN ('e', 'd')";
181
182impl ExtensionOids {
183    /// Builds a registry from `TYPE_DISCOVERY_SQL`'s rows, given as
184    /// `(oid, typtype, base_oid, typname)`.
185    pub(crate) fn from_discovery_rows(rows: impl IntoIterator<Item = (u32, String, u32, String)>) -> Self {
186        let mut out = Self::default();
187        for (oid, typtype, base_oid, typname) in rows {
188            match typtype.as_str() {
189                "e" => {
190                    out.enums.insert(oid);
191                }
192                "d" if base_oid != 0 => {
193                    out.domains.insert(oid, base_oid);
194                }
195                "A" => {
196                    out.arrays.insert(oid);
197                }
198                _ => {}
199            }
200            // `vector` is a base type (typtype 'b'), so it only matches the
201            // name arm of the discovery query.
202            if typname == "vector" {
203                out.vector = Some(oid);
204            }
205        }
206        out
207    }
208}
209
210/// Decodes one field's raw buffer (already length-stripped, matching what
211/// `postgres_types::FromSql::from_sql` receives) into a `DecodedValue`,
212/// given its Postgres type OID. NULL is handled by the caller (a `-1`
213/// field length never reaches this function) — see `decode_record`/
214/// `decode_array` for where that's checked.
215pub fn decode_value(oid: u32, data: &[u8], ext: &ExtensionOids) -> Result<DecodedValue> {
216    if let Some(vector_oid) = ext.vector
217        && oid == vector_oid
218    {
219        return Ok(DecodedValue::Array(decode_vector(data)?));
220    }
221    match oid {
222        OID_BOOL => Ok(DecodedValue::Bool(data.first().copied().unwrap_or(0) != 0)),
223        OID_INT2 => Ok(DecodedValue::I64(i16::from_be_bytes(data.try_into()?) as i64)),
224        OID_INT4 => Ok(DecodedValue::I64(i32::from_be_bytes(data.try_into()?) as i64)),
225        OID_INT8 => Ok(DecodedValue::I64(i64::from_be_bytes(data.try_into()?))),
226        OID_FLOAT4 => Ok(DecodedValue::F64(f32::from_be_bytes(data.try_into()?) as f64)),
227        OID_FLOAT8 => Ok(DecodedValue::F64(f64::from_be_bytes(data.try_into()?))),
228        OID_TEXT | OID_VARCHAR | OID_BPCHAR | OID_UNKNOWN | OID_NAME => {
229            Ok(DecodedValue::Str(std::str::from_utf8(data)?.to_string()))
230        }
231        OID_UUID => {
232            let mut bytes = [0u8; 16];
233            bytes.copy_from_slice(data);
234            Ok(DecodedValue::Uuid(bytes))
235        }
236        OID_BYTEA => Ok(DecodedValue::Bytes(data.to_vec())),
237        OID_NUMERIC => decode_numeric(data),
238        OID_INTERVAL => decode_interval(data),
239        OID_DATE => Ok(DecodedValue::Date(i32::from_be_bytes(data.try_into()?))),
240        OID_TIME => Ok(DecodedValue::Time(i64::from_be_bytes(data.try_into()?))),
241        OID_TIMESTAMP => Ok(DecodedValue::Timestamp(i64::from_be_bytes(data.try_into()?))),
242        OID_TIMESTAMPTZ => Ok(DecodedValue::Timestamptz(i64::from_be_bytes(data.try_into()?))),
243        OID_JSONB => decode_jsonb(data),
244        OID_RECORD => decode_record(data, ext),
245        OID_RECORD_ARRAY => decode_array(data, ext),
246        OID_BOOL_ARRAY | OID_BYTEA_ARRAY | OID_INT2_ARRAY | OID_INT4_ARRAY | OID_INT8_ARRAY | OID_TEXT_ARRAY
247        | OID_BPCHAR_ARRAY | OID_VARCHAR_ARRAY | OID_FLOAT4_ARRAY | OID_FLOAT8_ARRAY | OID_NUMERIC_ARRAY
248        | OID_UUID_ARRAY | OID_JSONB_ARRAY => decode_array(data, ext),
249        OID_INT4RANGE | OID_INT8RANGE | OID_NUMRANGE | OID_TSRANGE | OID_TSTZRANGE | OID_DATERANGE => {
250            decode_range(data, range_element_oid(oid).expect("range OID"), ext)
251        }
252        OID_INT4MULTIRANGE | OID_INT8MULTIRANGE | OID_NUMMULTIRANGE | OID_TSMULTIRANGE | OID_TSTZMULTIRANGE
253        | OID_DATEMULTIRANGE => decode_multirange(data, range_element_oid(oid).expect("multirange OID"), ext),
254        // Everything past this point has a database-assigned OID, so it can
255        // only be resolved through the registry discovered at connect time.
256        _ if ext.enums.contains(&oid) => {
257            // An enum's binary representation is its label, as text.
258            Ok(DecodedValue::Str(std::str::from_utf8(data)?.to_string()))
259        }
260        // The element OID is in the array's own header, so this only has to
261        // recognise that the type is an array at all.
262        _ if ext.arrays.contains(&oid) => decode_array(data, ext),
263        _ => match ext.domains.get(&oid) {
264            // A domain is a constrained alias: same wire format as its base.
265            Some(&base_oid) => decode_value(base_oid, data, ext),
266            None => Err(Error::UnknownTypeOid { oid }),
267        },
268    }
269}
270
271fn decode_numeric(data: &[u8]) -> Result<DecodedValue> {
272    Ok(DecodedValue::Decimal(numeric::decode(data)?))
273}
274
275/// PostgreSQL's binary `interval` wire format: `i64 microseconds, i32 days,
276/// i32 months`, in that order — see `interval_send` in Postgres's own
277/// `timestamp.c`. Backs both `std::duration` and `cal::relative_duration`
278/// (see `DecodedValue::Interval`'s own doc comment for why `months` isn't
279/// folded into `days`).
280fn decode_interval(data: &[u8]) -> Result<DecodedValue> {
281    if data.len() != 16 {
282        return Err(Error::message(format!(
283            "malformed interval: expected 16 bytes, got {}",
284            data.len()
285        )));
286    }
287    let microseconds = i64::from_be_bytes(data[0..8].try_into()?);
288    let days = i32::from_be_bytes(data[8..12].try_into()?);
289    let months = i32::from_be_bytes(data[12..16].try_into()?);
290    Ok(DecodedValue::Interval {
291        months,
292        days,
293        microseconds,
294    })
295}
296
297// PostgreSQL's range binary-format flag bits (`rangetypes.h`).
298const RANGE_EMPTY: u8 = 0x01;
299const RANGE_LB_INC: u8 = 0x02;
300const RANGE_UB_INC: u8 = 0x04;
301const RANGE_LB_INF: u8 = 0x08;
302const RANGE_UB_INF: u8 = 0x10;
303
304/// Binary range: `u8 flags`, then — only when not empty — a length-prefixed
305/// lower bound (skipped if `RANGE_LB_INF`) and a length-prefixed upper bound
306/// (skipped if `RANGE_UB_INF`), each bound decoded with `element_oid`'s own
307/// decoder (see `range_element_oid` for which element type backs which
308/// range OID).
309fn decode_range(data: &[u8], element_oid: u32, ext: &ExtensionOids) -> Result<DecodedValue> {
310    let flags = data[0];
311    let mut offset = 1usize;
312    if flags & RANGE_EMPTY != 0 {
313        return Ok(DecodedValue::Range {
314            lower: None,
315            upper: None,
316            inc_lower: false,
317            inc_upper: false,
318            empty: true,
319        });
320    }
321    let lower = if flags & RANGE_LB_INF != 0 {
322        None
323    } else {
324        let len = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
325        offset += 4;
326        let value = decode_value(element_oid, &data[offset..offset + len], ext)?;
327        offset += len;
328        Some(Box::new(value))
329    };
330    let upper = if flags & RANGE_UB_INF != 0 {
331        None
332    } else {
333        let len = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
334        offset += 4;
335        Some(Box::new(decode_value(element_oid, &data[offset..offset + len], ext)?))
336    };
337    Ok(DecodedValue::Range {
338        lower,
339        upper,
340        inc_lower: flags & RANGE_LB_INC != 0,
341        inc_upper: flags & RANGE_UB_INC != 0,
342        empty: false,
343    })
344}
345
346/// Binary multirange: `i32 range_count`, then per range an `i32 len` +
347/// `len` bytes of that range's own binary encoding (the same format
348/// `decode_range` reads). Decodes to a plain `Array` of `Range` values —
349/// see `DecodedValue::Range`'s own doc comment for why there's no separate
350/// multirange variant.
351fn decode_multirange(data: &[u8], element_oid: u32, ext: &ExtensionOids) -> Result<DecodedValue> {
352    let mut offset = 0usize;
353    let count = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
354    offset += 4;
355    let mut ranges = Vec::with_capacity(count);
356    for _ in 0..count {
357        let len = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
358        offset += 4;
359        ranges.push(decode_range(&data[offset..offset + len], element_oid, ext)?);
360        offset += len;
361    }
362    Ok(DecodedValue::Array(ranges))
363}
364
365/// Binary jsonb: a 1-byte format-version prefix (always `1` today) followed
366/// by the UTF-8 JSON text itself. Parsed into a `DecodedValue` tree (not
367/// left as an opaque string) so nested jsonb-backed named tuples decode
368/// the same way a composite field would.
369fn decode_jsonb(data: &[u8]) -> Result<DecodedValue> {
370    let text = std::str::from_utf8(&data[1..])?;
371    let value: serde_json::Value = serde_json::from_str(text)?;
372    Ok(json_to_cached(value))
373}
374
375fn json_to_cached(value: serde_json::Value) -> DecodedValue {
376    match value {
377        serde_json::Value::Null => DecodedValue::Null,
378        serde_json::Value::Bool(b) => DecodedValue::Bool(b),
379        serde_json::Value::Number(n) => {
380            if let Some(i) = n.as_i64() {
381                DecodedValue::I64(i)
382            } else {
383                DecodedValue::F64(n.as_f64().unwrap_or(f64::NAN))
384            }
385        }
386        serde_json::Value::String(s) => DecodedValue::Str(s),
387        serde_json::Value::Array(items) => DecodedValue::Array(items.into_iter().map(json_to_cached).collect()),
388        serde_json::Value::Object(map) => {
389            DecodedValue::Object(map.into_iter().map(|(k, v)| (k, json_to_cached(v))).collect())
390        }
391    }
392}
393
394/// `pgvector`'s binary format: `u16 ndim`, `u16 reserved` (always 0), then
395/// `ndim` big-endian `f32`s. Matches `_decode_vector_binary` exactly.
396fn decode_vector(data: &[u8]) -> Result<Vec<DecodedValue>> {
397    let ndim = u16::from_be_bytes(data[0..2].try_into()?) as usize;
398    let mut values = Vec::with_capacity(ndim);
399    for i in 0..ndim {
400        let start = 4 + i * 4;
401        let f = f32::from_be_bytes(data[start..start + 4].try_into()?);
402        values.push(DecodedValue::F64(f as f64));
403    }
404    Ok(values)
405}
406
407/// Encodes `items` (each expected to be `DecodedValue::F64`/`I64`) as
408/// pgvector's binary format — the inverse of `decode_vector`.
409fn encode_vector(items: &[DecodedValue], out: &mut bytes::BytesMut) -> Result<()> {
410    let ndim: u16 = items
411        .len()
412        .try_into()
413        .map_err(|_| Error::message("vector has too many dimensions to encode"))?;
414    out.put_u16(ndim);
415    out.put_u16(0); // reserved
416    for item in items {
417        let f = match item {
418            DecodedValue::F64(f) => *f as f32,
419            DecodedValue::I64(i) => *i as f32,
420            other => return Err(Error::message(format!("cannot encode {other:?} as a vector element"))),
421        };
422        out.put_f32(f);
423    }
424    Ok(())
425}
426
427/// Decodes a `record`-typed field: `i32 nfields`, then per field `u32
428/// type_oid` + `i32 field_len` (`-1` = NULL) + `field_len` bytes.
429fn decode_record(data: &[u8], ext: &ExtensionOids) -> Result<DecodedValue> {
430    let mut offset = 0usize;
431    let nfields = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
432    offset += 4;
433    let mut fields = Vec::with_capacity(nfields);
434    for _ in 0..nfields {
435        let type_oid = u32::from_be_bytes(data[offset..offset + 4].try_into()?);
436        offset += 4;
437        let field_len = i32::from_be_bytes(data[offset..offset + 4].try_into()?);
438        offset += 4;
439        if field_len == -1 {
440            fields.push(DecodedValue::Null);
441        } else {
442            let len = field_len as usize;
443            fields.push(decode_value(type_oid, &data[offset..offset + len], ext)?);
444            offset += len;
445        }
446    }
447    Ok(DecodedValue::Composite(fields))
448}
449
450/// Decodes any 1-dimensional array: `i32 ndim`, `i32 has_null_flag`, `u32
451/// element_oid`, one `(i32 dim_size, i32 lower_bound)` pair, then per
452/// element `i32 len` (`-1` = NULL) + `len` bytes. An empty array
453/// (`ndim == 0`) has no dimension pair to read.
454fn decode_array(data: &[u8], ext: &ExtensionOids) -> Result<DecodedValue> {
455    let mut offset = 0usize;
456    let ndim = i32::from_be_bytes(data[offset..offset + 4].try_into()?);
457    offset += 4;
458    offset += 4; // has-null flag — not needed, NULL is signaled per-element via len == -1
459    let element_oid = u32::from_be_bytes(data[offset..offset + 4].try_into()?);
460    offset += 4;
461    if ndim == 0 {
462        return Ok(DecodedValue::Array(vec![]));
463    }
464    let dim_size = i32::from_be_bytes(data[offset..offset + 4].try_into()?) as usize;
465    offset += 4;
466    offset += 4; // lower bound — Pylon arrays are always 1-based, not needed
467
468    let mut items = Vec::with_capacity(dim_size);
469    for _ in 0..dim_size {
470        let elem_len = i32::from_be_bytes(data[offset..offset + 4].try_into()?);
471        offset += 4;
472        if elem_len == -1 {
473            items.push(DecodedValue::Null);
474        } else {
475            let len = elem_len as usize;
476            items.push(decode_value(element_oid, &data[offset..offset + len], ext)?);
477            offset += len;
478        }
479    }
480    Ok(DecodedValue::Array(items))
481}
482
483// ── Parameter encoding (the inverse direction: DecodedValue -> wire bytes) ──
484//
485// Bound query parameters don't need pylon-core to supply explicit
486// per-parameter Postgres types up front: `Client::prepare` already asks
487// Postgres itself to analyze the SQL and report each `$1, $2, ...`'s
488// expected `Type` back (`Statement::params()`) — which drives
489// own extended-query-protocol binding already relies on today, just
490// encoding explicitly here instead of hiding it inside a codec
491// registry. So encoding is *type-directed*: given a `DecodedValue` and the
492// `Type` Postgres reported for that position, write the matching binary
493// representation. See `BoundParam` (in `lib.rs`) for the `ToSql` glue that
494// makes this pluggable into `tokio_postgres::Client::query`.
495
496use bytes::BufMut;
497use postgres_types::{IsNull, Kind, Type};
498
499/// Whether a string's own bytes are `ty`'s binary wire format.
500///
501/// True for the text-like types and for anything Postgres transmits as text:
502/// an enum travels as its label (which is what `decode_value` reads back), and
503/// `json` -- unlike `jsonb`, which carries a version byte first -- is its
504/// document verbatim. `uuid`, `jsonb` and `numeric` accept a string too, but by
505/// conversion rather than by copying, so they are handled at their own arms.
506/// `bytea` is here because a string bound to it has always meant those bytes.
507fn accepts_text_bytes(ty: &Type) -> bool {
508    match ty.kind() {
509        Kind::Enum(_) => true,
510        // A domain is a constrained alias, so it takes whatever its base does.
511        Kind::Domain(base) => accepts_text_bytes(base),
512        _ => matches!(
513            *ty,
514            Type::TEXT | Type::VARCHAR | Type::BPCHAR | Type::NAME | Type::JSON | Type::BYTEA | Type::UNKNOWN
515        ),
516    }
517}
518
519/// Encodes `value` as `ty`'s binary wire format into `out`. `ty` comes from
520/// `Statement::params()[i]` — Postgres's own analysis of the prepared SQL,
521/// not a guess — so this only needs to pick the right byte width/shape for
522/// whatever `DecodedValue` variant is actually being sent, not infer the
523/// target type itself.
524pub fn encode_value(value: &DecodedValue, ty: &Type, out: &mut bytes::BytesMut) -> Result<IsNull> {
525    let DecodedValue::Null = value else {
526        return encode_non_null(value, ty, out);
527    };
528    Ok(IsNull::Yes)
529}
530
531fn encode_non_null(value: &DecodedValue, ty: &Type, out: &mut bytes::BytesMut) -> Result<IsNull> {
532    // A `<std::decimal>$pN` cast makes Postgres report that parameter's
533    // type as `numeric` regardless of which `DecodedValue` variant the JSON
534    // request body produced (`I64`/`F64` for a JSON number, `Str` for a
535    // JSON string — `json_to_cached_value` in pylon-server has no visibility
536    // into the target PG type at parse time). Without this, the arms below
537    // write raw int8/float8 bytes or raw UTF-8 text straight into a
538    // numeric-typed slot, which Postgres's binary numeric decoder then reads
539    // as a corrupt header — "invalid sign in external representation" for
540    // I64/F64 (garbage sign field), "insufficient data left in message" for
541    // Str (too few bytes for the header). Route every numeric-ish variant
542    // through the same `numeric` encoding the `Decimal` arm below uses.
543    if *ty == Type::NUMERIC {
544        let text = match value {
545            DecodedValue::Decimal(s) | DecodedValue::Str(s) => s.clone(),
546            DecodedValue::I64(i) => i.to_string(),
547            // `{f}` is the shortest form that reads back as the same f64, so
548            // 0.1 binds as `0.1` rather than the full binary expansion.
549            DecodedValue::F64(f) => format!("{f}"),
550            _ => return Err(Error::message("cannot bind this value as a numeric parameter")),
551        };
552        numeric::encode(&text, out)?;
553        return Ok(IsNull::No);
554    }
555    match value {
556        DecodedValue::Null => unreachable!("caller already handled NULL"),
557        DecodedValue::Bool(b) => out.put_u8(*b as u8),
558        DecodedValue::I64(i) => {
559            if *ty == Type::INT2 {
560                out.put_i16(*i as i16);
561            } else if *ty == Type::INT4 {
562                out.put_i32(*i as i32);
563            } else {
564                out.put_i64(*i);
565            }
566        }
567        DecodedValue::F64(f) => {
568            if *ty == Type::FLOAT4 {
569                out.put_f32(*f as f32);
570            } else {
571                out.put_f64(*f);
572            }
573        }
574        DecodedValue::Str(s) => {
575            if *ty == Type::UUID {
576                // A JSON API request body necessarily carries a UUID query
577                // parameter as plain text (there's no JSON "uuid" type), so
578                // it arrives here as a `DecodedValue::Str`, not `::Uuid` —
579                // a `uuid` codec conventionally accepts a plain string the
580                // same way. Without this, the raw UTF-8 text bytes get sent
581                // for a binary-format `uuid` parameter, which Postgres
582                // rejects with "incorrect binary data format".
583                out.put_slice(&parse_uuid_str(s)?);
584            } else if *ty == Type::JSONB {
585                // A caller that already has serialized JSON text (e.g.
586                // `schema_to_db_state_json`'s output, bound as `$1::jsonb`
587                // in `migration apply`'s db_state snapshot update) arrives
588                // here as `DecodedValue::Str`, not `::Object` — treat it as
589                // already-valid JSON text and just add jsonb's binary
590                // version-byte prefix (see `decode_jsonb`/the `Object` arm
591                // below), rather than writing raw text bytes with no
592                // framing, which Postgres would reject.
593                out.put_u8(1);
594                out.put_slice(s.as_bytes());
595            } else if accepts_text_bytes(ty) {
596                out.put_slice(s.as_bytes());
597            } else {
598                // Without this the raw UTF-8 went into a binary-format slot of
599                // whatever type the parameter actually has, and Postgres reported
600                // its own reading of the bytes -- "insufficient data left in
601                // message" for an int8 or interval, "incorrect binary data
602                // format" for a bool or timestamptz. That came from the server,
603                // so it also left an open transaction aborted; refusing here
604                // names the type that was wanted and sends nothing.
605                return Err(Error::message(format!(
606                    "cannot bind a string as a parameter of type {:?}",
607                    ty.name()
608                )));
609            }
610        }
611        DecodedValue::Bytes(b) => out.put_slice(b),
612        DecodedValue::Uuid(bytes) => out.put_slice(bytes),
613        DecodedValue::Decimal(s) => numeric::encode(s, out)?,
614        DecodedValue::Array(items) if ty.name() == "vector" => {
615            // `$n::vector` casts the parameter directly (unlike
616            // `vector::search`'s `$n::float8[]::vector`, where the *inner*
617            // cast is what Postgres's prepare step reports as the param's
618            // type) — Postgres reports `$n` itself as `vector`, a scalar
619            // extension type, not `Kind::Array`. Bypass the generic array
620            // path entirely and write pgvector's own binary format.
621            encode_vector(items, out)?;
622        }
623        DecodedValue::Array(items) => {
624            let element_ty = match ty.kind() {
625                Kind::Array(inner) => inner.clone(),
626                // Not actually an array type per Postgres's own analysis —
627                // fall back to TEXT so encoding still proceeds deterministically
628                // rather than panicking; a real mismatch surfaces as a
629                // Postgres-side type error on execute, same as today.
630                _ => Type::TEXT,
631            };
632            encode_array(items, &element_ty, out)?;
633        }
634        DecodedValue::Composite(_) => {
635            // Composites only ever arise from *decoding* a query result
636            // (see `decode_record`) — PyQL never binds a raw composite as
637            // a query parameter, and encoding one correctly would need a
638            // per-field Postgres type that isn't available here (only the
639            // original compiled query's shape carries that). Erroring is
640            // safer than guessing wrong field types.
641            return Err(Error::message("cannot bind a composite value as a query parameter"));
642        }
643        DecodedValue::Object(fields) => {
644            let json = cached_object_to_json(fields);
645            out.put_u8(1); // jsonb binary format version prefix
646            out.put_slice(json.to_string().as_bytes());
647        }
648        DecodedValue::Interval {
649            months,
650            days,
651            microseconds,
652        } => {
653            // Same field order as `decode_interval`'s read.
654            out.put_i64(*microseconds);
655            out.put_i32(*days);
656            out.put_i32(*months);
657        }
658        DecodedValue::Date(days) => out.put_i32(*days),
659        DecodedValue::Time(us) => out.put_i64(*us),
660        DecodedValue::Timestamp(us) => out.put_i64(*us),
661        DecodedValue::Timestamptz(us) => out.put_i64(*us),
662        DecodedValue::Range {
663            lower,
664            upper,
665            inc_lower,
666            inc_upper,
667            empty,
668        } => {
669            if *empty {
670                out.put_u8(RANGE_EMPTY);
671                return Ok(IsNull::No);
672            }
673            let element_ty = match ty.kind() {
674                Kind::Range(inner) => inner.clone(),
675                // Not actually a range type per Postgres's own analysis —
676                // fall back to TEXT so encoding proceeds deterministically;
677                // a real mismatch surfaces as a Postgres-side error, same
678                // as the analogous fallback in the `Array` arm above.
679                _ => Type::TEXT,
680            };
681            let mut flags = 0u8;
682            if *inc_lower {
683                flags |= RANGE_LB_INC;
684            }
685            if *inc_upper {
686                flags |= RANGE_UB_INC;
687            }
688            if lower.is_none() {
689                flags |= RANGE_LB_INF;
690            }
691            if upper.is_none() {
692                flags |= RANGE_UB_INF;
693            }
694            out.put_u8(flags);
695            for bound in [lower, upper].into_iter().flatten() {
696                let mut buf = bytes::BytesMut::new();
697                encode_value(bound, &element_ty, &mut buf)?;
698                out.put_i32(buf.len() as i32);
699                out.put_slice(&buf);
700            }
701        }
702    }
703    Ok(IsNull::No)
704}
705
706fn cached_object_to_json(fields: &[(String, DecodedValue)]) -> serde_json::Value {
707    serde_json::Value::Object(fields.iter().map(|(k, v)| (k.clone(), cached_to_json(v))).collect())
708}
709
710fn cached_to_json(value: &DecodedValue) -> serde_json::Value {
711    match value {
712        DecodedValue::Null => serde_json::Value::Null,
713        DecodedValue::Bool(b) => serde_json::Value::Bool(*b),
714        DecodedValue::I64(i) => serde_json::Value::Number((*i).into()),
715        DecodedValue::F64(f) => serde_json::Number::from_f64(*f)
716            .map(serde_json::Value::Number)
717            .unwrap_or(serde_json::Value::Null),
718        DecodedValue::Str(s) => serde_json::Value::String(s.clone()),
719        DecodedValue::Bytes(b) => serde_json::Value::String(hex::encode(b)),
720        DecodedValue::Uuid(bytes) => serde_json::Value::String(format_uuid(bytes)),
721        DecodedValue::Decimal(s) => serde_json::Value::String(s.clone()),
722        DecodedValue::Array(items) | DecodedValue::Composite(items) => {
723            serde_json::Value::Array(items.iter().map(cached_to_json).collect())
724        }
725        DecodedValue::Object(fields) => cached_object_to_json(fields),
726        // No natural JSON scalar for an interval; only reachable if an
727        // Interval value ends up nested inside an Object being sent as a
728        // jsonb parameter — represented as its raw components so it's at
729        // least round-trippable, not silently dropped.
730        DecodedValue::Interval {
731            months,
732            days,
733            microseconds,
734        } => serde_json::json!({
735            "months": months, "days": days, "microseconds": microseconds,
736        }),
737        // Same rationale as Interval above — raw PG wire units, not a
738        // formatted calendar string (calendar math is deliberately left to
739        // Python's own `datetime` module at the `pgvalue.rs` boundary, not
740        // reimplemented here).
741        DecodedValue::Date(days) => serde_json::json!({ "days_since_2000_01_01": days }),
742        DecodedValue::Time(us) => serde_json::json!({ "microseconds_since_midnight": us }),
743        DecodedValue::Timestamp(us) => serde_json::json!({ "microseconds_since_2000_01_01": us }),
744        DecodedValue::Timestamptz(us) => serde_json::json!({ "microseconds_since_2000_01_01_utc": us }),
745        DecodedValue::Range {
746            lower,
747            upper,
748            inc_lower,
749            inc_upper,
750            empty,
751        } => serde_json::json!({
752            "lower": lower.as_deref().map(cached_to_json),
753            "upper": upper.as_deref().map(cached_to_json),
754            "inc_lower": inc_lower,
755            "inc_upper": inc_upper,
756            "empty": empty,
757        }),
758    }
759}
760
761fn format_uuid(bytes: &[u8; 16]) -> String {
762    let hex = hex::encode(bytes);
763    format!(
764        "{}-{}-{}-{}-{}",
765        &hex[0..8],
766        &hex[8..12],
767        &hex[12..16],
768        &hex[16..20],
769        &hex[20..32]
770    )
771}
772
773/// Parses a hyphenated UUID string into its 16 raw bytes — the inverse of
774/// `format_uuid`. Tolerates the hyphens being anywhere/absent (just strips
775/// every `-` and hex-decodes what's left) rather than validating the exact
776/// `8-4-4-4-12` grouping, since the only thing that matters here is
777/// recovering the right 16 bytes, not rejecting non-canonical formatting.
778fn parse_uuid_str(s: &str) -> Result<[u8; 16]> {
779    let hex_only: String = s.chars().filter(|c| *c != '-').collect();
780    let bytes = hex::decode(&hex_only).map_err(|_| Error::message(format!("invalid UUID string: {s:?}")))?;
781    bytes
782        .try_into()
783        .map_err(|_: Vec<u8>| Error::message(format!("invalid UUID string: {s:?}")))
784}
785
786/// 1-dimensional Postgres array binary format (see the module doc comment
787/// for the layout) — the encode-side mirror of `decode_array`.
788fn encode_array(items: &[DecodedValue], element_ty: &Type, out: &mut bytes::BytesMut) -> Result<()> {
789    if items.is_empty() {
790        out.put_i32(0); // ndim
791        out.put_i32(0); // has-null flag
792        out.put_u32(element_ty.oid());
793        return Ok(());
794    }
795    let has_null = items.iter().any(|v| matches!(v, DecodedValue::Null));
796    out.put_i32(1); // ndim — Pylon arrays are always 1-D
797    out.put_i32(has_null as i32);
798    out.put_u32(element_ty.oid());
799    out.put_i32(items.len() as i32); // dim size
800    out.put_i32(1); // lower bound
801
802    for item in items {
803        if matches!(item, DecodedValue::Null) {
804            out.put_i32(-1);
805            continue;
806        }
807        let start = out.len();
808        out.put_i32(0); // placeholder length, patched below
809        let is_null = encode_value(item, element_ty, out)?;
810        let len = (out.len() - start - 4) as i32;
811        let len = if matches!(is_null, IsNull::Yes) { -1 } else { len };
812        out[start..start + 4].copy_from_slice(&len.to_be_bytes());
813    }
814    Ok(())
815}
816
817#[cfg(test)]
818mod tests {
819    use super::*;
820
821    fn no_ext() -> ExtensionOids {
822        ExtensionOids::default()
823    }
824
825    #[test]
826    fn decodes_bool() {
827        assert_eq!(
828            decode_value(OID_BOOL, &[1], &no_ext()).unwrap(),
829            DecodedValue::Bool(true)
830        );
831        assert_eq!(
832            decode_value(OID_BOOL, &[0], &no_ext()).unwrap(),
833            DecodedValue::Bool(false)
834        );
835    }
836
837    #[test]
838    fn decodes_integers() {
839        assert_eq!(
840            decode_value(OID_INT2, &7i16.to_be_bytes(), &no_ext()).unwrap(),
841            DecodedValue::I64(7)
842        );
843        assert_eq!(
844            decode_value(OID_INT4, &(-42i32).to_be_bytes(), &no_ext()).unwrap(),
845            DecodedValue::I64(-42)
846        );
847        assert_eq!(
848            decode_value(OID_INT8, &9_223_372_036_854_775_807i64.to_be_bytes(), &no_ext()).unwrap(),
849            DecodedValue::I64(9_223_372_036_854_775_807)
850        );
851    }
852
853    #[test]
854    fn decodes_floats() {
855        assert_eq!(
856            decode_value(OID_FLOAT4, &1.5f32.to_be_bytes(), &no_ext()).unwrap(),
857            DecodedValue::F64(1.5)
858        );
859        assert_eq!(
860            decode_value(OID_FLOAT8, &2.25f64.to_be_bytes(), &no_ext()).unwrap(),
861            DecodedValue::F64(2.25)
862        );
863    }
864
865    #[test]
866    fn decodes_text_varchar_bpchar() {
867        for oid in [OID_TEXT, OID_VARCHAR, OID_BPCHAR] {
868            assert_eq!(
869                decode_value(oid, "hello".as_bytes(), &no_ext()).unwrap(),
870                DecodedValue::Str("hello".to_string())
871            );
872        }
873    }
874
875    #[test]
876    fn decodes_unicode_text() {
877        assert_eq!(
878            decode_value(OID_TEXT, "héllo wörld 🎉".as_bytes(), &no_ext()).unwrap(),
879            DecodedValue::Str("héllo wörld 🎉".to_string())
880        );
881    }
882
883    #[test]
884    fn decodes_uuid() {
885        let bytes: [u8; 16] = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16];
886        assert_eq!(
887            decode_value(OID_UUID, &bytes, &no_ext()).unwrap(),
888            DecodedValue::Uuid(bytes)
889        );
890    }
891
892    #[test]
893    fn decodes_bytea() {
894        assert_eq!(
895            decode_value(OID_BYTEA, &[1, 2, 3, 255], &no_ext()).unwrap(),
896            DecodedValue::Bytes(vec![1, 2, 3, 255])
897        );
898    }
899
900    #[test]
901    fn decodes_interval() {
902        // Regression: interval has no dedicated binary decoder — it used to
903        // fall through to the UTF-8-text fallback, which panics/errors on
904        // interval's actual binary payload (microseconds/days/months, not text).
905        let mut data = Vec::new();
906        data.extend_from_slice(&3_600_000_000i64.to_be_bytes()); // 1 hour, in microseconds
907        data.extend_from_slice(&2i32.to_be_bytes()); // 2 days
908        data.extend_from_slice(&1i32.to_be_bytes()); // 1 month
909        assert_eq!(
910            decode_value(OID_INTERVAL, &data, &no_ext()).unwrap(),
911            DecodedValue::Interval {
912                months: 1,
913                days: 2,
914                microseconds: 3_600_000_000
915            }
916        );
917    }
918
919    #[test]
920    fn encodes_interval() {
921        let value = DecodedValue::Interval {
922            months: 1,
923            days: 2,
924            microseconds: 3_600_000_000,
925        };
926        let mut out = bytes::BytesMut::new();
927        encode_value(&value, &postgres_types::Type::INTERVAL, &mut out).unwrap();
928        assert_eq!(decode_value(OID_INTERVAL, &out, &no_ext()).unwrap(), value);
929    }
930
931    #[test]
932    fn decodes_date_time_timestamp_timestamptz() {
933        // Regression: these had no binary decoder either — date silently
934        // returned garbage bytes (never even errored), the others panicked
935        // on the same UTF-8-text-fallback assumption interval did.
936        assert_eq!(
937            decode_value(OID_DATE, &9525i32.to_be_bytes(), &no_ext()).unwrap(),
938            DecodedValue::Date(9525)
939        );
940        assert_eq!(
941            decode_value(OID_TIME, &3_600_000_000i64.to_be_bytes(), &no_ext()).unwrap(),
942            DecodedValue::Time(3_600_000_000)
943        );
944        assert_eq!(
945            decode_value(OID_TIMESTAMP, &1_000_000_000i64.to_be_bytes(), &no_ext()).unwrap(),
946            DecodedValue::Timestamp(1_000_000_000)
947        );
948        assert_eq!(
949            decode_value(OID_TIMESTAMPTZ, &1_000_000_000i64.to_be_bytes(), &no_ext()).unwrap(),
950            DecodedValue::Timestamptz(1_000_000_000)
951        );
952    }
953
954    #[test]
955    fn encodes_date_time_timestamp_timestamptz() {
956        for (value, ty) in [
957            (DecodedValue::Date(9525), postgres_types::Type::DATE),
958            (DecodedValue::Time(3_600_000_000), postgres_types::Type::TIME),
959            (DecodedValue::Timestamp(1_000_000_000), postgres_types::Type::TIMESTAMP),
960            (
961                DecodedValue::Timestamptz(1_000_000_000),
962                postgres_types::Type::TIMESTAMPTZ,
963            ),
964        ] {
965            let mut out = bytes::BytesMut::new();
966            encode_value(&value, &ty, &mut out).unwrap();
967            let oid = match &value {
968                DecodedValue::Date(_) => OID_DATE,
969                DecodedValue::Time(_) => OID_TIME,
970                DecodedValue::Timestamp(_) => OID_TIMESTAMP,
971                DecodedValue::Timestamptz(_) => OID_TIMESTAMPTZ,
972                _ => unreachable!(),
973            };
974            assert_eq!(decode_value(oid, &out, &no_ext()).unwrap(), value);
975        }
976    }
977
978    #[test]
979    fn decodes_a_bounded_int8range() {
980        // flags = LB_INC | UB_INC-off = 0x02 (inclusive lower, exclusive upper)
981        let mut data = vec![RANGE_LB_INC];
982        data.extend_from_slice(&8i32.to_be_bytes());
983        data.extend_from_slice(&1i64.to_be_bytes());
984        data.extend_from_slice(&8i32.to_be_bytes());
985        data.extend_from_slice(&10i64.to_be_bytes());
986        assert_eq!(
987            decode_value(OID_INT8RANGE, &data, &no_ext()).unwrap(),
988            DecodedValue::Range {
989                lower: Some(Box::new(DecodedValue::I64(1))),
990                upper: Some(Box::new(DecodedValue::I64(10))),
991                inc_lower: true,
992                inc_upper: false,
993                empty: false,
994            }
995        );
996    }
997
998    #[test]
999    fn decodes_an_empty_range() {
1000        assert_eq!(
1001            decode_value(OID_INT8RANGE, &[RANGE_EMPTY], &no_ext()).unwrap(),
1002            DecodedValue::Range {
1003                lower: None,
1004                upper: None,
1005                inc_lower: false,
1006                inc_upper: false,
1007                empty: true
1008            }
1009        );
1010    }
1011
1012    #[test]
1013    fn decodes_an_unbounded_range() {
1014        // Both bounds infinite: flags = LB_INF | UB_INF, no bound payloads follow.
1015        let data = [RANGE_LB_INF | RANGE_UB_INF];
1016        assert_eq!(
1017            decode_value(OID_INT8RANGE, &data, &no_ext()).unwrap(),
1018            DecodedValue::Range {
1019                lower: None,
1020                upper: None,
1021                inc_lower: false,
1022                inc_upper: false,
1023                empty: false
1024            }
1025        );
1026    }
1027
1028    #[test]
1029    fn encodes_and_round_trips_an_int8range() {
1030        let value = DecodedValue::Range {
1031            lower: Some(Box::new(DecodedValue::I64(1))),
1032            upper: Some(Box::new(DecodedValue::I64(10))),
1033            inc_lower: true,
1034            inc_upper: false,
1035            empty: false,
1036        };
1037        let mut out = bytes::BytesMut::new();
1038        encode_value(&value, &postgres_types::Type::INT8_RANGE, &mut out).unwrap();
1039        assert_eq!(decode_value(OID_INT8RANGE, &out, &no_ext()).unwrap(), value);
1040    }
1041
1042    #[test]
1043    fn decodes_a_multirange_of_int8ranges() {
1044        let mut range1 = vec![RANGE_LB_INC];
1045        range1.extend_from_slice(&8i32.to_be_bytes());
1046        range1.extend_from_slice(&1i64.to_be_bytes());
1047        range1.extend_from_slice(&8i32.to_be_bytes());
1048        range1.extend_from_slice(&3i64.to_be_bytes());
1049
1050        let mut range2 = vec![RANGE_LB_INC];
1051        range2.extend_from_slice(&8i32.to_be_bytes());
1052        range2.extend_from_slice(&5i64.to_be_bytes());
1053        range2.extend_from_slice(&8i32.to_be_bytes());
1054        range2.extend_from_slice(&7i64.to_be_bytes());
1055
1056        let mut data = 2i32.to_be_bytes().to_vec();
1057        data.extend_from_slice(&(range1.len() as i32).to_be_bytes());
1058        data.extend_from_slice(&range1);
1059        data.extend_from_slice(&(range2.len() as i32).to_be_bytes());
1060        data.extend_from_slice(&range2);
1061
1062        let decoded = decode_value(OID_INT8MULTIRANGE, &data, &no_ext()).unwrap();
1063        assert_eq!(
1064            decoded,
1065            DecodedValue::Array(vec![
1066                DecodedValue::Range {
1067                    lower: Some(Box::new(DecodedValue::I64(1))),
1068                    upper: Some(Box::new(DecodedValue::I64(3))),
1069                    inc_lower: true,
1070                    inc_upper: false,
1071                    empty: false,
1072                },
1073                DecodedValue::Range {
1074                    lower: Some(Box::new(DecodedValue::I64(5))),
1075                    upper: Some(Box::new(DecodedValue::I64(7))),
1076                    inc_lower: true,
1077                    inc_upper: false,
1078                    empty: false,
1079                },
1080            ])
1081        );
1082    }
1083
1084    #[test]
1085    fn errors_on_an_oid_no_rule_or_discovery_covers() {
1086        // The old behaviour here was to decode the raw binary as UTF-8 text,
1087        // which silently produced mojibake for every non-text type whose OID
1088        // isn't a builtin (`vector` above all). An unclassified OID is now an
1089        // error rather than a guess.
1090        let err = decode_value(999_999, &[0xff, 0xfe], &no_ext()).unwrap_err();
1091        assert!(matches!(err, Error::UnknownTypeOid { oid: 999_999 }), "got {err:?}");
1092    }
1093
1094    #[test]
1095    fn decodes_a_discovered_enum_oid_as_its_label_text() {
1096        let ext = ExtensionOids {
1097            enums: std::collections::HashSet::from([50_001]),
1098            ..Default::default()
1099        };
1100        assert_eq!(
1101            decode_value(50_001, "Active".as_bytes(), &ext).unwrap(),
1102            DecodedValue::Str("Active".to_string())
1103        );
1104    }
1105
1106    #[test]
1107    fn decodes_a_discovered_domain_through_its_base_type() {
1108        // A domain over int8 has to decode as int8, not as text — the old
1109        // fallback would have run UTF-8 validation over these eight bytes.
1110        let ext = ExtensionOids {
1111            domains: std::collections::HashMap::from([(50_002, OID_INT8)]),
1112            ..Default::default()
1113        };
1114        assert_eq!(
1115            decode_value(50_002, &7i64.to_be_bytes(), &ext).unwrap(),
1116            DecodedValue::I64(7)
1117        );
1118    }
1119
1120    #[test]
1121    fn discovery_rows_populate_vector_enums_and_domains() {
1122        let ext = ExtensionOids::from_discovery_rows([
1123            (50_000, "b".to_string(), 0, "vector".to_string()),
1124            (50_001, "e".to_string(), 0, "status".to_string()),
1125            (50_002, "d".to_string(), OID_INT8, "positive_int".to_string()),
1126            // A domain with no resolvable base is skipped rather than
1127            // recorded as pointing at OID 0.
1128            (50_003, "d".to_string(), 0, "broken".to_string()),
1129        ]);
1130        assert_eq!(ext.vector, Some(50_000));
1131        assert!(ext.enums.contains(&50_001));
1132        assert_eq!(ext.domains.get(&50_002), Some(&OID_INT8));
1133        assert!(!ext.domains.contains_key(&50_003));
1134    }
1135
1136    #[test]
1137    fn discovery_rows_record_array_types() {
1138        let ext = ExtensionOids::from_discovery_rows([
1139            (50_001, "e".to_string(), 0, "status".to_string()),
1140            (50_010, "A".to_string(), 50_001, "_status".to_string()),
1141        ]);
1142        assert!(ext.enums.contains(&50_001));
1143        assert!(ext.arrays.contains(&50_010));
1144    }
1145
1146    #[test]
1147    fn decodes_an_array_of_a_discovered_enum() {
1148        // An `AuthenticationMethod[]` column: the array's own OID is
1149        // database-assigned, so without discovery it is an UnknownTypeOid
1150        // even though the element OID travels in the array header and every
1151        // label in it decodes.
1152        let ext = ExtensionOids {
1153            enums: std::collections::HashSet::from([50_001]),
1154            arrays: std::collections::HashSet::from([50_010]),
1155            ..Default::default()
1156        };
1157        let encoded = encode_array(50_001, &[Some(b"Password"), Some(b"Passkey")]);
1158        assert_eq!(
1159            decode_value(50_010, &encoded, &ext).unwrap(),
1160            DecodedValue::Array(vec![
1161                DecodedValue::Str("Password".to_string()),
1162                DecodedValue::Str("Passkey".to_string()),
1163            ])
1164        );
1165    }
1166
1167    #[test]
1168    fn an_array_of_an_undiscovered_type_is_still_an_error() {
1169        let ext = ExtensionOids::default();
1170        let encoded = encode_array(50_001, &[Some(b"Password")]);
1171        let err = decode_value(50_010, &encoded, &ext).unwrap_err();
1172        assert!(matches!(err, Error::UnknownTypeOid { oid: 50_010 }), "got {err:?}");
1173    }
1174
1175    #[test]
1176    fn a_vector_inside_a_record_decodes_when_discovery_ran() {
1177        // The real C-4 shape: pylon-core wraps every result in `SELECT (...)
1178        // AS result`, so a vector value arrives nested in a record and is
1179        // decoded through `decode_record`'s own per-field OID dispatch.
1180        let mut vec_bytes = 2u16.to_be_bytes().to_vec();
1181        vec_bytes.extend_from_slice(&0u16.to_be_bytes());
1182        vec_bytes.extend_from_slice(&1.5f32.to_be_bytes());
1183        vec_bytes.extend_from_slice(&2.5f32.to_be_bytes());
1184        let rec = encode_record(&[(OID_TEXT, Some(b"doc")), (50_000, Some(&vec_bytes))]);
1185
1186        let ext = ExtensionOids {
1187            vector: Some(50_000),
1188            ..Default::default()
1189        };
1190        let decoded = decode_value(OID_RECORD, &rec, &ext).unwrap();
1191        let DecodedValue::Composite(fields) = decoded else {
1192            panic!("expected Composite, got {decoded:?}")
1193        };
1194        assert_eq!(fields[0], DecodedValue::Str("doc".to_string()));
1195        assert_eq!(
1196            fields[1],
1197            DecodedValue::Array(vec![DecodedValue::F64(1.5), DecodedValue::F64(2.5)])
1198        );
1199
1200        // Without discovery the same bytes must fail loudly, not silently.
1201        let err = decode_value(OID_RECORD, &rec, &no_ext()).unwrap_err();
1202        assert!(matches!(err, Error::UnknownTypeOid { oid: 50_000 }), "got {err:?}");
1203    }
1204
1205    /// Builds the Postgres binary `numeric` wire format by hand: `u16
1206    /// ndigits`, `i16 weight`, `u16 sign`, `i16 dscale`, then `ndigits`
1207    /// base-10000 digit groups (each a `u16`, matching `NBASE = 10000`).
1208    fn encode_numeric(sign: u16, weight: i16, dscale: i16, digits: &[u16]) -> Vec<u8> {
1209        let mut buf = Vec::new();
1210        buf.extend_from_slice(&(digits.len() as u16).to_be_bytes());
1211        buf.extend_from_slice(&weight.to_be_bytes());
1212        buf.extend_from_slice(&sign.to_be_bytes());
1213        buf.extend_from_slice(&dscale.to_be_bytes());
1214        for d in digits {
1215            buf.extend_from_slice(&d.to_be_bytes());
1216        }
1217        buf
1218    }
1219
1220    #[test]
1221    fn decodes_numeric_integer() {
1222        // 12345 = digit groups [1, 2345] at weight 1 (10000^1 * 1 + 10000^0 * 2345)
1223        let data = encode_numeric(0x0000, 1, 0, &[1, 2345]);
1224        assert_eq!(
1225            decode_value(OID_NUMERIC, &data, &no_ext()).unwrap(),
1226            DecodedValue::Decimal("12345".to_string())
1227        );
1228    }
1229
1230    #[test]
1231    fn decodes_numeric_with_fraction() {
1232        // 12.50, dscale=2: digit groups [12, 5000] at weight 0
1233        let data = encode_numeric(0x0000, 0, 2, &[12, 5000]);
1234        assert_eq!(
1235            decode_value(OID_NUMERIC, &data, &no_ext()).unwrap(),
1236            DecodedValue::Decimal("12.50".to_string())
1237        );
1238    }
1239
1240    #[test]
1241    fn decodes_negative_numeric() {
1242        let data = encode_numeric(0x4000, 0, 2, &[12, 5000]);
1243        assert_eq!(
1244            decode_value(OID_NUMERIC, &data, &no_ext()).unwrap(),
1245            DecodedValue::Decimal("-12.50".to_string())
1246        );
1247    }
1248
1249    #[test]
1250    fn refuses_a_string_for_a_parameter_whose_binary_form_is_not_text() {
1251        // The raw UTF-8 used to go straight into the binary slot, leaving
1252        // Postgres to report its own reading of the bytes -- and aborting the
1253        // caller's transaction, because the complaint came from the server.
1254        for ty in [
1255            Type::INT8,
1256            Type::INT4,
1257            Type::BOOL,
1258            Type::INTERVAL,
1259            Type::TIMESTAMPTZ,
1260            Type::DATE,
1261        ] {
1262            let mut buffer = bytes::BytesMut::new();
1263            let Err(err) = encode_value(&DecodedValue::Str("25 days".into()), &ty, &mut buffer) else {
1264                panic!("{} must refuse a string", ty.name());
1265            };
1266            assert!(
1267                err.to_string().contains(ty.name()),
1268                "the message must name the type that was wanted: {err}"
1269            );
1270            assert!(buffer.is_empty(), "nothing may be written for a refused parameter");
1271        }
1272    }
1273
1274    #[test]
1275    fn a_string_still_reaches_the_types_it_is_the_wire_form_of() {
1276        for ty in [
1277            Type::TEXT,
1278            Type::VARCHAR,
1279            Type::BPCHAR,
1280            Type::NAME,
1281            Type::JSON,
1282            Type::BYTEA,
1283            Type::UNKNOWN,
1284        ] {
1285            let mut buffer = bytes::BytesMut::new();
1286            encode_value(&DecodedValue::Str("hello".into()), &ty, &mut buffer)
1287                .unwrap_or_else(|e| panic!("{} must take a string: {e}", ty.name()));
1288            assert_eq!(&buffer[..], b"hello", "{}", ty.name());
1289        }
1290    }
1291
1292    #[test]
1293    fn a_string_still_converts_for_uuid_jsonb_and_numeric() {
1294        let mut buffer = bytes::BytesMut::new();
1295        encode_value(
1296            &DecodedValue::Str("00000000-0000-0000-0000-000000000001".into()),
1297            &Type::UUID,
1298            &mut buffer,
1299        )
1300        .expect("uuid takes a string");
1301        assert_eq!(buffer.len(), 16, "a uuid is converted, not copied");
1302
1303        let mut buffer = bytes::BytesMut::new();
1304        encode_value(&DecodedValue::Str(r#"{"a":1}"#.into()), &Type::JSONB, &mut buffer).expect("jsonb takes a string");
1305        assert_eq!(buffer[0], 1, "jsonb needs its version byte");
1306
1307        let mut buffer = bytes::BytesMut::new();
1308        encode_value(&DecodedValue::Str("12.50".into()), &Type::NUMERIC, &mut buffer).expect("numeric takes a string");
1309        assert_eq!(numeric::decode(&buffer).unwrap(), "12.50");
1310    }
1311
1312    #[test]
1313    fn decodes_jsonb_object() {
1314        let mut data = vec![1u8]; // version prefix
1315        data.extend_from_slice(br#"{"a":1,"b":"two","c":[1,2,3],"d":null}"#);
1316        let decoded = decode_value(OID_JSONB, &data, &no_ext()).unwrap();
1317        assert_eq!(
1318            decoded,
1319            DecodedValue::Object(vec![
1320                ("a".into(), DecodedValue::I64(1)),
1321                ("b".into(), DecodedValue::Str("two".into())),
1322                (
1323                    "c".into(),
1324                    DecodedValue::Array(vec![DecodedValue::I64(1), DecodedValue::I64(2), DecodedValue::I64(3)])
1325                ),
1326                ("d".into(), DecodedValue::Null),
1327            ])
1328        );
1329    }
1330
1331    #[test]
1332    fn decodes_jsonb_scalar_and_array() {
1333        let mut data = vec![1u8];
1334        data.extend_from_slice(b"42");
1335        assert_eq!(
1336            decode_value(OID_JSONB, &data, &no_ext()).unwrap(),
1337            DecodedValue::I64(42)
1338        );
1339
1340        let mut data2 = vec![1u8];
1341        data2.extend_from_slice(b"[1.5, 2.5]");
1342        assert_eq!(
1343            decode_value(OID_JSONB, &data2, &no_ext()).unwrap(),
1344            DecodedValue::Array(vec![DecodedValue::F64(1.5), DecodedValue::F64(2.5)])
1345        );
1346    }
1347
1348    /// Builds a `record`-field's binary payload by hand: `i32 nfields`,
1349    /// then per field `u32 type_oid` + `i32 field_len` (`-1` = NULL) +
1350    /// bytes — mirroring exactly what `decode_record` reads.
1351    fn encode_record(fields: &[(u32, Option<&[u8]>)]) -> Vec<u8> {
1352        let mut buf = Vec::new();
1353        buf.extend_from_slice(&(fields.len() as i32).to_be_bytes());
1354        for (oid, data) in fields {
1355            buf.extend_from_slice(&oid.to_be_bytes());
1356            match data {
1357                None => buf.extend_from_slice(&(-1i32).to_be_bytes()),
1358                Some(bytes) => {
1359                    buf.extend_from_slice(&(bytes.len() as i32).to_be_bytes());
1360                    buf.extend_from_slice(bytes);
1361                }
1362            }
1363        }
1364        buf
1365    }
1366
1367    #[test]
1368    fn decodes_flat_record() {
1369        let data = encode_record(&[
1370            (OID_INT8, Some(&42i64.to_be_bytes())),
1371            (OID_TEXT, Some(b"alice")),
1372            (OID_BOOL, None),
1373        ]);
1374        let decoded = decode_value(OID_RECORD, &data, &no_ext()).unwrap();
1375        assert_eq!(
1376            decoded,
1377            DecodedValue::Composite(vec![
1378                DecodedValue::I64(42),
1379                DecodedValue::Str("alice".into()),
1380                DecodedValue::Null
1381            ])
1382        );
1383    }
1384
1385    #[test]
1386    fn decodes_nested_record() {
1387        let inner = encode_record(&[(OID_INT8, Some(&1i64.to_be_bytes()))]);
1388        let outer = encode_record(&[(OID_RECORD, Some(&inner)), (OID_TEXT, Some(b"outer"))]);
1389        let decoded = decode_value(OID_RECORD, &outer, &no_ext()).unwrap();
1390        assert_eq!(
1391            decoded,
1392            DecodedValue::Composite(vec![
1393                DecodedValue::Composite(vec![DecodedValue::I64(1)]),
1394                DecodedValue::Str("outer".into()),
1395            ])
1396        );
1397    }
1398
1399    /// Builds a Postgres array's binary payload by hand: `i32 ndim`, `i32
1400    /// has_null`, `u32 element_oid`, `(i32 dim, i32 lbound)`, then per
1401    /// element `i32 len` (`-1` = NULL) + bytes.
1402    fn encode_array(element_oid: u32, elements: &[Option<&[u8]>]) -> Vec<u8> {
1403        if elements.is_empty() {
1404            let mut buf = Vec::new();
1405            buf.extend_from_slice(&0i32.to_be_bytes());
1406            buf.extend_from_slice(&0i32.to_be_bytes());
1407            buf.extend_from_slice(&element_oid.to_be_bytes());
1408            return buf;
1409        }
1410        let mut buf = Vec::new();
1411        buf.extend_from_slice(&1i32.to_be_bytes());
1412        buf.extend_from_slice(&0i32.to_be_bytes());
1413        buf.extend_from_slice(&element_oid.to_be_bytes());
1414        buf.extend_from_slice(&(elements.len() as i32).to_be_bytes());
1415        buf.extend_from_slice(&1i32.to_be_bytes());
1416        for data in elements {
1417            match data {
1418                None => buf.extend_from_slice(&(-1i32).to_be_bytes()),
1419                Some(bytes) => {
1420                    buf.extend_from_slice(&(bytes.len() as i32).to_be_bytes());
1421                    buf.extend_from_slice(bytes);
1422                }
1423            }
1424        }
1425        buf
1426    }
1427
1428    #[test]
1429    fn decodes_array_of_scalars() {
1430        let data = encode_array(OID_TEXT, &[Some(b"a"), Some(b"b"), None]);
1431        let decoded = decode_value(OID_TEXT_ARRAY, &data, &no_ext()).unwrap();
1432        assert_eq!(
1433            decoded,
1434            DecodedValue::Array(vec![
1435                DecodedValue::Str("a".into()),
1436                DecodedValue::Str("b".into()),
1437                DecodedValue::Null
1438            ])
1439        );
1440    }
1441
1442    #[test]
1443    fn decodes_empty_array() {
1444        let data = encode_array(OID_TEXT, &[]);
1445        assert_eq!(
1446            decode_value(OID_TEXT_ARRAY, &data, &no_ext()).unwrap(),
1447            DecodedValue::Array(vec![])
1448        );
1449    }
1450
1451    #[test]
1452    fn decodes_array_of_records() {
1453        let rec1 = encode_record(&[(OID_INT8, Some(&1i64.to_be_bytes()))]);
1454        let rec2 = encode_record(&[(OID_INT8, Some(&2i64.to_be_bytes()))]);
1455        let data = encode_array(OID_RECORD, &[Some(&rec1), Some(&rec2)]);
1456        let decoded = decode_value(OID_RECORD_ARRAY, &data, &no_ext()).unwrap();
1457        assert_eq!(
1458            decoded,
1459            DecodedValue::Array(vec![
1460                DecodedValue::Composite(vec![DecodedValue::I64(1)]),
1461                DecodedValue::Composite(vec![DecodedValue::I64(2)]),
1462            ])
1463        );
1464    }
1465
1466    #[test]
1467    fn decodes_vector_when_extension_oid_known() {
1468        let mut data = 2u16.to_be_bytes().to_vec(); // ndim = 2
1469        data.extend_from_slice(&0u16.to_be_bytes()); // reserved
1470        data.extend_from_slice(&1.5f32.to_be_bytes());
1471        data.extend_from_slice(&2.5f32.to_be_bytes());
1472
1473        let ext = ExtensionOids {
1474            vector: Some(50_000),
1475            ..Default::default()
1476        };
1477        let decoded = decode_value(50_000, &data, &ext).unwrap();
1478        assert_eq!(
1479            decoded,
1480            DecodedValue::Array(vec![DecodedValue::F64(1.5), DecodedValue::F64(2.5)])
1481        );
1482    }
1483
1484    #[test]
1485    fn unknown_oid_without_vector_extension_is_an_error_not_a_guess() {
1486        // Same OID as the vector test above, but with nothing discovered —
1487        // it must neither be read as vector binary data nor blindly
1488        // stringified, since which of those is right is exactly what
1489        // discovery exists to establish.
1490        let err = decode_value(50_000, "some-domain-value".as_bytes(), &no_ext()).unwrap_err();
1491        assert!(matches!(err, Error::UnknownTypeOid { oid: 50_000 }), "got {err:?}");
1492    }
1493
1494    #[test]
1495    fn encodes_a_str_value_as_uuid_binary_when_the_target_type_is_uuid() {
1496        // Regression test: a JSON API request body carries a UUID query
1497        // parameter as plain text (there's no JSON "uuid" type), so it
1498        // arrives as `DecodedValue::Str` — binding it directly against a
1499        // `uuid`-typed parameter must produce the 16-byte binary form, not
1500        // the raw 36-character text bytes (which Postgres rejects with
1501        // "incorrect binary data format").
1502        let value = DecodedValue::Str("11111111-2222-3333-4444-555555555555".to_string());
1503        let mut out = bytes::BytesMut::new();
1504        encode_value(&value, &postgres_types::Type::UUID, &mut out).unwrap();
1505        assert_eq!(
1506            out.as_ref(),
1507            &[
1508                0x11, 0x11, 0x11, 0x11, 0x22, 0x22, 0x33, 0x33, 0x44, 0x44, 0x55, 0x55, 0x55, 0x55, 0x55, 0x55
1509            ]
1510        );
1511    }
1512
1513    #[test]
1514    fn a_str_value_still_encodes_as_plain_text_for_a_text_target() {
1515        let value = DecodedValue::Str("11111111-2222-3333-4444-555555555555".to_string());
1516        let mut out = bytes::BytesMut::new();
1517        encode_value(&value, &postgres_types::Type::TEXT, &mut out).unwrap();
1518        assert_eq!(out.as_ref(), "11111111-2222-3333-4444-555555555555".as_bytes());
1519    }
1520
1521    #[test]
1522    fn rejects_a_malformed_uuid_string_instead_of_sending_garbage_bytes() {
1523        let value = DecodedValue::Str("not-a-uuid".to_string());
1524        let mut out = bytes::BytesMut::new();
1525        assert!(encode_value(&value, &postgres_types::Type::UUID, &mut out).is_err());
1526    }
1527
1528    fn vector_type() -> Type {
1529        // `vector` is a pgvector extension type, not a `postgres_types`
1530        // builtin — construct it the way `Statement::params()` would
1531        // report it back (any OID works here; encoding only inspects the
1532        // name via `ty.name()`, matching `$n::vector`'s cast-reported type).
1533        Type::new(
1534            "vector".to_string(),
1535            50_000,
1536            postgres_types::Kind::Simple,
1537            "public".to_string(),
1538        )
1539    }
1540
1541    #[test]
1542    fn encodes_an_array_value_as_pgvector_binary_when_the_target_type_is_vector() {
1543        // Regression test: `$n::vector` reports the parameter's type as
1544        // the scalar `vector` type itself (unlike `$n::float8[]::vector`,
1545        // where the *inner* cast makes Postgres report `float8[]`) — an
1546        // `Array` value bound against it must produce pgvector's own
1547        // binary format (`u16 ndim`, `u16 reserved`, then big-endian
1548        // `f32`s), not the generic Postgres array wire format.
1549        let value = DecodedValue::Array(vec![
1550            DecodedValue::F64(1.5),
1551            DecodedValue::F64(-2.25),
1552            DecodedValue::F64(0.0),
1553        ]);
1554        let mut out = bytes::BytesMut::new();
1555        encode_value(&value, &vector_type(), &mut out).unwrap();
1556        let mut expected = vec![0u8, 3, 0, 0];
1557        expected.extend_from_slice(&1.5f32.to_be_bytes());
1558        expected.extend_from_slice(&(-2.25f32).to_be_bytes());
1559        expected.extend_from_slice(&0.0f32.to_be_bytes());
1560        assert_eq!(out.as_ref(), expected.as_slice());
1561    }
1562
1563    #[test]
1564    fn a_vector_encoded_value_round_trips_through_decode_vector() {
1565        let value = DecodedValue::Array(vec![
1566            DecodedValue::F64(1.0),
1567            DecodedValue::F64(2.0),
1568            DecodedValue::F64(3.0),
1569        ]);
1570        let mut out = bytes::BytesMut::new();
1571        encode_value(&value, &vector_type(), &mut out).unwrap();
1572        let decoded = decode_vector(out.as_ref()).unwrap();
1573        assert_eq!(
1574            decoded,
1575            vec![DecodedValue::F64(1.0), DecodedValue::F64(2.0), DecodedValue::F64(3.0)]
1576        );
1577    }
1578
1579    #[test]
1580    fn an_array_value_still_encodes_as_a_plain_postgres_array_for_a_non_vector_target() {
1581        let value = DecodedValue::Array(vec![DecodedValue::F64(1.0), DecodedValue::F64(2.0)]);
1582        let mut out = bytes::BytesMut::new();
1583        encode_value(&value, &postgres_types::Type::FLOAT8_ARRAY, &mut out).unwrap();
1584        // Generic array format starts with ndim=1 (i32), not pgvector's
1585        // ndim=2 (u16) — first four bytes distinguish the two encodings.
1586        assert_eq!(&out.as_ref()[0..4], &1i32.to_be_bytes());
1587    }
1588
1589    #[test]
1590    fn encodes_a_str_value_as_jsonb_binary_when_the_target_type_is_jsonb() {
1591        // Regression test for the same class of bug as the UUID case above:
1592        // a caller with already-serialized JSON text (e.g.
1593        // `schema_to_db_state_json`'s output) arrives as `DecodedValue::Str`,
1594        // not `::Object` — binding it against a `jsonb` parameter must add
1595        // the binary version-byte prefix, not send raw unframed text.
1596        let value = DecodedValue::Str(r#"{"a":1}"#.to_string());
1597        let mut out = bytes::BytesMut::new();
1598        encode_value(&value, &postgres_types::Type::JSONB, &mut out).unwrap();
1599        assert_eq!(out.as_ref(), [&[1u8][..], br#"{"a":1}"#].concat());
1600        // And decodes back correctly through the normal jsonb decode path.
1601        assert_eq!(
1602            decode_value(OID_JSONB, &out, &no_ext()).unwrap(),
1603            DecodedValue::Object(vec![("a".into(), DecodedValue::I64(1))])
1604        );
1605    }
1606}