#[must_use]
pub fn default_name_for_pk(table: &str) -> String {
format!("{table}_pkey")
}
#[must_use]
pub fn default_name_for_fk(
table: &str,
columns: &[String],
_table_to: &str,
_columns_to: &[String],
) -> String {
let first_column = columns.first().map_or("", String::as_str);
let desired = format!("{table}_{first_column}_fkey");
if desired.len() > 63 {
let hash = hash_string(&desired);
if table.len() < 63 - 18 {
format!("{table}_{hash}_fkey")
} else {
format!("{hash}_fkey")
}
} else {
desired
}
}
#[must_use]
pub fn default_name_for_unique(table: &str, columns: &[String]) -> String {
truncate_identifier(&format!("{}_{}_key", table, columns.join("_")), "_key")
}
#[must_use]
pub fn default_name_for_index(table: &str, columns: &[String]) -> String {
truncate_identifier(&format!("{}_{}_idx", table, columns.join("_")), "_idx")
}
#[must_use]
pub fn default_name_for_identity_sequence(table: &str, column: &str) -> String {
format!("{table}_{column}_seq")
}
#[must_use]
pub fn default_name_for_check(table: &str, index: usize) -> String {
format!("{table}_check{}", index + 1)
}
fn hash_string(s: &str) -> String {
use sha2::{Digest, Sha256};
let digest = Sha256::digest(s.as_bytes());
let mut out = String::with_capacity(12);
for byte in digest.iter().take(6) {
use std::fmt::Write;
let _ = write!(out, "{byte:02x}");
}
out
}
fn truncate_identifier(name: &str, suffix: &str) -> String {
const MAX_IDENTIFIER_LEN: usize = 63;
if name.len() <= MAX_IDENTIFIER_LEN {
return name.to_string();
}
let hash = hash_string(name);
let budget = MAX_IDENTIFIER_LEN - hash.len() - 1 - suffix.len();
let mut cutoff = budget.min(name.len());
while !name.is_char_boundary(cutoff) {
cutoff -= 1;
}
format!("{}_{hash}{suffix}", &name[..cutoff])
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PgTypeCategory {
SmallInt,
Integer,
BigInt,
Numeric,
Real,
DoublePrecision,
Boolean,
Char,
Varchar,
Text,
Json,
Jsonb,
Time,
TimeTz,
Timestamp,
TimestampTz,
Date,
Uuid,
Interval,
Inet,
Cidr,
MacAddr,
MacAddr8,
Vector,
HalfVec,
SparseVec,
Bit,
Point,
Line,
Geometry,
Serial,
SmallSerial,
BigSerial,
Enum,
Custom,
}
impl PgTypeCategory {
fn type_name_rest<'a>(s: &'a str, type_name: &str) -> Option<&'a str> {
if !s.starts_with(type_name) {
return None;
}
let rest = &s[type_name.len()..];
if rest
.chars()
.next()
.is_some_and(|ch| ch.is_ascii_alphanumeric() || ch == '_')
{
return None;
}
Some(rest)
}
fn first_type_argument<'a>(s: &'a str, type_name: &str) -> Option<&'a str> {
let rest = Self::type_name_rest(s, type_name)?.trim_start();
let body = rest.strip_prefix('(')?;
let end = body.find([',', ')'])?;
Some(body[..end].trim())
}
fn match_numeric(s: &str) -> Option<Self> {
if s.starts_with("smallserial") {
return Some(Self::SmallSerial);
}
if s.starts_with("bigserial") {
return Some(Self::BigSerial);
}
if s.starts_with("serial") {
return Some(Self::Serial);
}
if s.starts_with("smallint") || s == "int2" {
return Some(Self::SmallInt);
}
if s.starts_with("integer") || s == "int" || s == "int4" {
return Some(Self::Integer);
}
if s.starts_with("bigint") || s == "int8" {
return Some(Self::BigInt);
}
if s.starts_with("numeric") || s.starts_with("decimal") {
return Some(Self::Numeric);
}
if s.starts_with("real") || s == "float4" {
return Some(Self::Real);
}
if s.starts_with("double") {
return Some(Self::DoublePrecision);
}
if s.starts_with("boolean") || s == "bool" {
return Some(Self::Boolean);
}
None
}
fn match_string_or_json(s: &str) -> Option<Self> {
if s.starts_with("varchar") || s.starts_with("character varying") {
return Some(Self::Varchar);
}
if s.starts_with("char") || s.starts_with("character") {
return Some(Self::Char);
}
if s.starts_with("text") {
return Some(Self::Text);
}
if s.starts_with("jsonb") {
return Some(Self::Jsonb);
}
if s.starts_with("json") {
return Some(Self::Json);
}
None
}
fn match_temporal(s: &str) -> Option<Self> {
if s.starts_with("timestamp") && s.contains("with time zone") {
return Some(Self::TimestampTz);
}
if s.starts_with("timestamp") {
return Some(Self::Timestamp);
}
if s.starts_with("time") && s.contains("with time zone") {
return Some(Self::TimeTz);
}
if s.starts_with("time") {
return Some(Self::Time);
}
if s.starts_with("date") {
return Some(Self::Date);
}
if s.starts_with("interval") {
return Some(Self::Interval);
}
None
}
fn match_specialized(s: &str) -> Option<Self> {
if s.starts_with("uuid") {
return Some(Self::Uuid);
}
if s.starts_with("inet") {
return Some(Self::Inet);
}
if s.starts_with("cidr") {
return Some(Self::Cidr);
}
if s.starts_with("macaddr8") {
return Some(Self::MacAddr8);
}
if s.starts_with("macaddr") {
return Some(Self::MacAddr);
}
if s.starts_with("vector") {
return Some(Self::Vector);
}
if s.starts_with("halfvec") {
return Some(Self::HalfVec);
}
if s.starts_with("sparsevec") {
return Some(Self::SparseVec);
}
if s.starts_with("bit") {
return Some(Self::Bit);
}
if Self::type_name_rest(s, "geometry").is_some() {
return Some(match Self::first_type_argument(s, "geometry") {
Some("point") => Self::Geometry,
_ => Self::Custom,
});
}
if Self::type_name_rest(s, "geography").is_some()
|| Self::type_name_rest(s, "box2d").is_some()
|| Self::type_name_rest(s, "box3d").is_some()
|| Self::type_name_rest(s, "raster").is_some()
{
return Some(Self::Custom);
}
if s.starts_with("point") {
return Some(Self::Point);
}
if s.starts_with("line") {
return Some(Self::Line);
}
None
}
#[must_use]
pub fn from_sql_type(sql_type: &str) -> Self {
let s = sql_type.trim().to_lowercase();
Self::match_numeric(&s)
.or_else(|| Self::match_string_or_json(&s))
.or_else(|| Self::match_temporal(&s))
.or_else(|| Self::match_specialized(&s))
.unwrap_or(Self::Custom)
}
#[must_use]
pub const fn drizzle_import(&self) -> &'static str {
match self {
Self::SmallInt => "smallint",
Self::Integer => "integer",
Self::BigInt => "bigint",
Self::Numeric => "numeric",
Self::Real => "real",
Self::DoublePrecision => "doublePrecision",
Self::Boolean => "boolean",
Self::Char => "char",
Self::Varchar => "varchar",
Self::Text => "text",
Self::Json => "json",
Self::Jsonb => "jsonb",
Self::Time | Self::TimeTz => "time",
Self::Timestamp | Self::TimestampTz => "timestamp",
Self::Date => "date",
Self::Uuid => "uuid",
Self::Interval => "interval",
Self::Inet => "inet",
Self::Cidr => "cidr",
Self::MacAddr => "macaddr",
Self::MacAddr8 => "macaddr8",
Self::Vector => "vector",
Self::HalfVec => "halfvec",
Self::SparseVec => "sparsevec",
Self::Bit => "bit",
Self::Point => "point",
Self::Line => "line",
Self::Geometry => "geometry",
Self::Serial => "serial",
Self::SmallSerial => "smallserial",
Self::BigSerial => "bigserial",
Self::Enum => "pgEnum",
Self::Custom => "customType",
}
}
#[must_use]
pub const fn is_serial(&self) -> bool {
matches!(self, Self::Serial | Self::SmallSerial | Self::BigSerial)
}
}
#[must_use]
pub fn parse_type_params(sql_type: &str) -> Option<(String, Option<String>)> {
let start = sql_type.find('(')?;
let end = sql_type.find(')')?;
let params = &sql_type[start + 1..end];
let parts: Vec<&str> = params.split(',').map(str::trim).collect();
match parts.len() {
1 => Some((parts[0].to_string(), None)),
2 => Some((parts[0].to_string(), Some(parts[1].to_string()))),
_ => None,
}
}
#[must_use]
pub fn is_serial_expression(expr: &str, schema: &str) -> bool {
let schema_prefix = if schema == "public" {
String::new()
} else {
format!("{schema}.")
};
(expr.starts_with(&format!("nextval('{schema_prefix}"))
|| expr.starts_with(&format!("nextval('\"{schema_prefix}")))
&& (expr.ends_with("_seq'::regclass)") || expr.ends_with("_seq\"'::regclass)"))
}
#[must_use]
pub fn extract_nextval_sequence(expr: &str) -> Option<String> {
let inner = expr
.strip_prefix("nextval('")?
.strip_suffix("'::regclass)")?;
let name_part = inner.rfind('.').map_or(inner, |pos| &inner[pos + 1..]);
let name = name_part.trim_matches('"');
if name.is_empty() {
return None;
}
Some(name.to_string())
}
pub struct IdentityDefaults;
impl IdentityDefaults {
pub const START_WITH: &'static str = "1";
pub const INCREMENT: &'static str = "1";
pub const MIN: &'static str = "1";
pub const CACHE: i32 = 1;
pub const CYCLE: bool = false;
#[must_use]
pub fn max_for(column_type: &str) -> &'static str {
match column_type {
"smallint" => "32767",
"bigint" => "9223372036854775807",
_ => "2147483647",
}
}
#[must_use]
pub fn min_for(column_type: &str) -> &'static str {
match column_type {
"smallint" => "-32768",
"bigint" => "-9223372036854775808",
_ => "-2147483648",
}
}
}
pub const SYSTEM_NAMESPACE_NAMES: &[&str] = &["pg_toast", "pg_catalog", "information_schema"];
#[must_use]
pub fn is_system_namespace(name: &str) -> bool {
name.starts_with("pg_toast")
|| name == "pg_default"
|| name == "pg_global"
|| name.starts_with("pg_temp_")
|| SYSTEM_NAMESPACE_NAMES.contains(&name)
}
#[must_use]
pub fn is_system_role(name: &str) -> bool {
name == "postgres" || name.starts_with("pg_")
}
pub struct PgDefaults;
impl PgDefaults {
pub const TABLESPACE: &'static str = "pg_default";
pub const ACCESS_METHOD: &'static str = "heap";
pub const NULLS_NOT_DISTINCT: bool = false;
pub const INDEX_METHOD: &'static str = "btree";
pub const GEOMETRY_SRID: i32 = 0;
}
pub const VECTOR_OPS: &[&str] = &[
"vector_l2_ops",
"vector_ip_ops",
"vector_cosine_ops",
"vector_l1_ops",
"bit_hamming_ops",
"bit_jaccard_ops",
"halfvec_l2_ops",
"sparsevec_l2_ops",
];
#[must_use]
pub fn parse_check_definition(value: &str) -> String {
let trimmed = value.trim();
let rest = trimmed
.strip_prefix("CHECK")
.or_else(|| trimmed.strip_prefix("check"))
.map_or(trimmed, str::trim_start);
rest.to_string()
}
#[must_use]
pub fn parse_view_definition(value: &str) -> String {
value
.split_whitespace()
.collect::<Vec<_>>()
.join(" ")
.trim_end_matches(';')
.to_string()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_name_for_pk() {
assert_eq!(default_name_for_pk("users"), "users_pkey");
}
#[test]
fn test_default_name_for_fk() {
let name = default_name_for_fk(
"posts",
&["author_id".to_string()],
"users",
&["id".to_string()],
);
assert_eq!(name, "posts_author_id_fkey");
}
#[test]
fn test_default_name_for_composite_fk_uses_first_column() {
let name = default_name_for_fk(
"order_lines",
&["order_id".to_string(), "tenant_id".to_string()],
"orders",
&["id".to_string(), "tenant_id".to_string()],
);
assert_eq!(name, "order_lines_order_id_fkey");
}
#[test]
fn test_default_name_for_unique() {
let name = default_name_for_unique("users", &["email".to_string()]);
assert_eq!(name, "users_email_key");
}
#[test]
fn test_default_name_for_index() {
let name = default_name_for_index("users", &["email".to_string(), "name".to_string()]);
assert_eq!(name, "users_email_name_idx");
}
#[test]
fn test_parse_type_params() {
assert_eq!(
parse_type_params("varchar(255)"),
Some(("255".to_string(), None))
);
assert_eq!(
parse_type_params("numeric(10,2)"),
Some(("10".to_string(), Some("2".to_string())))
);
assert_eq!(parse_type_params("text"), None);
}
#[test]
fn test_is_system_namespace() {
assert!(is_system_namespace("pg_catalog"));
assert!(is_system_namespace("pg_toast_12345"));
assert!(!is_system_namespace("public"));
assert!(!is_system_namespace("myschema"));
}
#[test]
fn test_identity_defaults() {
assert_eq!(IdentityDefaults::max_for("smallint"), "32767");
assert_eq!(IdentityDefaults::max_for("integer"), "2147483647");
assert_eq!(IdentityDefaults::max_for("bigint"), "9223372036854775807");
}
#[test]
fn test_from_sql_type_serial_vs_integer() {
assert_eq!(
PgTypeCategory::from_sql_type("integer"),
PgTypeCategory::Integer
);
assert_eq!(
PgTypeCategory::from_sql_type("int"),
PgTypeCategory::Integer
);
assert_eq!(
PgTypeCategory::from_sql_type("int4"),
PgTypeCategory::Integer
);
assert_eq!(
PgTypeCategory::from_sql_type("bigint"),
PgTypeCategory::BigInt
);
assert_eq!(
PgTypeCategory::from_sql_type("int8"),
PgTypeCategory::BigInt
);
assert_eq!(
PgTypeCategory::from_sql_type("smallint"),
PgTypeCategory::SmallInt
);
assert_eq!(
PgTypeCategory::from_sql_type("int2"),
PgTypeCategory::SmallInt
);
assert_eq!(
PgTypeCategory::from_sql_type("serial"),
PgTypeCategory::Serial
);
assert_eq!(
PgTypeCategory::from_sql_type("SERIAL"),
PgTypeCategory::Serial
);
assert_eq!(
PgTypeCategory::from_sql_type("bigserial"),
PgTypeCategory::BigSerial
);
assert_eq!(
PgTypeCategory::from_sql_type("smallserial"),
PgTypeCategory::SmallSerial
);
assert!(PgTypeCategory::Serial.is_serial());
assert!(PgTypeCategory::BigSerial.is_serial());
assert!(PgTypeCategory::SmallSerial.is_serial());
assert!(!PgTypeCategory::Integer.is_serial());
assert!(!PgTypeCategory::BigInt.is_serial());
}
#[test]
fn test_from_sql_type_common() {
assert_eq!(PgTypeCategory::from_sql_type("text"), PgTypeCategory::Text);
assert_eq!(
PgTypeCategory::from_sql_type("varchar(255)"),
PgTypeCategory::Varchar
);
assert_eq!(
PgTypeCategory::from_sql_type("boolean"),
PgTypeCategory::Boolean
);
assert_eq!(
PgTypeCategory::from_sql_type("bool"),
PgTypeCategory::Boolean
);
assert_eq!(PgTypeCategory::from_sql_type("uuid"), PgTypeCategory::Uuid);
assert_eq!(
PgTypeCategory::from_sql_type("jsonb"),
PgTypeCategory::Jsonb
);
assert_eq!(PgTypeCategory::from_sql_type("json"), PgTypeCategory::Json);
assert_eq!(
PgTypeCategory::from_sql_type("timestamp with time zone"),
PgTypeCategory::TimestampTz
);
assert_eq!(
PgTypeCategory::from_sql_type("timestamp without time zone"),
PgTypeCategory::Timestamp
);
assert_eq!(
PgTypeCategory::from_sql_type("timestamp"),
PgTypeCategory::Timestamp
);
assert_eq!(
PgTypeCategory::from_sql_type("time without time zone"),
PgTypeCategory::Time
);
assert_eq!(PgTypeCategory::from_sql_type("date"), PgTypeCategory::Date);
assert_eq!(
PgTypeCategory::from_sql_type("numeric(10,2)"),
PgTypeCategory::Numeric
);
assert_eq!(PgTypeCategory::from_sql_type("real"), PgTypeCategory::Real);
assert_eq!(
PgTypeCategory::from_sql_type("double precision"),
PgTypeCategory::DoublePrecision
);
assert_eq!(
PgTypeCategory::from_sql_type("macaddr8"),
PgTypeCategory::MacAddr8
);
assert_eq!(
PgTypeCategory::from_sql_type("macaddr"),
PgTypeCategory::MacAddr
);
}
#[test]
fn test_from_sql_type_postgis_surface() {
assert_eq!(
PgTypeCategory::from_sql_type("geometry(point)"),
PgTypeCategory::Geometry
);
assert_eq!(
PgTypeCategory::from_sql_type("geometry(point, 4326)"),
PgTypeCategory::Geometry
);
assert_eq!(
PgTypeCategory::from_sql_type("geometry(polygon, 4326)"),
PgTypeCategory::Custom
);
assert_eq!(
PgTypeCategory::from_sql_type("geography(point)"),
PgTypeCategory::Custom
);
assert_eq!(
PgTypeCategory::from_sql_type("box2d"),
PgTypeCategory::Custom
);
assert_eq!(
PgTypeCategory::from_sql_type("box3d"),
PgTypeCategory::Custom
);
assert_eq!(
PgTypeCategory::from_sql_type("raster"),
PgTypeCategory::Custom
);
}
#[test]
fn test_extract_nextval_sequence() {
assert_eq!(
extract_nextval_sequence("nextval('users_id_seq'::regclass)"),
Some("users_id_seq".to_string())
);
assert_eq!(
extract_nextval_sequence("nextval('public.users_id_seq'::regclass)"),
Some("users_id_seq".to_string())
);
assert_eq!(
extract_nextval_sequence("nextval('\"myschema\".\"users_id_seq\"'::regclass)"),
Some("users_id_seq".to_string())
);
assert_eq!(extract_nextval_sequence("not_a_nextval"), None);
assert_eq!(extract_nextval_sequence("nextval(''::regclass)"), None);
}
}