use crate::relational::error::{RelError, Result};
use crate::relational::types::SqlType;
use serde::de::Error as _;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use std::collections::{BTreeMap, HashMap};
pub const FIRST_USER_OID: u32 = 16384;
const SIDECAR_MARKER: &str = "@sidecar";
fn strip_sidecar_marker(version: &str) -> &str {
version.strip_suffix(SIDECAR_MARKER).unwrap_or(version)
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct QualifiedName {
pub schema: String,
pub name: String,
}
impl Serialize for QualifiedName {
fn serialize<S: Serializer>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error> {
serializer.serialize_str(&format!("{}\u{1f}{}", self.schema, self.name))
}
}
impl<'de> Deserialize<'de> for QualifiedName {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> std::result::Result<Self, D::Error> {
let s = String::deserialize(deserializer)?;
match s.split_once('\u{1f}') {
Some((schema, name)) => Ok(QualifiedName::new(schema, name)),
None => Err(D::Error::custom("malformed qualified name key")),
}
}
}
impl QualifiedName {
pub fn new(schema: impl Into<String>, name: impl Into<String>) -> Self {
Self {
schema: schema.into(),
name: name.into(),
}
}
pub fn to_string_qualified(&self) -> String {
format!("{}.{}", self.schema, self.name)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Schema {
pub name: String,
pub oid: u32,
pub owner: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Column {
pub name: String,
pub ty: SqlType,
pub nullable: bool,
pub default: Option<String>,
pub identity_sequence: Option<String>,
pub ordinal: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PrimaryKey {
pub name: String,
pub columns: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct UniqueConstraint {
pub name: String,
pub columns: Vec<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ReferentialAction {
NoAction,
Restrict,
Cascade,
SetNull,
SetDefault,
}
impl ReferentialAction {
pub fn as_sql(&self) -> &'static str {
match self {
ReferentialAction::NoAction => "NO ACTION",
ReferentialAction::Restrict => "RESTRICT",
ReferentialAction::Cascade => "CASCADE",
ReferentialAction::SetNull => "SET NULL",
ReferentialAction::SetDefault => "SET DEFAULT",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
pub enum MatchType {
#[default]
Simple,
Full,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
pub enum Deferrable {
#[default]
NotDeferrable,
DeferrableImmediate,
DeferrableDeferred,
}
impl Deferrable {
pub fn is_deferrable(self) -> bool {
!matches!(self, Deferrable::NotDeferrable)
}
pub fn initially_deferred(self) -> bool {
matches!(self, Deferrable::DeferrableDeferred)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ForeignKey {
pub name: String,
pub columns: Vec<String>,
pub ref_schema: String,
pub ref_table: String,
pub ref_columns: Vec<String>,
pub on_delete: ReferentialAction,
pub on_update: ReferentialAction,
#[serde(default)]
pub match_type: MatchType,
#[serde(default)]
pub deferrable: Deferrable,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CheckConstraint {
pub name: String,
pub expr: String,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum PolicyCmd {
All,
Select,
Insert,
Update,
Delete,
}
impl PolicyCmd {
pub fn as_sql(&self) -> &'static str {
match self {
PolicyCmd::All => "ALL",
PolicyCmd::Select => "SELECT",
PolicyCmd::Insert => "INSERT",
PolicyCmd::Update => "UPDATE",
PolicyCmd::Delete => "DELETE",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Policy {
pub name: String,
pub cmd: PolicyCmd,
pub roles: Vec<String>,
pub using_expr: Option<String>,
pub check_expr: Option<String>,
pub permissive: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum TriggerTiming {
Before,
After,
InsteadOf,
}
impl TriggerTiming {
pub fn as_sql(&self) -> &'static str {
match self {
TriggerTiming::Before => "BEFORE",
TriggerTiming::After => "AFTER",
TriggerTiming::InsteadOf => "INSTEAD OF",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum TriggerLevel {
Row,
Statement,
}
impl TriggerLevel {
pub fn as_sql(&self) -> &'static str {
match self {
TriggerLevel::Row => "ROW",
TriggerLevel::Statement => "STATEMENT",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum TriggerEventDef {
Insert,
Update {
columns: Vec<String>,
},
Delete,
Truncate,
}
fn default_true() -> bool {
true
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TriggerDef {
pub oid: u32,
pub name: String,
pub timing: TriggerTiming,
pub events: Vec<TriggerEventDef>,
pub level: TriggerLevel,
pub when_expr: Option<String>,
pub function_schema: String,
pub function_name: String,
#[serde(default = "default_true")]
pub enabled: bool,
#[serde(default)]
pub is_constraint: bool,
#[serde(default)]
pub deferrable: bool,
#[serde(default)]
pub initially_deferred: bool,
#[serde(default)]
pub referencing_old: Option<String>,
#[serde(default)]
pub referencing_new: Option<String>,
#[serde(default)]
pub on_view: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(from = "TableDe")]
pub struct Table {
pub oid: u32,
pub schema: String,
pub name: String,
pub columns: Vec<Column>,
pub primary_key: Option<PrimaryKey>,
pub uniques: Vec<UniqueConstraint>,
pub foreign_keys: Vec<ForeignKey>,
pub checks: Vec<CheckConstraint>,
pub storage_collection: String,
#[serde(default)]
pub rls_enabled: bool,
#[serde(default)]
pub rls_forced: bool,
#[serde(default)]
pub policies: Vec<Policy>,
#[serde(default)]
pub triggers: Vec<TriggerDef>,
#[serde(skip)]
pub(crate) column_map: HashMap<String, usize>,
}
#[derive(Deserialize)]
struct TableDe {
oid: u32,
schema: String,
name: String,
columns: Vec<Column>,
primary_key: Option<PrimaryKey>,
uniques: Vec<UniqueConstraint>,
foreign_keys: Vec<ForeignKey>,
checks: Vec<CheckConstraint>,
storage_collection: String,
#[serde(default)]
rls_enabled: bool,
#[serde(default)]
rls_forced: bool,
#[serde(default)]
policies: Vec<Policy>,
#[serde(default)]
triggers: Vec<TriggerDef>,
}
impl From<TableDe> for Table {
fn from(d: TableDe) -> Self {
let mut t = Table {
oid: d.oid,
schema: d.schema,
name: d.name,
columns: d.columns,
primary_key: d.primary_key,
uniques: d.uniques,
foreign_keys: d.foreign_keys,
checks: d.checks,
storage_collection: d.storage_collection,
rls_enabled: d.rls_enabled,
rls_forced: d.rls_forced,
policies: d.policies,
triggers: d.triggers,
column_map: HashMap::new(),
};
t.rebuild_column_map();
t
}
}
impl Table {
pub fn rebuild_column_map(&mut self) {
self.column_map.clear();
for (i, c) in self.columns.iter().enumerate() {
self.column_map.insert(c.name.clone(), i);
}
}
pub fn column(&self, name: &str) -> Option<&Column> {
self.column_map.get(name).map(|&i| &self.columns[i])
}
pub fn policy(&self, name: &str) -> Option<&Policy> {
self.policies.iter().find(|p| p.name == name)
}
pub fn trigger(&self, name: &str) -> Option<&TriggerDef> {
self.triggers.iter().find(|t| t.name == name)
}
pub fn column_mut(&mut self, name: &str) -> Option<&mut Column> {
let idx = *self.column_map.get(name)?;
Some(&mut self.columns[idx])
}
pub fn column_index(&self, name: &str) -> Option<usize> {
self.column_map.get(name).copied()
}
pub fn qualified(&self) -> QualifiedName {
QualifiedName::new(self.schema.clone(), self.name.clone())
}
pub fn pk_columns(&self) -> Vec<String> {
self.primary_key
.as_ref()
.map(|pk| pk.columns.clone())
.unwrap_or_default()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Index {
pub oid: u32,
pub name: String,
pub schema: String,
pub table: String,
pub columns: Vec<String>,
pub unique: bool,
pub primary: bool,
pub method: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Sequence {
pub schema: String,
pub name: String,
pub current: i64,
pub increment: i64,
pub start: i64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct View {
pub oid: u32,
pub schema: String,
pub name: String,
pub query: String,
pub columns: Vec<String>,
#[serde(default)]
pub triggers: Vec<TriggerDef>,
}
impl View {
pub fn trigger(&self, name: &str) -> Option<&TriggerDef> {
self.triggers.iter().find(|t| t.name == name)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum FunctionLanguage {
Sql,
PlPgSql,
}
impl FunctionLanguage {
pub fn as_sql(&self) -> &'static str {
match self {
FunctionLanguage::Sql => "sql",
FunctionLanguage::PlPgSql => "plpgsql",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum FunctionVolatility {
Immutable,
Stable,
Volatile,
}
impl FunctionVolatility {
pub fn as_char(&self) -> char {
match self {
FunctionVolatility::Immutable => 'i',
FunctionVolatility::Stable => 's',
FunctionVolatility::Volatile => 'v',
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FunctionArgDef {
pub name: String,
pub ty: SqlType,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FunctionDef {
pub oid: u32,
pub schema: String,
pub name: String,
pub args: Vec<FunctionArgDef>,
pub return_type: SqlType,
pub language: FunctionLanguage,
pub volatility: FunctionVolatility,
pub strict: bool,
pub body: String,
#[serde(default)]
pub returns_trigger: bool,
}
impl FunctionDef {
pub fn arity(&self) -> usize {
self.args.len()
}
pub fn qualified(&self) -> QualifiedName {
QualifiedName::new(self.schema.clone(), self.name.clone())
}
}
pub enum DropFunctionByName {
Removed,
NotFound,
Ambiguous,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TsDictionaryDef {
pub name: String,
pub schema: String,
pub oid: u32,
pub synonyms: BTreeMap<String, Vec<String>>,
#[serde(default)]
pub thesaurus_entries: BTreeMap<String, Vec<String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(from = "CatalogRaw")]
pub struct Catalog {
pub database: String,
schemas: BTreeMap<String, Schema>,
tables: BTreeMap<QualifiedName, Table>,
indexes: BTreeMap<QualifiedName, Index>,
sequences: BTreeMap<QualifiedName, Sequence>,
views: BTreeMap<QualifiedName, View>,
next_oid: u32,
pub search_path: Vec<String>,
#[serde(default)]
extensions: BTreeMap<String, String>,
#[serde(default)]
functions: Vec<FunctionDef>,
#[serde(default)]
ts_dictionaries: Vec<TsDictionaryDef>,
#[serde(skip)]
table_to_indexes: HashMap<(String, String), Vec<QualifiedName>>,
}
#[derive(Deserialize)]
struct CatalogRaw {
database: String,
schemas: BTreeMap<String, Schema>,
tables: BTreeMap<QualifiedName, Table>,
indexes: BTreeMap<QualifiedName, Index>,
sequences: BTreeMap<QualifiedName, Sequence>,
views: BTreeMap<QualifiedName, View>,
next_oid: u32,
search_path: Vec<String>,
#[serde(default)]
extensions: BTreeMap<String, String>,
#[serde(default)]
functions: Vec<FunctionDef>,
#[serde(default)]
ts_dictionaries: Vec<TsDictionaryDef>,
}
impl From<CatalogRaw> for Catalog {
fn from(raw: CatalogRaw) -> Self {
let mut table_to_indexes: HashMap<(String, String), Vec<QualifiedName>> = HashMap::new();
for (q, idx) in &raw.indexes {
table_to_indexes
.entry((idx.schema.clone(), idx.table.clone()))
.or_default()
.push(q.clone());
}
Catalog {
database: raw.database,
schemas: raw.schemas,
tables: raw.tables,
indexes: raw.indexes,
sequences: raw.sequences,
views: raw.views,
next_oid: raw.next_oid,
search_path: raw.search_path,
extensions: raw.extensions,
functions: raw.functions,
ts_dictionaries: raw.ts_dictionaries,
table_to_indexes,
}
}
}
impl Catalog {
pub fn new(database: impl Into<String>) -> Self {
let mut catalog = Self {
database: database.into(),
schemas: BTreeMap::new(),
tables: BTreeMap::new(),
indexes: BTreeMap::new(),
sequences: BTreeMap::new(),
views: BTreeMap::new(),
next_oid: FIRST_USER_OID,
search_path: vec!["public".to_string()],
extensions: BTreeMap::new(),
functions: Vec::new(),
ts_dictionaries: Vec::new(),
table_to_indexes: HashMap::new(),
};
catalog
.extensions
.insert("plpgsql".to_string(), "1.0".to_string());
for sys in ["pg_catalog", "information_schema"] {
let oid = catalog.allocate_oid();
catalog.schemas.insert(
sys.to_string(),
Schema {
name: sys.to_string(),
oid,
owner: "guardian".into(),
},
);
}
let oid = catalog.allocate_oid();
catalog.schemas.insert(
"public".to_string(),
Schema {
name: "public".into(),
oid,
owner: "guardian".into(),
},
);
catalog
}
pub fn extensions(&self) -> impl Iterator<Item = (&str, &str)> {
self.extensions
.iter()
.map(|(k, v)| (k.as_str(), strip_sidecar_marker(v)))
}
pub fn extension_version(&self, name: &str) -> Option<&str> {
self.extensions.get(name).map(|v| strip_sidecar_marker(v))
}
pub fn extension_installed(&self, name: &str) -> bool {
self.extensions.contains_key(name)
}
pub fn extension_is_sidecar(&self, name: &str) -> bool {
self.extensions
.get(name)
.map(|v| v.ends_with(SIDECAR_MARKER))
.unwrap_or(false)
}
pub fn install_extension(&mut self, name: &str, version: &str) -> bool {
self.extensions
.insert(name.to_string(), version.to_string())
.is_none()
}
pub fn install_sidecar_extension(&mut self, name: &str, version: &str) -> bool {
self.extensions
.insert(name.to_string(), format!("{version}{SIDECAR_MARKER}"))
.is_none()
}
pub fn uninstall_extension(&mut self, name: &str) -> bool {
self.extensions.remove(name).is_some()
}
pub fn set_extension_version(&mut self, name: &str, version: &str) -> bool {
match self.extensions.get_mut(name) {
Some(v) => {
*v = if v.ends_with(SIDECAR_MARKER) {
format!("{version}{SIDECAR_MARKER}")
} else {
version.to_string()
};
true
}
None => false,
}
}
pub fn allocate_oid(&mut self) -> u32 {
let oid = self.next_oid;
self.next_oid += 1;
oid
}
pub fn has_schema(&self, name: &str) -> bool {
self.schemas.contains_key(name)
}
pub fn schemas(&self) -> impl Iterator<Item = &Schema> {
self.schemas.values()
}
pub fn create_schema(&mut self, name: &str, if_not_exists: bool) -> Result<()> {
if self.schemas.contains_key(name) {
if if_not_exists {
return Ok(());
}
return Err(RelError::DuplicateSchema(name.to_string()));
}
let oid = self.allocate_oid();
self.schemas.insert(
name.to_string(),
Schema {
name: name.to_string(),
oid,
owner: "guardian".into(),
},
);
Ok(())
}
pub fn drop_schema(&mut self, name: &str, if_exists: bool, cascade: bool) -> Result<()> {
if !self.schemas.contains_key(name) {
if if_exists {
return Ok(());
}
return Err(RelError::UndefinedSchema(name.to_string()));
}
let table_names: Vec<QualifiedName> = self
.tables
.keys()
.filter(|k| k.schema == name)
.cloned()
.collect();
if !table_names.is_empty() && !cascade {
return Err(RelError::FeatureNotSupported(format!(
"cannot drop schema {name} because it contains objects (use CASCADE)"
)));
}
for t in table_names {
self.drop_table_qualified(&t)?;
}
self.schemas.remove(name);
Ok(())
}
pub fn resolve_table_name(&self, schema: Option<&str>, name: &str) -> Option<QualifiedName> {
if let Some(schema) = schema {
let q = QualifiedName::new(schema, name);
if self.tables.contains_key(&q) || self.views.contains_key(&q) {
return Some(q);
}
return None;
}
for schema in &self.search_path {
let q = QualifiedName::new(schema.clone(), name);
if self.tables.contains_key(&q) || self.views.contains_key(&q) {
return Some(q);
}
}
None
}
pub fn creation_schema(&self, schema: Option<&str>) -> Result<String> {
match schema {
Some(s) => {
if !self.schemas.contains_key(s) {
return Err(RelError::UndefinedSchema(s.to_string()));
}
Ok(s.to_string())
}
None => Ok(self
.search_path
.first()
.cloned()
.unwrap_or_else(|| "public".to_string())),
}
}
pub fn tables(&self) -> impl Iterator<Item = &Table> {
self.tables.values()
}
pub fn get_table(&self, q: &QualifiedName) -> Option<&Table> {
self.tables.get(q)
}
pub fn get_table_mut(&mut self, q: &QualifiedName) -> Option<&mut Table> {
self.tables.get_mut(q)
}
pub fn require_table(&self, q: &QualifiedName) -> Result<&Table> {
self.tables
.get(q)
.ok_or_else(|| RelError::UndefinedTable(q.to_string_qualified()))
}
pub fn has_table(&self, q: &QualifiedName) -> bool {
self.tables.contains_key(q)
}
pub fn insert_table(&mut self, mut table: Table) -> Result<()> {
let q = table.qualified();
if self.tables.contains_key(&q) || self.views.contains_key(&q) {
return Err(RelError::DuplicateTable(q.to_string_qualified()));
}
if !self.schemas.contains_key(&table.schema) {
return Err(RelError::UndefinedSchema(table.schema.clone()));
}
if table.storage_collection.is_empty() {
table.storage_collection = format!("__gdb_sql_rows_{}", table.oid);
}
table.rebuild_column_map();
self.tables.insert(q, table);
Ok(())
}
pub fn referencing_foreign_keys(&self, q: &QualifiedName) -> Vec<(QualifiedName, ForeignKey)> {
let mut out = Vec::new();
for table in self.tables.values() {
for fk in &table.foreign_keys {
if fk.ref_schema == q.schema && fk.ref_table == q.name {
out.push((table.qualified(), fk.clone()));
}
}
}
out
}
pub fn foreign_keys_named(
&self,
schema: Option<&str>,
name: &str,
) -> Vec<(QualifiedName, ForeignKey)> {
let schemas: Vec<&str> = match schema {
Some(s) => vec![s],
None => self.search_path.iter().map(String::as_str).collect(),
};
for s in schemas {
let found: Vec<(QualifiedName, ForeignKey)> = self
.tables
.values()
.filter(|t| t.schema == s)
.flat_map(|t| {
t.foreign_keys
.iter()
.filter(|fk| fk.name == name)
.map(move |fk| (t.qualified(), fk.clone()))
})
.collect();
if !found.is_empty() {
return found;
}
}
Vec::new()
}
pub fn drop_table_qualified(&mut self, q: &QualifiedName) -> Result<Table> {
let table = self
.tables
.remove(q)
.ok_or_else(|| RelError::UndefinedTable(q.to_string_qualified()))?;
for t in self.tables.values_mut() {
t.foreign_keys
.retain(|fk| !(fk.ref_schema == q.schema && fk.ref_table == q.name));
}
let idx_keys: Vec<QualifiedName> = self
.indexes
.iter()
.filter(|(_, i)| i.schema == q.schema && i.table == q.name)
.map(|(k, _)| k.clone())
.collect();
for k in idx_keys {
self.indexes.remove(&k);
}
self.table_to_indexes
.remove(&(q.schema.clone(), q.name.clone()));
for col in &table.columns {
if let Some(seq) = &col.identity_sequence {
let sk = QualifiedName::new(q.schema.clone(), seq.clone());
self.sequences.remove(&sk);
}
}
Ok(table)
}
pub fn indexes(&self) -> impl Iterator<Item = &Index> {
self.indexes.values()
}
pub fn indexes_for_table(&self, schema: &str, table: &str) -> Vec<&Index> {
let key = (schema.to_string(), table.to_string());
self.table_to_indexes
.get(&key)
.map(|names| names.iter().filter_map(|q| self.indexes.get(q)).collect())
.unwrap_or_default()
}
pub fn get_index(&self, q: &QualifiedName) -> Option<&Index> {
self.indexes.get(q)
}
pub fn insert_index(&mut self, index: Index) -> Result<()> {
let q = QualifiedName::new(index.schema.clone(), index.name.clone());
if self.indexes.contains_key(&q) {
return Err(RelError::DuplicateIndex(q.to_string_qualified()));
}
self.table_to_indexes
.entry((index.schema.clone(), index.table.clone()))
.or_default()
.push(q.clone());
self.indexes.insert(q, index);
Ok(())
}
pub fn drop_index(&mut self, schema: Option<&str>, name: &str, if_exists: bool) -> Result<()> {
let q = match schema {
Some(s) => QualifiedName::new(s, name),
None => {
let found = self
.search_path
.iter()
.map(|s| QualifiedName::new(s.clone(), name))
.find(|q| self.indexes.contains_key(q));
match found {
Some(q) => q,
None => {
if if_exists {
return Ok(());
}
return Err(RelError::UndefinedIndex(name.to_string()));
}
}
}
};
if let Some(removed) = self.indexes.remove(&q) {
let tkey = (removed.schema.clone(), removed.table.clone());
if let Some(v) = self.table_to_indexes.get_mut(&tkey) {
v.retain(|iq| iq != &q);
if v.is_empty() {
self.table_to_indexes.remove(&tkey);
}
}
} else if !if_exists {
return Err(RelError::UndefinedIndex(q.to_string_qualified()));
}
Ok(())
}
pub fn sequences(&self) -> impl Iterator<Item = &Sequence> {
self.sequences.values()
}
pub fn create_sequence(&mut self, schema: &str, name: &str) -> Result<()> {
let q = QualifiedName::new(schema, name);
self.sequences.entry(q).or_insert(Sequence {
schema: schema.to_string(),
name: name.to_string(),
current: 0,
increment: 1,
start: 1,
});
Ok(())
}
pub fn next_sequence_value(&mut self, schema: &str, name: &str) -> Result<i64> {
let q = QualifiedName::new(schema, name);
let seq = self
.sequences
.get_mut(&q)
.ok_or_else(|| RelError::UndefinedObject(format!("sequence {schema}.{name}")))?;
let next = if seq.current == 0 {
seq.start
} else {
seq.current + seq.increment
};
seq.current = next;
Ok(next)
}
pub fn observe_sequence_value(&mut self, schema: &str, name: &str, value: i64) {
let q = QualifiedName::new(schema, name);
if let Some(seq) = self.sequences.get_mut(&q)
&& value > seq.current
{
seq.current = value;
}
}
pub fn views(&self) -> impl Iterator<Item = &View> {
self.views.values()
}
pub fn get_view(&self, q: &QualifiedName) -> Option<&View> {
self.views.get(q)
}
pub fn get_view_mut(&mut self, q: &QualifiedName) -> Option<&mut View> {
self.views.get_mut(q)
}
pub fn insert_view(&mut self, view: View) -> Result<()> {
let q = QualifiedName::new(view.schema.clone(), view.name.clone());
if self.tables.contains_key(&q) || self.views.contains_key(&q) {
return Err(RelError::DuplicateTable(q.to_string_qualified()));
}
self.views.insert(q, view);
Ok(())
}
pub fn drop_view(&mut self, q: &QualifiedName, if_exists: bool) -> Result<()> {
if self.views.remove(q).is_none() && !if_exists {
return Err(RelError::UndefinedTable(q.to_string_qualified()));
}
Ok(())
}
pub fn functions(&self) -> impl Iterator<Item = &FunctionDef> {
self.functions.iter()
}
fn resolve_function_index(
&self,
schema: Option<&str>,
name: &str,
arity: usize,
) -> Option<usize> {
if let Some(s) = schema {
return self
.functions
.iter()
.position(|f| f.schema == s && f.name == name && f.arity() == arity);
}
for s in &self.search_path {
if let Some(i) = self
.functions
.iter()
.position(|f| &f.schema == s && f.name == name && f.arity() == arity)
{
return Some(i);
}
}
None
}
pub fn find_function(
&self,
schema: Option<&str>,
name: &str,
arity: usize,
) -> Option<&FunctionDef> {
self.resolve_function_index(schema, name, arity)
.map(|i| &self.functions[i])
}
pub fn insert_function(&mut self, def: FunctionDef) -> Result<()> {
if self
.functions
.iter()
.any(|f| f.schema == def.schema && f.name == def.name && f.arity() == def.arity())
{
return Err(RelError::DuplicateFunction(format!(
"function \"{}\" already exists with same argument count",
def.name
)));
}
self.functions.push(def);
Ok(())
}
pub fn replace_function(&mut self, mut def: FunctionDef) {
match self
.functions
.iter()
.position(|f| f.schema == def.schema && f.name == def.name && f.arity() == def.arity())
{
Some(i) => {
def.oid = self.functions[i].oid;
self.functions[i] = def;
}
None => self.functions.push(def),
}
}
pub fn drop_function(&mut self, schema: Option<&str>, name: &str, arity: usize) -> bool {
match self.resolve_function_index(schema, name, arity) {
Some(i) => {
self.functions.remove(i);
true
}
None => false,
}
}
pub fn drop_function_by_name(
&mut self,
schema: Option<&str>,
name: &str,
) -> DropFunctionByName {
let matches: Vec<usize> = self
.functions
.iter()
.enumerate()
.filter(|(_, f)| f.name == name && schema.map(|s| f.schema == s).unwrap_or(true))
.map(|(i, _)| i)
.collect();
match matches.len() {
0 => DropFunctionByName::NotFound,
1 => {
self.functions.remove(matches[0]);
DropFunctionByName::Removed
}
_ => DropFunctionByName::Ambiguous,
}
}
pub fn ts_dictionaries(&self) -> impl Iterator<Item = &TsDictionaryDef> {
self.ts_dictionaries.iter()
}
pub fn insert_ts_dictionary(
&mut self,
def: TsDictionaryDef,
if_not_exists: bool,
) -> Result<()> {
if self
.ts_dictionaries
.iter()
.any(|d| d.schema == def.schema && d.name == def.name)
{
if if_not_exists {
return Ok(());
}
return Err(RelError::DuplicateObject(format!(
"text search dictionary \"{}.{}\"",
def.schema, def.name
)));
}
self.ts_dictionaries.push(def);
Ok(())
}
pub fn drop_ts_dictionary(
&mut self,
schema: Option<&str>,
name: &str,
if_exists: bool,
) -> Result<bool> {
let pos = self
.ts_dictionaries
.iter()
.position(|d| d.name == name && schema.map(|s| d.schema == s).unwrap_or(true));
match pos {
Some(i) => {
self.ts_dictionaries.remove(i);
Ok(true)
}
None => {
if if_exists {
Ok(false)
} else {
Err(RelError::UndefinedObject(format!(
"text search dictionary \"{name}\""
)))
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_table(cat: &mut Catalog) -> Table {
let oid = cat.allocate_oid();
Table {
oid,
schema: "public".into(),
name: "users".into(),
columns: vec![
Column {
name: "id".into(),
ty: SqlType::Integer,
nullable: false,
default: None,
identity_sequence: None,
ordinal: 0,
},
Column {
name: "email".into(),
ty: SqlType::Text,
nullable: false,
default: None,
identity_sequence: None,
ordinal: 1,
},
],
primary_key: Some(PrimaryKey {
name: "users_pkey".into(),
columns: vec!["id".into()],
}),
uniques: vec![],
foreign_keys: vec![],
checks: vec![],
storage_collection: String::new(),
rls_enabled: false,
rls_forced: false,
policies: vec![],
triggers: vec![],
column_map: HashMap::new(),
}
}
#[test]
fn create_and_resolve_table() {
let mut cat = Catalog::new("app");
let t = sample_table(&mut cat);
cat.insert_table(t).unwrap();
let q = cat.resolve_table_name(None, "users").unwrap();
assert_eq!(q.schema, "public");
assert!(
cat.get_table(&q)
.unwrap()
.storage_collection
.starts_with("__gdb_sql_rows_")
);
}
#[test]
fn duplicate_table_errors() {
let mut cat = Catalog::new("app");
let t = sample_table(&mut cat);
cat.insert_table(t.clone()).unwrap();
let t2 = sample_table(&mut cat);
assert!(matches!(
cat.insert_table(t2),
Err(RelError::DuplicateTable(_))
));
}
#[test]
fn sequence_advances() {
let mut cat = Catalog::new("app");
cat.create_sequence("public", "users_id_seq").unwrap();
assert_eq!(
cat.next_sequence_value("public", "users_id_seq").unwrap(),
1
);
assert_eq!(
cat.next_sequence_value("public", "users_id_seq").unwrap(),
2
);
cat.observe_sequence_value("public", "users_id_seq", 10);
assert_eq!(
cat.next_sequence_value("public", "users_id_seq").unwrap(),
11
);
}
#[test]
fn drop_schema_requires_cascade() {
let mut cat = Catalog::new("app");
cat.create_schema("app", false).unwrap();
let oid = cat.allocate_oid();
let mut t = sample_table(&mut cat);
t.schema = "app".into();
t.oid = oid;
cat.insert_table(t).unwrap();
assert!(cat.drop_schema("app", false, false).is_err());
assert!(cat.drop_schema("app", false, true).is_ok());
}
#[test]
fn sidecar_extension_marker_round_trips() {
let mut cat = Catalog::new("app");
cat.install_sidecar_extension("pg_stat_statements", "1.10");
assert!(cat.extension_installed("pg_stat_statements"));
assert!(cat.extension_is_sidecar("pg_stat_statements"));
assert_eq!(cat.extension_version("pg_stat_statements"), Some("1.10"));
assert!(
cat.extensions()
.any(|(n, v)| n == "pg_stat_statements" && v == "1.10")
);
cat.set_extension_version("pg_stat_statements", "1.11");
assert!(cat.extension_is_sidecar("pg_stat_statements"));
assert_eq!(cat.extension_version("pg_stat_statements"), Some("1.11"));
assert!(!cat.extension_is_sidecar("plpgsql"));
let json = serde_json::to_value(&cat).unwrap();
let back: Catalog = serde_json::from_value(json).unwrap();
assert!(back.extension_is_sidecar("pg_stat_statements"));
assert_eq!(back.extension_version("plpgsql"), Some("1.0"));
}
#[test]
fn catalog_round_trips_json() {
let mut cat = Catalog::new("app");
let t = sample_table(&mut cat);
cat.insert_table(t).unwrap();
let json = serde_json::to_value(&cat).unwrap();
let back: Catalog = serde_json::from_value(json).unwrap();
assert!(back.resolve_table_name(None, "users").is_some());
}
}