use std::sync::Arc;
use pgwire::api::results::{DataRowEncoder, FieldInfo, QueryResponse, Response, Tag};
use pgwire::error::PgWireResult;
use serde_json::Value as JsonValue;
use crate::control::server::response_shape::types::{DdlColType, ShapedRows};
use crate::control::server::shared::ddl::result::{DdlError, DdlResult};
use super::types::{
bool_field, bytea_field, float4_array_field, float4_field, float8_array_field, float8_field,
int2_field, int4_field, int8_field, json_field, jsonb_field, sqlstate_error, text_field,
timestamp_field, timestamptz_field, varchar_field,
};
pub fn ddl_results_to_pgwire(
result: Result<Vec<DdlResult>, DdlError>,
) -> PgWireResult<Vec<Response>> {
let results = match result {
Ok(results) => results,
Err(DdlError { sqlstate, message }) => return Err(sqlstate_error(&sqlstate, &message)),
};
let mut responses = Vec::with_capacity(results.len());
for ddl in results {
responses.push(ddl_result_to_response(ddl)?);
}
Ok(responses)
}
fn ddl_result_to_response(ddl: DdlResult) -> PgWireResult<Response> {
match ddl {
DdlResult::Status {
command,
rows_affected,
} => {
let tag = match rows_affected {
Some(n) => Tag::new(&command).with_rows(n as usize),
None => Tag::new(&command),
};
Ok(Response::Execution(tag))
}
DdlResult::Empty => Ok(Response::EmptyQuery),
DdlResult::Rows(shaped) => rows_to_response(shaped),
}
}
fn rows_to_response(shaped: ShapedRows) -> PgWireResult<Response> {
let ShapedRows {
columns,
column_types,
rows,
..
} = shaped;
let fields: Vec<FieldInfo> = columns
.iter()
.enumerate()
.map(|(i, name)| {
let ct = column_types.get(i).copied().unwrap_or(DdlColType::Text);
col_type_to_field(name, ct)
})
.collect();
let schema = Arc::new(fields);
let mut encoded_rows: Vec<PgWireResult<pgwire::messages::data::DataRow>> =
Vec::with_capacity(rows.len());
for row in &rows {
let mut encoder = DataRowEncoder::new(schema.clone());
for (idx, name) in columns.iter().enumerate() {
let ct = column_types.get(idx).copied().unwrap_or(DdlColType::Text);
match row.get(name) {
Some(JsonValue::String(s)) => encoder.encode_field(&s)?,
Some(JsonValue::Null) | None => encoder.encode_field(&None::<&str>)?,
Some(JsonValue::Number(n)) => match ct {
DdlColType::Float8 => match n.as_f64() {
Some(f) => encoder.encode_field(&f)?,
None => encoder.encode_field(&None::<f64>)?,
},
DdlColType::Float4 => match n.as_f64() {
Some(f) => encoder.encode_field(&(f as f32))?,
None => encoder.encode_field(&None::<f32>)?,
},
_ => encoder.encode_field(&n.to_string())?,
},
Some(other) => encoder.encode_field(&other.to_string())?,
}
}
encoded_rows.push(Ok(encoder.take_row()));
}
Ok(Response::Query(QueryResponse::new(
schema,
futures::stream::iter(encoded_rows),
)))
}
pub(in crate::control::server::pgwire) fn col_type_to_field(
name: &str,
ct: DdlColType,
) -> FieldInfo {
match ct {
DdlColType::Text => text_field(name),
DdlColType::Int8 => int8_field(name),
DdlColType::Int4 => int4_field(name),
DdlColType::Int2 => int2_field(name),
DdlColType::Float8 => float8_field(name),
DdlColType::Float4 => float4_field(name),
DdlColType::Bool => bool_field(name),
DdlColType::Bytea => bytea_field(name),
DdlColType::Json => json_field(name),
DdlColType::Jsonb => jsonb_field(name),
DdlColType::Timestamp => timestamp_field(name),
DdlColType::Timestamptz => timestamptz_field(name),
DdlColType::Varchar => varchar_field(name),
DdlColType::Float4Array => float4_array_field(name),
DdlColType::Float8Array => float8_array_field(name),
}
}
pub(in crate::control::server::pgwire) fn col_type_to_field_with_format(
name: &str,
ct: DdlColType,
format: pgwire::api::results::FieldFormat,
) -> FieldInfo {
let base = col_type_to_field(name, ct);
FieldInfo::new(name.to_owned(), None, None, base.datatype().clone(), format)
}