use super::expr::parse_scalar_expr;
use crate::catalog::SchemaName;
use crate::engine::{DataType, IdentitySpec, ObjName, ScalarExpr, Value, fe};
use pg_query::protobuf::a_const::Val;
use pg_query::protobuf::{AConst, ColumnDef, TypeName};
use pg_query::{Node, NodeEnum};
use pgwire::error::PgWireResult;
type ColumnDefSpec = (
String,
DataType,
bool,
Option<ScalarExpr>,
Option<IdentitySpec>,
);
pub(super) fn const_to_value(c: &AConst) -> PgWireResult<Value> {
if c.val.is_none() {
return Ok(Value::Null);
}
let v = c.val.as_ref().unwrap();
match v {
Val::Ival(i) => Ok(Value::Int64(i.ival as i64)),
Val::Fval(f) => {
Ok(Value::from_f64(f.fval.parse::<f64>().map_err(|e| {
pgwire::error::PgWireError::ApiError(Box::new(e))
})?))
}
Val::Boolval(b) => Ok(Value::Bool(b.boolval)),
Val::Sval(s) => Ok(Value::Text(s.sval.clone())),
Val::Bsval(_) => Err(fe("bitstring const not yet supported")),
}
}
pub(super) fn map_type(cd: &ColumnDef) -> PgWireResult<DataType> {
let typ = cd.type_name.as_ref().ok_or_else(|| fe("missing type"))?;
parse_type_name(typ)
}
pub(super) fn parse_type_name(typ: &TypeName) -> PgWireResult<DataType> {
let mut tokens: Vec<String> = typ
.names
.iter()
.filter_map(|n| {
n.node.as_ref().and_then(|nn| {
if let NodeEnum::String(s) = nn {
Some(s.sval.to_ascii_lowercase())
} else {
None
}
})
})
.collect();
tokens.retain(|t| t != "pg_catalog" && t != "public");
if tokens.is_empty() {
return Err(fe("bad type name"));
}
let last = tokens.last().unwrap().as_str();
let dt = if tokens.len() >= 2
&& tokens[tokens.len() - 2] == "double"
&& tokens[tokens.len() - 1] == "precision"
{
DataType::Float8
} else if tokens.len() >= 4
&& tokens[tokens.len() - 4] == "timestamp"
&& tokens[tokens.len() - 3] == "without"
&& tokens[tokens.len() - 2] == "time"
&& tokens[tokens.len() - 1] == "zone"
{
DataType::Timestamp
} else if tokens.len() >= 4
&& tokens[tokens.len() - 4] == "timestamp"
&& tokens[tokens.len() - 3] == "with"
&& tokens[tokens.len() - 2] == "time"
&& tokens[tokens.len() - 1] == "zone"
{
DataType::Timestamptz
} else {
match last {
"int" | "int4" | "integer" => DataType::Int4,
"bigint" | "int8" => DataType::Int8,
"float8" | "double" => DataType::Float8,
"text" | "varchar" => DataType::Text,
"json" => DataType::Json,
"jsonb" => DataType::Jsonb,
"bool" | "boolean" => DataType::Bool,
"oid" => DataType::Int8,
"date" => DataType::Date,
"timestamp" => DataType::Timestamp,
"timestamptz" => DataType::Timestamptz,
"bytea" => DataType::Bytea,
"interval" => DataType::Interval,
"regtype" => DataType::Text,
"void" => DataType::Void,
other => return Err(fe(format!("unsupported type: {other}"))),
}
};
Ok(dt)
}
pub(super) fn parse_column_def(cd: &ColumnDef) -> PgWireResult<ColumnDefSpec> {
let dt = map_type(cd)?;
let default_node = cd
.raw_default
.as_ref()
.and_then(|n| n.node.as_ref())
.or_else(|| cd.cooked_default.as_ref().and_then(|n| n.node.as_ref()))
.or_else(|| {
cd.constraints.iter().find_map(|c| {
let Some(NodeEnum::Constraint(cons)) = c.node.as_ref() else {
return None;
};
if cons.contype == pg_query::protobuf::ConstrType::ConstrDefault as i32 {
cons.raw_expr.as_ref().and_then(|n| n.node.as_ref())
} else {
None
}
})
});
let mut nullable = !cd
.constraints
.iter()
.any(|c| matches!(c.node.as_ref(), Some(NodeEnum::Constraint(cons)) if cons.contype == pg_query::protobuf::ConstrType::ConstrNotnull as i32));
let default = match default_node {
Some(node) => {
let expr = parse_scalar_expr(node)?;
ensure_default_expr_is_const(&expr)?;
Some(expr)
}
None => None,
};
let identity = parse_identity_spec(cd)?;
if let Some(spec) = &identity {
if !matches!(dt, DataType::Int4 | DataType::Int8) {
return Err(fe("IDENTITY columns must be INT or BIGINT"));
}
if spec.increment_by == 0 {
return Err(fe("IDENTITY INCREMENT BY cannot be zero"));
}
}
let name = cd.colname.clone();
if name.is_empty() {
return Err(fe("column must have a name"));
}
if identity.is_some() && default.is_some() {
return Err(fe(format!(
"identity column {name} cannot have an explicit DEFAULT"
)));
}
if identity.is_some() {
nullable = false;
}
Ok((name, dt, nullable, default, identity))
}
fn ensure_default_expr_is_const(expr: &ScalarExpr) -> PgWireResult<()> {
match expr {
ScalarExpr::Literal(_) => Ok(()),
ScalarExpr::Column(..) | ScalarExpr::ColumnIdx(_) | ScalarExpr::ExcludedIdx(_) => {
Err(fe("DEFAULT expressions cannot reference columns"))
}
ScalarExpr::Param { .. } => Err(fe("DEFAULT expressions cannot reference parameters")),
ScalarExpr::BinaryOp { left, right, .. } => {
ensure_default_expr_is_const(left)?;
ensure_default_expr_is_const(right)
}
ScalarExpr::UnaryOp { expr, .. } | ScalarExpr::Cast { expr, .. } => {
ensure_default_expr_is_const(expr)
}
ScalarExpr::Func { args, .. } => {
for arg in args {
ensure_default_expr_is_const(arg)?;
}
Ok(())
}
ScalarExpr::Predicate(expr) => ensure_default_bool_expr_is_const(expr),
ScalarExpr::Subquery(_) => Err(fe("DEFAULT expressions cannot contain subqueries")),
ScalarExpr::Case {
when_then,
else_expr,
} => {
for (cond, result) in when_then {
ensure_default_bool_expr_is_const(cond)?;
ensure_default_expr_is_const(result)?;
}
if let Some(expr) = else_expr {
ensure_default_expr_is_const(expr)?;
}
Ok(())
}
}
}
fn ensure_default_bool_expr_is_const(expr: &crate::engine::BoolExpr) -> PgWireResult<()> {
use crate::engine::BoolExpr;
match expr {
BoolExpr::Literal(_) => Ok(()),
BoolExpr::Comparison { lhs, rhs, .. } => {
ensure_default_expr_is_const(lhs)?;
ensure_default_expr_is_const(rhs)
}
BoolExpr::And(parts) | BoolExpr::Or(parts) => {
for part in parts {
ensure_default_bool_expr_is_const(part)?;
}
Ok(())
}
BoolExpr::Not(inner) => ensure_default_bool_expr_is_const(inner),
BoolExpr::IsNull { expr, .. } => ensure_default_expr_is_const(expr),
BoolExpr::InSubquery { .. } | BoolExpr::InListValues { .. } => {
Err(fe("DEFAULT expressions cannot contain subqueries"))
}
}
}
fn parse_identity_spec(cd: &ColumnDef) -> PgWireResult<Option<IdentitySpec>> {
let mut spec: Option<IdentitySpec> = None;
for constraint in &cd.constraints {
let Some(NodeEnum::Constraint(cons)) = constraint.node.as_ref() else {
continue;
};
if cons.contype != pg_query::protobuf::ConstrType::ConstrIdentity as i32 {
continue;
}
if spec.is_some() {
return Err(fe(format!(
"column {} specifies IDENTITY more than once",
cd.colname
)));
}
let always = match cons.generated_when.as_str() {
"a" | "A" => true,
"" | "d" | "D" => false,
other => {
return Err(fe(format!(
"unsupported IDENTITY generation mode {other:?}"
)));
}
};
let mut start_with = None;
let mut increment_by = None;
for opt in &cons.options {
let Some(NodeEnum::DefElem(def)) = opt.node.as_ref() else {
continue;
};
match def.defname.as_str() {
"start" => start_with = Some(parse_identity_option_value(&def.arg)?),
"increment" => increment_by = Some(parse_identity_option_value(&def.arg)?),
_ => {}
}
}
spec = Some(IdentitySpec {
always,
start_with: start_with.unwrap_or(1),
increment_by: increment_by.unwrap_or(1),
});
}
Ok(spec)
}
fn parse_identity_option_value(arg: &Option<Box<Node>>) -> PgWireResult<i128> {
let node = arg
.as_ref()
.and_then(|n| n.node.as_ref())
.ok_or_else(|| fe("IDENTITY option requires a value"))?;
match node {
NodeEnum::Integer(i) => Ok(i.ival as i128),
NodeEnum::AConst(c) => match const_to_value(c)? {
Value::Int64(v) => Ok(v as i128),
Value::Text(s) => s
.parse::<i128>()
.map_err(|_| fe("IDENTITY option requires integer")),
_ => Err(fe("IDENTITY option requires integer literal")),
},
_ => Err(fe("IDENTITY option requires integer literal")),
}
}
pub(super) fn parse_index_columns(params: &[pg_query::Node]) -> PgWireResult<Vec<String>> {
if params.is_empty() {
return Err(fe("index requires at least one column"));
}
let mut cols = Vec::with_capacity(params.len());
for p in params {
let node = p.node.as_ref().ok_or_else(|| fe("bad index column"))?;
let NodeEnum::IndexElem(elem) = node else {
return Err(fe("index expressions not supported"));
};
if elem.expr.is_some() {
return Err(fe("expression indexes not supported"));
}
if elem.name.is_empty() {
return Err(fe("index column name required"));
}
cols.push(elem.name.clone());
}
Ok(cols)
}
pub(super) fn parse_obj_name_from_list(node: &NodeEnum) -> PgWireResult<ObjName> {
let mut parts = Vec::new();
match node {
NodeEnum::List(list) => {
for item in &list.items {
let Some(NodeEnum::String(s)) = item.node.as_ref() else {
return Err(fe("bad qualified name component"));
};
parts.push(s.sval.clone());
}
}
NodeEnum::String(s) => parts.push(s.sval.clone()),
_ => return Err(fe("bad qualified name")),
}
if parts.is_empty() {
return Err(fe("empty name"));
}
let name = parts.pop().unwrap();
let schema = if parts.is_empty() {
None
} else {
Some(SchemaName::new(parts.join(".")))
};
Ok(ObjName { schema, name })
}
pub(super) fn parse_set_value(args: &[pg_query::Node]) -> PgWireResult<Vec<String>> {
if args.is_empty() {
return Err(fe("SET requires value"));
}
let mut values = Vec::with_capacity(args.len());
for arg in args {
let node = arg.node.as_ref().ok_or_else(|| fe("bad SET value"))?;
let Some(v) = try_parse_literal(node)? else {
return Err(fe("unsupported SET value"));
};
values.push(literal_value_to_string(v)?);
}
Ok(values)
}
pub(super) fn literal_value_to_string(value: Value) -> PgWireResult<String> {
Ok(match value {
Value::Text(s) => s,
Value::Int64(i) => i.to_string(),
Value::Bool(b) => {
if b {
"true".into()
} else {
"false".into()
}
}
_ => return Err(fe("SET literal type not supported")),
})
}
pub(super) fn try_parse_literal(node: &NodeEnum) -> PgWireResult<Option<Value>> {
match node {
NodeEnum::AConst(c) => Ok(Some(const_to_value(c)?)),
NodeEnum::AExpr(ax) => {
let is_minus = ax.name.iter().any(|nn| {
matches!(
nn.node.as_ref(),
Some(NodeEnum::String(s)) if s.sval == "-"
)
});
if is_minus {
let rhs = ax
.rexpr
.as_ref()
.and_then(|n| n.node.as_ref())
.ok_or_else(|| fe("bad unary minus"))?;
match rhs {
NodeEnum::AConst(c) => match const_to_value(c)? {
Value::Int64(i) => Ok(Some(Value::Int64(-i))),
Value::Float64Bits(b) => Ok(Some(Value::from_f64(-f64::from_bits(b)))),
Value::Null => Err(fe("minus over null")),
Value::Text(_)
| Value::Bool(_)
| Value::Date(_)
| Value::TimestampMicros(_)
| Value::TimestamptzMicros(_)
| Value::Bytes(_)
| Value::IntervalMicros(_) => Err(fe("minus over non-numeric literal")),
},
_ => Err(fe("minus over non-const")),
}
} else {
Ok(None)
}
}
_ => Ok(None),
}
}