Skip to main content

systemprompt_database/services/postgres/
conversion.rs

1//! Conversions between `SQLx` `PostgreSQL` rows, [`QueryResult`] /
2//! [`crate::models::JsonRow`] values, and parameter binders.
3//!
4//! Part of the documented sqlx allowlist: the binder operates on the dynamic
5//! `sqlx::query::Query` value passed in by the trait implementation.
6//!
7//! Copyright (c) systemprompt.io — Business Source License 1.1.
8//! See <https://systemprompt.io> for licensing details.
9
10use sqlx::{Column, Row};
11use std::collections::HashMap;
12
13use crate::models::{DbValue, QueryResult, ToDbValue};
14
15pub fn rows_to_result(rows: Vec<sqlx::postgres::PgRow>, start: std::time::Instant) -> QueryResult {
16    let mut columns = Vec::new();
17    let mut result_rows = Vec::new();
18
19    if let Some(first_row) = rows.first() {
20        columns = first_row
21            .columns()
22            .iter()
23            .map(|c| c.name().to_owned())
24            .collect();
25    }
26
27    for row in rows {
28        result_rows.push(row_to_json(&row));
29    }
30
31    let row_count = result_rows.len();
32    let execution_time_ms = u64::try_from(start.elapsed().as_millis()).unwrap_or(u64::MAX);
33
34    QueryResult {
35        columns,
36        rows: result_rows,
37        row_count,
38        execution_time_ms,
39    }
40}
41
42pub fn row_to_json(row: &sqlx::postgres::PgRow) -> HashMap<String, serde_json::Value> {
43    row.columns()
44        .iter()
45        .map(|col| (col.name().to_owned(), column_to_json(row, col.ordinal())))
46        .collect()
47}
48
49fn column_to_json(row: &sqlx::postgres::PgRow, ordinal: usize) -> serde_json::Value {
50    if let Ok(val) = row.try_get::<Option<chrono::DateTime<chrono::Utc>>, _>(ordinal) {
51        return val.map_or(serde_json::Value::Null, |v| {
52            serde_json::Value::String(v.to_rfc3339())
53        });
54    }
55    if let Ok(val) = row.try_get::<Option<uuid::Uuid>, _>(ordinal) {
56        return val.map_or(serde_json::Value::Null, |v| {
57            serde_json::Value::String(v.to_string())
58        });
59    }
60    if let Ok(val) = row.try_get::<Option<String>, _>(ordinal) {
61        return val.map_or(serde_json::Value::Null, serde_json::Value::String);
62    }
63    if let Ok(val) = row.try_get::<Option<i64>, _>(ordinal) {
64        return val.map_or(serde_json::Value::Null, |v| {
65            serde_json::Value::Number(v.into())
66        });
67    }
68    if let Ok(val) = row.try_get::<Option<i32>, _>(ordinal) {
69        return val.map_or(serde_json::Value::Null, |v| {
70            serde_json::Value::Number(i64::from(v).into())
71        });
72    }
73    if let Ok(val) = row.try_get::<Option<f64>, _>(ordinal) {
74        return val.map_or(serde_json::Value::Null, |v| serde_json::json!(v));
75    }
76    if let Ok(val) = row.try_get::<Option<rust_decimal::Decimal>, _>(ordinal) {
77        return val.map_or(serde_json::Value::Null, |v| {
78            v.to_string().parse::<f64>().map_or_else(
79                |_| serde_json::Value::String(v.to_string()),
80                |f| serde_json::json!(f),
81            )
82        });
83    }
84    if let Ok(val) = row.try_get::<Option<bool>, _>(ordinal) {
85        return val.map_or(serde_json::Value::Null, serde_json::Value::Bool);
86    }
87    if let Ok(val) = row.try_get::<Option<Vec<String>>, _>(ordinal) {
88        return val.map_or(serde_json::Value::Null, |v| {
89            serde_json::Value::Array(v.into_iter().map(serde_json::Value::String).collect())
90        });
91    }
92    if let Ok(val) = row.try_get::<Option<serde_json::Value>, _>(ordinal) {
93        return val.unwrap_or(serde_json::Value::Null);
94    }
95    if let Ok(val) = row.try_get::<Option<Vec<u8>>, _>(ordinal) {
96        return val.map_or(serde_json::Value::Null, |bytes| {
97            use base64::Engine;
98            use base64::engine::general_purpose::STANDARD;
99            serde_json::Value::String(STANDARD.encode(&bytes))
100        });
101    }
102    row.try_get_raw(ordinal)
103        .ok()
104        .map_or(serde_json::Value::Null, |value| raw_value_to_json(&value))
105}
106
107/// A column of a type none of the typed decoders accept — `regclass`,
108/// `"char"`, `oid`, `name`, `interval` and the other catalog types an
109/// operator meets in `pg_*` queries. Rendering nothing hid whole columns; the
110/// wire bytes are always representable as text or, for the object-id family,
111/// a number.
112// JSON: the raw text, the oid as a number, or base64 when neither applies.
113fn raw_value_to_json(value: &sqlx::postgres::PgValueRef<'_>) -> serde_json::Value {
114    use sqlx::{TypeInfo, ValueRef};
115    if value.is_null() {
116        return serde_json::Value::Null;
117    }
118    let type_name = value.type_info().name().to_ascii_uppercase();
119    match value.format() {
120        sqlx::postgres::PgValueFormat::Text => {
121            value.as_str().map_or(serde_json::Value::Null, |s| {
122                serde_json::Value::String(s.to_owned())
123            })
124        },
125        sqlx::postgres::PgValueFormat::Binary => {
126            let bytes = value.as_bytes().unwrap_or_default();
127            if bytes.len() == 4
128                && (type_name == "OID" || type_name.starts_with("REG"))
129                && let Ok(raw) = <[u8; 4]>::try_from(bytes)
130            {
131                return serde_json::Value::Number(u64::from(u32::from_be_bytes(raw)).into());
132            }
133            std::str::from_utf8(bytes).map_or_else(
134                |_| {
135                    use base64::Engine;
136                    use base64::engine::general_purpose::STANDARD;
137                    serde_json::Value::String(STANDARD.encode(bytes))
138                },
139                |s| serde_json::Value::String(s.to_owned()),
140            )
141        },
142    }
143}
144
145pub fn bind_params<'q>(
146    mut query: sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments>,
147    params: &[&dyn ToDbValue],
148) -> sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments> {
149    for param in params {
150        let value = param.to_db_value();
151        query = match value {
152            DbValue::String(s) => query.bind(s),
153            DbValue::Int(i) => query.bind(i),
154            DbValue::Float(f) => query.bind(f),
155            DbValue::Bool(b) => query.bind(b),
156            DbValue::Bytes(b) => query.bind(b),
157            DbValue::Timestamp(dt) => query.bind(dt),
158            DbValue::StringArray(arr) => query.bind(arr),
159            DbValue::NullString => query.bind(None::<String>),
160            DbValue::NullInt => query.bind(None::<i64>),
161            DbValue::NullFloat => query.bind(None::<f64>),
162            DbValue::NullBool => query.bind(None::<bool>),
163            DbValue::NullBytes => query.bind(None::<Vec<u8>>),
164            DbValue::NullTimestamp => query.bind(None::<chrono::DateTime<chrono::Utc>>),
165            DbValue::NullStringArray => query.bind(None::<Vec<String>>),
166        };
167    }
168    query
169}