use std::fmt::Debug;
use bytes::Bytes;
use futures::sink::Sink;
use pgwire::api::portal::Portal;
use pgwire::api::results::Response;
use pgwire::api::{ClientInfo, ClientPortalStore, Type};
use pgwire::error::{ErrorInfo, PgWireError, PgWireResult};
use pgwire::messages::PgWireBackendMessage;
use crate::control::server::response_shape::schema::{OutputColumn, OutputSchema};
use super::super::core::NodeDbPgHandler;
use super::super::routing::result_shaping::ResultShaping;
use super::result_format::{pg_type_to_ddl_col_type, resolve_result_formats};
use super::statement::ParsedStatement;
impl NodeDbPgHandler {
pub(crate) async fn execute_prepared<C>(
&self,
client: &mut C,
portal: &Portal<ParsedStatement>,
_max_rows: usize,
) -> PgWireResult<Response>
where
C: ClientInfo + ClientPortalStore + Sink<PgWireBackendMessage> + Unpin + Send + Sync,
C::Error: Debug,
PgWireError: From<<C as Sink<PgWireBackendMessage>>::Error>,
{
let addr = client.socket_addr();
let identity = self.resolve_identity(client, &addr)?;
self.authorize_session_database(&identity, &addr)?;
let stmt = &portal.statement.statement;
let tenant_id = identity.tenant_id;
let _audit_scope = crate::control::server::shared::session::audit_context::AuditScope::new(
crate::control::server::shared::session::audit_context::AuditCtx {
auth_user_id: identity.user_id.to_string(),
auth_user_name: identity.username.clone(),
sql_text: stmt.sql.clone(),
},
);
if let Some(intent) = crate::control::backup::detect(&stmt.sql) {
return self.intent_to_response(&identity, addr, intent).await;
}
let params = convert_portal_params(
&portal.parameters,
&stmt.param_types,
&portal.parameter_format,
)?;
if stmt.is_dsl {
let bound = nodedb_sql::dsl_bind::bind_dsl(&stmt.sql, ¶ms).map_err(|e| {
PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".into(),
"42601".into(),
format!("DSL parameter bind: {e}"),
)))
})?;
let mut results = self.execute_sql(&identity, &addr, bound.as_str()).await?;
return Ok(results.pop().unwrap_or(Response::EmptyQuery));
}
let result_formats =
resolve_result_formats(&stmt.result_fields, &portal.result_column_format);
let projection: Option<OutputSchema> = if stmt.result_fields.is_empty() {
None
} else {
Some(OutputSchema {
columns: stmt
.result_fields
.iter()
.map(|f| OutputColumn {
display_name: f.name().into(),
lookup_key: f.name().into(),
ty: pg_type_to_ddl_col_type(f.datatype()),
})
.collect(),
is_star: false,
})
};
let mut results = self
.execute_planned_sql_with_params(
&identity,
&stmt.sql,
tenant_id,
&addr,
¶ms,
ResultShaping {
projection: projection.as_ref(),
formats: &result_formats,
},
)
.await?;
Ok(results.pop().unwrap_or(Response::EmptyQuery))
}
}
fn convert_portal_params(
params: &[Option<Bytes>],
param_types: &[Option<Type>],
param_format: &pgwire::api::portal::Format,
) -> PgWireResult<Vec<nodedb_sql::ParamValue>> {
let mut result = Vec::with_capacity(params.len());
for (i, param) in params.iter().enumerate() {
let pg_type = param_types
.get(i)
.and_then(|t| t.as_ref())
.unwrap_or(&Type::UNKNOWN);
let pv = match param {
None => nodedb_sql::ParamValue::Null,
Some(bytes) => {
if param_format.is_binary(i) {
let type_name = if *pg_type == Type::NUMERIC {
Some("NUMERIC")
} else if *pg_type == Type::TIMESTAMP {
Some("TIMESTAMP")
} else if *pg_type == Type::TIMESTAMPTZ {
Some("TIMESTAMPTZ")
} else {
None
};
if let Some(name) = type_name {
return Err(PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"0A000".to_owned(),
format!(
"binary {name} parameter format is not supported for \
parameter ${n}; use text format",
n = i + 1
),
))));
}
}
let text = std::str::from_utf8(bytes).map_err(|_| {
PgWireError::UserError(Box::new(ErrorInfo::new(
"ERROR".to_owned(),
"22021".to_owned(),
format!("invalid UTF-8 in parameter ${}", i + 1),
)))
})?;
pgwire_text_to_param(text, pg_type)
}
};
result.push(pv);
}
Ok(result)
}
fn pgwire_text_to_param(text: &str, pg_type: &Type) -> nodedb_sql::ParamValue {
match *pg_type {
Type::BOOL => {
let lower = text.to_lowercase();
if lower == "t" || lower == "true" || lower == "1" {
return nodedb_sql::ParamValue::Bool(true);
}
if lower == "f" || lower == "false" || lower == "0" {
return nodedb_sql::ParamValue::Bool(false);
}
nodedb_sql::ParamValue::Text(text.to_string())
}
Type::INT2 | Type::INT4 | Type::INT8 => {
if let Ok(n) = text.parse::<i64>() {
return nodedb_sql::ParamValue::Int64(n);
}
nodedb_sql::ParamValue::Text(text.to_string())
}
Type::FLOAT4 | Type::FLOAT8 => {
if let Ok(f) = text.parse::<f64>() {
return nodedb_sql::ParamValue::Float64(f);
}
nodedb_sql::ParamValue::Text(text.to_string())
}
Type::NUMERIC => {
if let Ok(d) = rust_decimal::Decimal::from_str_exact(text) {
return nodedb_sql::ParamValue::Decimal(d);
}
nodedb_sql::ParamValue::Text(text.to_string())
}
Type::TIMESTAMP => {
if let Some(dt) = nodedb_types::datetime::NdbDateTime::parse(text) {
return nodedb_sql::ParamValue::Timestamp(dt);
}
nodedb_sql::ParamValue::Text(text.to_string())
}
Type::TIMESTAMPTZ => {
if let Some(dt) = nodedb_types::datetime::NdbDateTime::parse(text) {
return nodedb_sql::ParamValue::Timestamptz(dt);
}
nodedb_sql::ParamValue::Text(text.to_string())
}
_ => nodedb_sql::ParamValue::Text(text.to_string()),
}
}
#[cfg(test)]
mod tests {
use pgwire::api::portal::Format;
use super::*;
fn text_format() -> Format {
Format::UnifiedText
}
fn binary_format() -> Format {
Format::UnifiedBinary
}
#[test]
fn convert_null_param() {
let params = vec![None];
let types = vec![Some(Type::INT8)];
let result = convert_portal_params(¶ms, &types, &text_format()).unwrap();
assert_eq!(result.len(), 1);
assert!(matches!(result[0], nodedb_sql::ParamValue::Null));
}
#[test]
fn convert_typed_params() {
let params = vec![
Some(Bytes::from_static(b"42")),
Some(Bytes::from_static(b"hello")),
Some(Bytes::from_static(b"true")),
];
let types = vec![Some(Type::INT8), Some(Type::TEXT), Some(Type::BOOL)];
let result = convert_portal_params(¶ms, &types, &text_format()).unwrap();
assert!(matches!(result[0], nodedb_sql::ParamValue::Int64(42)));
assert!(matches!(&result[1], nodedb_sql::ParamValue::Text(s) if s == "hello"));
assert!(matches!(result[2], nodedb_sql::ParamValue::Bool(true)));
}
#[test]
fn convert_float_param() {
let params = vec![Some(Bytes::from_static(b"2.78"))];
let types = vec![Some(Type::FLOAT8)];
let result = convert_portal_params(¶ms, &types, &text_format()).unwrap();
assert!(
matches!(result[0], nodedb_sql::ParamValue::Float64(f) if (f - 2.78).abs() < f64::EPSILON)
);
}
#[test]
fn convert_numeric_text_to_decimal() {
let params = vec![Some(Bytes::from_static(b"123.45"))];
let types = vec![Some(Type::NUMERIC)];
let result = convert_portal_params(¶ms, &types, &text_format()).unwrap();
match &result[0] {
nodedb_sql::ParamValue::Decimal(decimal) => assert_eq!(decimal.to_string(), "123.45"),
other => panic!("expected Decimal, got {other:?}"),
}
}
fn assert_binary_type_rejected(ty: Type, bytes: &'static [u8], name: &str) {
let params = vec![Some(Bytes::from_static(bytes))];
let types = vec![Some(ty)];
let error = convert_portal_params(¶ms, &types, &binary_format()).unwrap_err();
let message = error.to_string();
assert!(message.contains(name) || message.contains("0A000"));
}
#[test]
fn convert_numeric_binary_returns_error() {
assert_binary_type_rejected(Type::NUMERIC, &[0x00, 0x03, 0x00, 0x02], "NUMERIC");
}
#[test]
fn convert_timestamp_binary_returns_error() {
assert_binary_type_rejected(Type::TIMESTAMP, &[0; 8], "TIMESTAMP");
}
#[test]
fn convert_timestamptz_binary_returns_error() {
assert_binary_type_rejected(Type::TIMESTAMPTZ, &[0; 8], "TIMESTAMPTZ");
}
fn assert_text_param(
input: &'static [u8],
ty: Type,
expected: fn(&nodedb_sql::ParamValue) -> bool,
) {
let params = vec![Some(Bytes::from_static(input))];
let types = vec![Some(ty)];
let result = convert_portal_params(¶ms, &types, &text_format()).unwrap();
assert!(expected(&result[0]));
}
#[test]
fn convert_timestamp_text_to_typed() {
assert_text_param(b"2024-01-01 00:00:00", Type::TIMESTAMP, |value| {
matches!(value, nodedb_sql::ParamValue::Timestamp(_))
});
}
#[test]
fn convert_timestamptz_text_to_typed() {
assert_text_param(b"2024-01-01 00:00:00+00", Type::TIMESTAMPTZ, |value| {
matches!(value, nodedb_sql::ParamValue::Timestamptz(_))
});
}
#[test]
fn convert_bool_variants() {
for (input, expected) in [("t", true), ("f", false), ("1", true), ("0", false)] {
let params = vec![Some(Bytes::from(input))];
let types = vec![Some(Type::BOOL)];
let result = convert_portal_params(¶ms, &types, &text_format()).unwrap();
assert!(matches!(result[0], nodedb_sql::ParamValue::Bool(value) if value == expected));
}
}
#[test]
fn passthrough_date_text() {
let value = pgwire_text_to_param("2026-04-19", &Type::DATE);
assert!(matches!(&value, nodedb_sql::ParamValue::Text(text) if text == "2026-04-19"));
}
#[test]
fn timestamp_text_parses_to_typed() {
let value = pgwire_text_to_param("2026-04-19 12:00:00", &Type::TIMESTAMP);
assert!(matches!(value, nodedb_sql::ParamValue::Timestamp(_)));
}
#[test]
fn timestamptz_text_parses_to_typed() {
let value = pgwire_text_to_param("2026-04-19 12:00:00+00", &Type::TIMESTAMPTZ);
assert!(matches!(value, nodedb_sql::ParamValue::Timestamptz(_)));
}
#[test]
fn passthrough_uuid_text() {
let uuid = "550e8400-e29b-41d4-a716-446655440000";
let value = pgwire_text_to_param(uuid, &Type::UUID);
assert!(matches!(&value, nodedb_sql::ParamValue::Text(text) if text == uuid));
}
#[test]
fn passthrough_jsonb_text() {
let json = r#"{"a":1}"#;
let value = pgwire_text_to_param(json, &Type::JSONB);
assert!(matches!(&value, nodedb_sql::ParamValue::Text(text) if text == json));
}
#[test]
fn passthrough_bytea_hex_text() {
let value = pgwire_text_to_param("\\xDEADBEEF", &Type::BYTEA);
assert!(matches!(&value, nodedb_sql::ParamValue::Text(text) if text == "\\xDEADBEEF"));
}
#[test]
fn int_parse_failure_falls_back_to_text() {
let value = pgwire_text_to_param("abc", &Type::INT8);
assert!(matches!(&value, nodedb_sql::ParamValue::Text(text) if text == "abc"));
}
#[test]
fn unknown_type_routes_to_text() {
let value = pgwire_text_to_param("42", &Type::UNKNOWN);
assert!(matches!(&value, nodedb_sql::ParamValue::Text(text) if text == "42"));
}
}