fix(postgres): include extensions in database exports

This commit is contained in:
t8y2 2026-07-19 12:48:15 +08:00
parent 9679d48824
commit 17e54bd5eb
2 changed files with 219 additions and 16 deletions

View File

@ -93,6 +93,18 @@ struct PostgresExportSequence {
owner_column: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct PostgresExportExtension {
name: String,
schema: String,
}
#[derive(Debug, Default)]
struct PostgresExtensionMembers {
relation_names: HashSet<String>,
function_keys: HashSet<(String, String)>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ExportedTableSql {
@ -888,6 +900,43 @@ fn generate_postgres_sequence_setval_sql(sequence: &PostgresExportSequence, sche
}
}
fn generate_postgres_extension_ddl(extension: &PostgresExportExtension) -> String {
// Match pg_dump: omit VERSION so the target installation selects its
// default compatible version, while preserving the source schema.
format!(
"CREATE EXTENSION IF NOT EXISTS {} WITH SCHEMA {};",
quote_identifier(&extension.name, &DatabaseType::Postgres),
quote_identifier(&extension.schema, &DatabaseType::Postgres)
)
}
async fn list_postgres_extension_members(
state: &crate::connection::AppState,
pool_key: &str,
schema: &str,
) -> Result<PostgresExtensionMembers, String> {
let pool = {
let connections = state.connections.read().await;
match connections.get(pool_key) {
Some(crate::connection::PoolKind::Postgres(pool)) => pool.clone(),
_ => return Ok(PostgresExtensionMembers::default()),
}
};
let mut members = PostgresExtensionMembers::default();
for (kind, name, signature) in crate::db::postgres::list_extension_member_objects(&pool, schema).await? {
if kind == "RELATION" {
members.relation_names.insert(name);
} else if kind == "FUNCTION" {
members.function_keys.insert((name, signature));
}
}
Ok(members)
}
fn is_postgres_extension_member_routine(object: &crate::types::ObjectInfo, members: &PostgresExtensionMembers) -> bool {
members.function_keys.contains(&(object.name.clone(), object.signature.clone().unwrap_or_default()))
}
async fn list_postgres_export_sequences(
state: &crate::connection::AppState,
pool_key: &str,
@ -1144,8 +1193,6 @@ pub async fn export_database_sql_core(
None,
)
.await?;
let all_tables = filter_export_table_infos(all_tables, &request.selected_tables, &request.excluded_tables);
// 4. Create file
let mut file = std::fs::File::create(&request.file_path).map_err(|e| format!("Failed to write file: {e}"))?;
@ -1162,6 +1209,41 @@ pub async fn export_database_sql_core(
}
// 7. Separate tables and views
let postgres_extension_members =
if matches!(db_type, DatabaseType::Postgres) && (request.include_structure || request.include_objects) {
match list_postgres_extension_members(state, &pool_key, &request.schema).await {
Ok(members) => members,
Err(e) => {
record_export_error(&mut file, request.fail_on_error, format!("reading extension members: {e}"))?;
PostgresExtensionMembers::default()
}
}
} else {
PostgresExtensionMembers::default()
};
let postgres_extensions = if request.include_structure && matches!(db_type, DatabaseType::Postgres) {
match crate::schema::list_extensions_core(state, &request.connection_id, &request.database, &request.schema)
.await
{
Ok(extensions) => extensions
.into_iter()
.map(|extension| PostgresExportExtension {
name: extension.name,
schema: extension.schema.unwrap_or_else(|| request.schema.clone()),
})
.collect(),
Err(e) => {
record_export_error(&mut file, request.fail_on_error, format!("exporting extensions: {e}"))?;
Vec::new()
}
}
} else {
Vec::new()
};
let all_tables = filter_export_table_infos(all_tables, &request.selected_tables, &request.excluded_tables)
.into_iter()
.filter(|table| !postgres_extension_members.relation_names.contains(&table.name))
.collect::<Vec<_>>();
let mut tables: Vec<_> = all_tables.iter().filter(|t| !t.table_type.contains("VIEW")).collect();
let views: Vec<_> = all_tables.iter().filter(|t| t.table_type.contains("VIEW")).collect();
let postgres_sequences = if request.include_structure && matches!(db_type, DatabaseType::Postgres) {
@ -1211,11 +1293,11 @@ pub async fn export_database_sql_core(
}
// 8. Calculate total objects
let mut total_objects = tables.len() + views.len() + postgres_sequences.len();
let mut total_objects = tables.len() + views.len() + postgres_sequences.len() + postgres_extensions.len();
// We'll add procedures/functions count later if include_objects
let mut procedures: Vec<String> = Vec::new();
let mut functions: Vec<String> = Vec::new();
let mut procedures: Vec<crate::types::ObjectInfo> = Vec::new();
let mut functions: Vec<crate::types::ObjectInfo> = Vec::new();
if request.include_objects && request.selected_tables.is_empty() {
match crate::schema::list_objects_core(
@ -1233,10 +1315,13 @@ pub async fn export_database_sql_core(
Ok(objects) => {
for obj in &objects {
let ot = obj.object_type.to_uppercase();
if is_postgres_extension_member_routine(obj, &postgres_extension_members) {
continue;
}
if ot.contains("PROCEDURE") {
procedures.push(obj.name.clone());
procedures.push(obj.clone());
} else if ot.contains("FUNCTION") {
functions.push(obj.name.clone());
functions.push(obj.clone());
}
}
}
@ -1251,6 +1336,25 @@ pub async fn export_database_sql_core(
// Export tables
let batch_size = if request.batch_size == 0 { 1000 } else { request.batch_size };
for extension in &postgres_extensions {
if is_export_cancelled(&request.export_id).await {
return Err("Export cancelled".to_string());
}
on_progress(ExportProgress {
export_id: request.export_id.clone(),
current_object: extension.name.clone(),
object_index,
total_objects,
rows_exported: 0,
total_rows: None,
status: ExportStatus::Running,
error: None,
});
writeln!(file, "{}\n", generate_postgres_extension_ddl(extension))
.map_err(|e| format!("Failed to write file: {e}"))?;
object_index += 1;
}
for sequence in postgres_sequences.iter().filter(|sequence| sequence.owner_table.is_none()) {
if is_export_cancelled(&request.export_id).await {
return Err("Export cancelled".to_string());
@ -1669,11 +1773,13 @@ pub async fn export_database_sql_core(
}
// Export procedures
for proc_name in &procedures {
for procedure in &procedures {
if is_export_cancelled(&request.export_id).await {
return Err("Export cancelled".to_string());
}
let proc_name = &procedure.name;
on_progress(ExportProgress {
export_id: request.export_id.clone(),
current_object: proc_name.clone(),
@ -1692,7 +1798,7 @@ pub async fn export_database_sql_core(
&request.schema,
proc_name,
crate::db::ObjectSourceKind::Procedure,
None,
procedure.signature.as_deref(),
)
.await
{
@ -1719,11 +1825,13 @@ pub async fn export_database_sql_core(
}
// Export functions
for func_name in &functions {
for function in &functions {
if is_export_cancelled(&request.export_id).await {
return Err("Export cancelled".to_string());
}
let func_name = &function.name;
on_progress(ExportProgress {
export_id: request.export_id.clone(),
current_object: func_name.clone(),
@ -1742,7 +1850,7 @@ pub async fn export_database_sql_core(
&request.schema,
func_name,
crate::db::ObjectSourceKind::Function,
None,
function.signature.as_deref(),
)
.await
{
@ -1812,13 +1920,14 @@ mod tests {
use super::concurrent_metadata_prefetch_allowed;
use super::{
build_database_sql_export, build_export_insert_statements, drop_table_if_exists_sql, filter_export_table_infos,
format_export_sql_literal, generate_postgres_sequence_create_ddl, generate_postgres_sequence_owner_ddl,
generate_postgres_sequence_setval_sql, normalize_export_table_ddl, record_export_error,
BuildDatabaseSqlExportOptions, BuildExportInsertStatementsOptions, ExportedTableSql, PostgresExportSequence,
DATABASE_EXPORT_INSERT_BATCH_SIZE, DATABASE_EXPORT_ROW_LIMIT,
format_export_sql_literal, generate_postgres_extension_ddl, generate_postgres_sequence_create_ddl,
generate_postgres_sequence_owner_ddl, generate_postgres_sequence_setval_sql,
is_postgres_extension_member_routine, normalize_export_table_ddl, record_export_error,
BuildDatabaseSqlExportOptions, BuildExportInsertStatementsOptions, ExportedTableSql, PostgresExportExtension,
PostgresExportSequence, PostgresExtensionMembers, DATABASE_EXPORT_INSERT_BATCH_SIZE, DATABASE_EXPORT_ROW_LIMIT,
};
use crate::models::connection::DatabaseType;
use crate::types::TableInfo;
use crate::types::{ObjectInfo, TableInfo};
use serde_json::{json, Value};
fn table(name: &str, table_type: &str) -> TableInfo {
@ -1831,6 +1940,41 @@ mod tests {
}
}
fn routine(name: &str, signature: &str) -> ObjectInfo {
ObjectInfo {
name: name.to_string(),
object_type: "FUNCTION".to_string(),
schema: Some("public".to_string()),
valid: None,
signature: Some(signature.to_string()),
comment: None,
created_at: None,
updated_at: None,
parent_schema: None,
parent_name: None,
}
}
#[test]
fn postgres_extension_ddl_uses_target_default_version_and_source_schema() {
let extension = PostgresExportExtension { name: "pg_trgm".to_string(), schema: "addons".to_string() };
let ddl = generate_postgres_extension_ddl(&extension);
assert_eq!(ddl, "CREATE EXTENSION IF NOT EXISTS \"pg_trgm\" WITH SCHEMA \"addons\";");
assert!(!ddl.contains("VERSION"));
}
#[test]
fn postgres_extension_member_filter_keeps_user_overload_with_same_name() {
let mut members = PostgresExtensionMembers::default();
members.function_keys.insert(("similarity".to_string(), "text, text".to_string()));
assert!(is_postgres_extension_member_routine(&routine("similarity", "text, text"), &members));
assert!(!is_postgres_extension_member_routine(&routine("similarity", "integer, integer"), &members));
assert!(!is_postgres_extension_member_routine(&routine("user_similarity", "text, text"), &members));
}
#[test]
fn concurrent_prefetch_only_allowed_for_multi_connection_pools() {
use crate::connection::PoolKind;

View File

@ -3243,6 +3243,53 @@ pub async fn list_extensions(pool: &Pool, schema: &str) -> Result<Vec<ExtensionI
.collect())
}
fn list_extension_member_objects_sql() -> &'static str {
"SELECT 'RELATION'::text AS object_kind, c.relname, ''::text AS signature \
FROM pg_catalog.pg_class c \
JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace \
WHERE n.nspname = $1 \
AND EXISTS ( \
SELECT 1 FROM pg_catalog.pg_depend d \
WHERE d.classid = 'pg_catalog.pg_class'::regclass \
AND d.objid = c.oid \
AND d.refclassid = 'pg_catalog.pg_extension'::regclass \
AND d.deptype = 'e' \
) \
UNION ALL \
SELECT 'FUNCTION'::text, p.proname, pg_get_function_identity_arguments(p.oid) \
FROM pg_catalog.pg_proc p \
JOIN pg_catalog.pg_namespace n ON n.oid = p.pronamespace \
WHERE n.nspname = $1 \
AND EXISTS ( \
SELECT 1 FROM pg_catalog.pg_depend d \
WHERE d.classid = 'pg_catalog.pg_proc'::regclass \
AND d.objid = p.oid \
AND d.refclassid = 'pg_catalog.pg_extension'::regclass \
AND d.deptype = 'e' \
)"
}
pub async fn list_extension_member_objects(pool: &Pool, schema: &str) -> Result<Vec<(String, String, String)>, String> {
let client = checkout_postgres_client(pool, None, super::connection_timeout()).await?;
let rows = match postgres_query_cached(&client, list_extension_member_objects_sql(), &[&schema]).await {
Ok(rows) => rows,
Err(primary_error) => {
// PostgreSQL-compatible servers before the identity-argument
// formatter can still be filtered using their legacy formatter.
let fallback_sql = list_extension_member_objects_sql()
.replace("pg_get_function_identity_arguments(p.oid)", "pg_get_function_arguments(p.oid)");
postgres_query_cached(&client, &fallback_sql, &[&schema])
.await
.map_err(|fallback_error| format!("{primary_error}; legacy fallback failed: {fallback_error}"))?
}
};
Ok(rows
.iter()
.map(|row| (pg_row_try_string(row, 0), pg_row_try_string(row, 1), pg_row_try_string(row, 2)))
.collect())
}
pub async fn list_available_extensions(pool: &Pool) -> Result<Vec<ExtensionInfo>, String> {
let client = checkout_postgres_client(pool, None, super::connection_timeout()).await?;
let rows = postgres_query_cached(
@ -3972,6 +4019,18 @@ mod tests {
assert!(!POSTGRES_COLUMNS_INFORMATION_SCHEMA_SQL.contains("regclass"));
}
#[test]
fn extension_member_query_filters_only_owned_relations_and_routines() {
let sql = list_extension_member_objects_sql();
assert!(sql.contains("d.classid = 'pg_catalog.pg_class'::regclass"));
assert!(sql.contains("d.classid = 'pg_catalog.pg_proc'::regclass"));
assert!(sql.contains("d.refclassid = 'pg_catalog.pg_extension'::regclass"));
assert!(sql.contains("d.deptype = 'e'"));
assert!(sql.contains("pg_get_function_identity_arguments(p.oid)"));
assert!(!sql.contains("d.deptype = 'x'"));
}
#[tokio::test]
async fn postgres_column_metadata_query_returns_enum_values_against_real_postgres() {
let Some(container) = start_docker_postgres().await else {