diff --git a/docs/README.md b/docs/README.md index 1c70c2d..9daef1b 100644 --- a/docs/README.md +++ b/docs/README.md @@ -84,10 +84,15 @@ Features: * **multiple file-sets** supported. This means you can choose to only sync your database, but not your binary resources/assets. * **Speed Optimized**: publicly available binary assets are not zipped extra; but the already-public files are simply downloaded. Resources which already exist locally and have the same file size and modification date are never re-downloaded. -* **no extra SQL client needed**: We package a custom implementation of `mysqldump` into the binary. +* **no extra SQL client needed**: We package a custom implementation of `mysqldump`/`pg_dump` into the binary - no + `mysqldump`, `pg_dump` or `psql` binary is ever shelled out to. * currently supported databases: - * **MySQL** - * (Postgres support planned) + * **MySQL / MariaDB** + * **PostgreSQL** - full schema reconstruction (columns/types, NOT NULL, defaults, PK/UNIQUE/CHECK/FK constraints, + indexes, sequence/identity state) from `pg_catalog`, with data streamed via `COPY ... FROM stdin`. Not (yet) + dumped: `CREATE TYPE`/`CREATE EXTENSION` (enum *values* round-trip, but the type/extension must already exist + on the target), non-`public` schemas, materialized views, partitioned tables, table inheritance, + triggers/functions, comments, and GRANTs/ownership. * **auto-cleanup**: remove dumps when tool is stopped # Installation diff --git a/docs/architecture.md b/docs/architecture.md index 6366d3c..5f75d18 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -56,8 +56,11 @@ The destination server will only fetch data from the source; and never the other on the source server (e.g. the production system), the user needs to log in, and then invoke the `synco-lite` executable. This does: * install synco-source via a shell script -* Detect which framework is used. E.g. for Flow/Neos or Symfony, synco then knows how to create a database dump (e.g. for MySQL, using go-mysqldump or - pingcap dumpling; and for Postgres some pgx based solution??) +* Detect which framework is used. E.g. for Flow/Neos or Symfony, synco then knows how to create a database dump + (for MySQL/MariaDB, using our vendored `go_mysqldump`; for PostgreSQL, using `pkg/util/postgres` + `pgx`, which + reconstructs full schema DDL from `pg_catalog` and streams data via `COPY ... FROM stdin` - see the "known + limitations" note in `docs/README.md` for what isn't (yet) covered on the Postgres side: custom types/extensions, + non-`public` schemas, materialized views, partitioned tables, triggers/functions, comments, GRANTs) * Publish a metadata file and encrypt it which shows the current status. * Create the database dump and encrypt it. * Create a file mapping for data/persistent in flow - as we do not need to re-compress static assets which are available online. diff --git a/go.mod b/go.mod index 8f49029..aac5e39 100644 --- a/go.mod +++ b/go.mod @@ -10,6 +10,7 @@ require ( github.com/dop251/goja v0.0.0-20221003171542-5ea1285e6c91 github.com/dustin/go-humanize v1.0.1 github.com/go-sql-driver/mysql v1.10.0 + github.com/jackc/pgx/v5 v5.10.0 github.com/jamf/go-mysqldump v0.7.1 github.com/logrusorgru/aurora v2.0.3+incompatible github.com/manifoldco/promptui v0.9.0 @@ -57,8 +58,12 @@ require ( github.com/google/uuid v1.6.0 // indirect github.com/gookit/color v1.6.0 // indirect github.com/inconshreveable/mousetrap v1.0.1 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect github.com/josharian/intern v1.0.0 // indirect github.com/json-iterator/go v1.1.12 // indirect + github.com/lib/pq v1.10.7 // indirect github.com/lithammer/fuzzysearch v1.1.8 // indirect github.com/lucasb-eyer/go-colorful v1.2.0 // indirect github.com/mailru/easyjson v0.7.7 // indirect diff --git a/go.sum b/go.sum index 1b46982..f5eb736 100644 --- a/go.sum +++ b/go.sum @@ -108,6 +108,14 @@ github.com/gookit/color v1.6.0 h1:JjJXBTk1ETNyqyilJhkTXJYYigHG24TM9Xa2M1xAhRA= github.com/gookit/color v1.6.0/go.mod h1:9ACFc7/1IpHGBW8RwuDm/0YEnhg3dwwXpoMsmtyHfjs= github.com/inconshreveable/mousetrap v1.0.1 h1:U3uMjPSQEBMNp1lFxmllqCPM6P5u/Xq7Pgzkat/bFNc= github.com/inconshreveable/mousetrap v1.0.1/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0= +github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= github.com/jamf/go-mysqldump v0.7.1 h1:JuEjzzKX51Bn9urjciXSqvmCGAxAwH3IaN3+nuplf+o= github.com/jamf/go-mysqldump v0.7.1/go.mod h1:YWqhOv9PfioqsO59t/DziO8gFEHw8G2vV6qBlFCdHIM= github.com/josharian/intern v1.0.0 h1:vlS4z54oSdjm0bgjRigI+G1HpF+tI+9rE5LLzOg8HmY= @@ -130,6 +138,8 @@ github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/lib/pq v1.10.7 h1:p7ZhMD+KsSRozJr34udlUrhboJwWAgCg34+/ZZNvZZw= +github.com/lib/pq v1.10.7/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/lithammer/fuzzysearch v1.1.8 h1:/HIuJnjHuXS8bKaiTMeeDlW2/AyIWk2brx1V8LFgLN4= github.com/lithammer/fuzzysearch v1.1.8/go.mod h1:IdqeyBClc3FFqSzYq/MXESsS4S0FsZ5ajtkr5xPLts4= github.com/logrusorgru/aurora v2.0.3+incompatible h1:tOpm7WcpBTn4fjmVfgpQq0EfczGlG91VSDkswnjF5A8= diff --git a/pkg/common/commonServe/databaseDump.go b/pkg/common/commonServe/databaseDump.go index c6473f5..f82a508 100644 --- a/pkg/common/commonServe/databaseDump.go +++ b/pkg/common/commonServe/databaseDump.go @@ -7,31 +7,56 @@ import ( "github.com/sandstorm/synco/v2/pkg/common/dto" "github.com/sandstorm/synco/v2/pkg/serve" "github.com/sandstorm/synco/v2/pkg/util/mysql" + "github.com/sandstorm/synco/v2/pkg/util/postgres" ) +// fileSetTypeFor maps a DB driver to the dto.FileSetType used to record its +// dump in the transfer session's metadata. Unrecognized/unset drivers +// default to mysql, matching the pre-Postgres-support behavior. +func fileSetTypeFor(driver common.DbDriver) dto.FileSetType { + if driver == common.DbDriverPostgres { + return dto.TYPE_POSTGRESDUMP + } + return dto.TYPE_MYSQLDUMP +} + func DatabaseDump(transferSession *serve.TransferSession, dbCredentials *common.DbCredentials, whereClauseForTables map[string]string) *sql.DB { // 2) DATABASE DUMP // basically the way it works is: - // mysql.CreateDump --> age.Encrypt --> write to file. + // mysql.CreateDump / postgres.CreateDump --> age.Encrypt --> write to file. // but because this is based on streams, we need to construct it the other way around: // 1st: open the target file // 2nd: init age.Encrypt - // 3rd: do mysql dump (which feeds the Writer) + // 3rd: do the DB dump (which feeds the Writer) wc, err := transferSession.EncryptToFile("dump.sql.enc") + if err != nil { + pterm.Fatal.Printfln("could not open encrypted dump file: %s", err) + } fileSet := &dto.FileSet{ Name: "dbDump", - Type: dto.TYPE_MYSQLDUMP, - MysqlDump: &dto.FileSetMysqlDump{ - FileName: "dump.sql.enc", - }, + Type: fileSetTypeFor(dbCredentials.Driver), } // 2b) the actual DB dump. also finishes writing. - db, err := mysql.CreateDump(dbCredentials, wc, whereClauseForTables) + var db *sql.DB + switch fileSet.Type { + case dto.TYPE_POSTGRESDUMP: + fileSet.PostgresDump = &dto.FileSetPostgresDump{FileName: "dump.sql.enc"} + db, err = postgres.CreateDump(dbCredentials, wc, whereClauseForTables) + default: + fileSet.MysqlDump = &dto.FileSetMysqlDump{FileName: "dump.sql.enc"} + db, err = mysql.CreateDump(dbCredentials, wc, whereClauseForTables) + } if err != nil { pterm.Fatal.Printfln("could not create SQL dump: %s", err) } - fileSet.MysqlDump.SizeBytes = wc.Size() + + sizeBytes := wc.Size() + if fileSet.PostgresDump != nil { + fileSet.PostgresDump.SizeBytes = sizeBytes + } else { + fileSet.MysqlDump.SizeBytes = sizeBytes + } transferSession.Meta.FileSets = append(transferSession.Meta.FileSets, fileSet) err = transferSession.UpdateMetadata() if err != nil { diff --git a/pkg/common/commonServe/databaseDump_test.go b/pkg/common/commonServe/databaseDump_test.go new file mode 100644 index 0000000..549c9a3 --- /dev/null +++ b/pkg/common/commonServe/databaseDump_test.go @@ -0,0 +1,29 @@ +package commonServe + +import ( + "testing" + + "github.com/sandstorm/synco/v2/pkg/common" + "github.com/sandstorm/synco/v2/pkg/common/dto" +) + +func TestFileSetTypeForDriver(t *testing.T) { + tests := []struct { + name string + driver common.DbDriver + want dto.FileSetType + }{ + {name: "postgres", driver: common.DbDriverPostgres, want: dto.TYPE_POSTGRESDUMP}, + {name: "mysql", driver: common.DbDriverMysql, want: dto.TYPE_MYSQLDUMP}, + {name: "unset driver defaults to mysql", driver: "", want: dto.TYPE_MYSQLDUMP}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := fileSetTypeFor(tt.driver) + if got != tt.want { + t.Errorf("fileSetTypeFor(%q) = %q, want %q", tt.driver, got, tt.want) + } + }) + } +} diff --git a/pkg/common/types.go b/pkg/common/types.go index 5f276ee..b026149 100644 --- a/pkg/common/types.go +++ b/pkg/common/types.go @@ -16,7 +16,15 @@ type ServeFramework interface { Serve(metadata *serve.TransferSession) } +type DbDriver string + +const ( + DbDriverMysql DbDriver = "mysql" + DbDriverPostgres DbDriver = "postgres" +) + type DbCredentials struct { + Driver DbDriver Host string Port int User string diff --git a/pkg/frameworks/flowServe/flowServe.go b/pkg/frameworks/flowServe/flowServe.go index 3850211..1b1fcad 100644 --- a/pkg/frameworks/flowServe/flowServe.go +++ b/pkg/frameworks/flowServe/flowServe.go @@ -121,11 +121,24 @@ type flowPersistenceBackendOptions struct { } func (fp *flowPersistenceBackendOptions) ToDbCredentials() *common.DbCredentials { - port := 3306 + driver := common.DbDriverMysql + defaultPort := 3306 + switch fp.Driver { + case "pdo_pgsql": + driver = common.DbDriverPostgres + defaultPort = 5432 + case "pdo_mysql", "": + // already defaulted above + default: + pterm.Warning.Printfln("unrecognized Neos.Flow.persistence.backendOptions.driver '%s', falling back to mysql", fp.Driver) + } + + port := defaultPort if len(fp.Port) != 0 { port, _ = strconv.Atoi(fp.Port) } return &common.DbCredentials{ + Driver: driver, Host: fp.Host, Port: port, User: fp.User, diff --git a/pkg/frameworks/flowServe/flowServer_test.go b/pkg/frameworks/flowServe/flowServer_test.go index b805396..037dafc 100644 --- a/pkg/frameworks/flowServe/flowServer_test.go +++ b/pkg/frameworks/flowServe/flowServer_test.go @@ -1,6 +1,10 @@ package flowServe -import "testing" +import ( + "testing" + + "github.com/sandstorm/synco/v2/pkg/common" +) type generateS3ResourcesPathTest struct { filename string @@ -54,6 +58,34 @@ var generateS3ResourcesPathTests = []generateS3ResourcesPathTest{ }, } +func TestToDbCredentialsDetectsDriverAndPortDefault(t *testing.T) { + tests := []struct { + name string + driver string + port string + wantDriver common.DbDriver + wantPort int + }{ + {name: "pdo_pgsql defaults to port 5432", driver: "pdo_pgsql", port: "", wantDriver: common.DbDriverPostgres, wantPort: 5432}, + {name: "pdo_mysql defaults to port 3306", driver: "pdo_mysql", port: "", wantDriver: common.DbDriverMysql, wantPort: 3306}, + {name: "unrecognized driver falls back to mysql", driver: "pdo_sqlite", port: "", wantDriver: common.DbDriverMysql, wantPort: 3306}, + {name: "explicit port overrides the postgres default", driver: "pdo_pgsql", port: "6543", wantDriver: common.DbDriverPostgres, wantPort: 6543}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + opts := flowPersistenceBackendOptions{Driver: tt.driver, Port: tt.port} + creds := opts.ToDbCredentials() + if creds.Driver != tt.wantDriver { + t.Errorf("Driver = %q, want %q", creds.Driver, tt.wantDriver) + } + if creds.Port != tt.wantPort { + t.Errorf("Port = %d, want %d", creds.Port, tt.wantPort) + } + }) + } +} + func TestGenerateS3ResourcesPaths(t *testing.T) { for _, tt := range generateS3ResourcesPathTests { path := generateS3ResourcePublicPath(&tt.targetConfiguration, tt.resourceSha1, tt.filename) diff --git a/pkg/frameworks/laravelServe/laravelServe.go b/pkg/frameworks/laravelServe/laravelServe.go index bd457d2..2d1c7be 100644 --- a/pkg/frameworks/laravelServe/laravelServe.go +++ b/pkg/frameworks/laravelServe/laravelServe.go @@ -58,11 +58,24 @@ func (ldo *laravelDatabaseOptions) ToDbCredentials() *common.DbCredentials { pterm.Warning.Printfln("Could not extract DB connection, WILL NOT INCLUDE DB DUMP.") return nil } - port := 3306 + driver := common.DbDriverMysql + defaultPort := 3306 + switch connection.Driver { + case "pgsql": + driver = common.DbDriverPostgres + defaultPort = 5432 + case "mysql", "": + // already defaulted above + default: + pterm.Warning.Printfln("unrecognized database driver '%s', falling back to mysql", connection.Driver) + } + + port := defaultPort if len(connection.Port) != 0 { port, _ = strconv.Atoi(connection.Port) } return &common.DbCredentials{ + Driver: driver, Host: connection.Host, Port: port, User: connection.Username, diff --git a/pkg/frameworks/laravelServe/laravelServe_test.go b/pkg/frameworks/laravelServe/laravelServe_test.go new file mode 100644 index 0000000..8f7090d --- /dev/null +++ b/pkg/frameworks/laravelServe/laravelServe_test.go @@ -0,0 +1,40 @@ +package laravelServe + +import ( + "testing" + + "github.com/sandstorm/synco/v2/pkg/common" +) + +func TestToDbCredentialsDetectsDriverAndPortDefault(t *testing.T) { + tests := []struct { + name string + driver string + port string + wantDriver common.DbDriver + wantPort int + }{ + {name: "pgsql defaults to port 5432", driver: "pgsql", port: "", wantDriver: common.DbDriverPostgres, wantPort: 5432}, + {name: "mysql defaults to port 3306", driver: "mysql", port: "", wantDriver: common.DbDriverMysql, wantPort: 3306}, + {name: "unrecognized driver falls back to mysql", driver: "sqlite", port: "", wantDriver: common.DbDriverMysql, wantPort: 3306}, + {name: "explicit port overrides the postgres default", driver: "pgsql", port: "6543", wantDriver: common.DbDriverPostgres, wantPort: 6543}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + opts := laravelDatabaseOptions{ + Default: "default", + Connections: map[string]laravelDatabaseConnectionOptions{ + "default": {Driver: tt.driver, Port: tt.port}, + }, + } + creds := opts.ToDbCredentials() + if creds.Driver != tt.wantDriver { + t.Errorf("Driver = %q, want %q", creds.Driver, tt.wantDriver) + } + if creds.Port != tt.wantPort { + t.Errorf("Port = %d, want %d", creds.Port, tt.wantPort) + } + }) + } +} diff --git a/pkg/receive/cmd/receive-cmd.go b/pkg/receive/cmd/receive-cmd.go index a66bc50..7d4ab0b 100644 --- a/pkg/receive/cmd/receive-cmd.go +++ b/pkg/receive/cmd/receive-cmd.go @@ -76,6 +76,8 @@ var ReceiveCmd = &cobra.Command{ switch fileSet.Type { case dto.TYPE_MYSQLDUMP: err = downloadMysqldump(receiveSession, fileSet) + case dto.TYPE_POSTGRESDUMP: + err = downloadPostgresdump(receiveSession, fileSet) case dto.TYPE_PUBLICFILES: err = downloadPublicFiles(receiveSession, fileSet) case dto.TYPE_PRIVATE_ENCRYPTED_FILES: @@ -218,6 +220,10 @@ func downloadMysqldump(receiveSession *receive.ReceiveSession, fileSet *dto.File return receiveSession.DumpAndDecryptFileWithProgressBar(fileSet.MysqlDump.FileName, fileSet.Name+".sql") } +func downloadPostgresdump(receiveSession *receive.ReceiveSession, fileSet *dto.FileSet) error { + return receiveSession.DumpAndDecryptFileWithProgressBar(fileSet.PostgresDump.FileName, fileSet.Name+".sql") +} + func downloadPublicFiles(receiveSession *receive.ReceiveSession, fileSet *dto.FileSet) error { indexFileName := fileSet.Name + ".index.json" err := receiveSession.DumpAndDecryptFileWithProgressBar(fileSet.PublicFiles.IndexFileName, indexFileName) diff --git a/pkg/util/postgres/go_pgdump/catalog.go b/pkg/util/postgres/go_pgdump/catalog.go new file mode 100644 index 0000000..43e788e --- /dev/null +++ b/pkg/util/postgres/go_pgdump/catalog.go @@ -0,0 +1,298 @@ +package go_pgdump + +import ( + "database/sql" + "fmt" + "sort" + "strings" +) + +// queryer is satisfied by both *sql.DB and *sql.Tx, letting catalog helpers +// run either against a live connection or an in-flight read-only transaction. +type queryer interface { + Query(query string, args ...interface{}) (*sql.Rows, error) + QueryRow(query string, args ...interface{}) *sql.Row +} + +const getTablesQuery = ` +SELECT c.relname +FROM pg_catalog.pg_class c +JOIN pg_catalog.pg_namespace n ON n.oid = c.relnamespace +WHERE n.nspname = 'public' + AND c.relkind = 'r' + AND NOT c.relispartition +ORDER BY c.relname; +` + +// getTables lists ordinary (base) tables in the public schema, excluding +// views, materialized views, partitioned parents, foreign tables and +// partition children. +func getTables(q queryer) ([]string, error) { + rows, err := q.Query(getTablesQuery) + if err != nil { + return nil, err + } + defer rows.Close() + + var tables []string + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return nil, err + } + tables = append(tables, name) + } + return tables, rows.Err() +} + +const getColumnsQuery = ` +SELECT a.attname, a.attnum, + pg_catalog.format_type(a.atttypid, a.atttypmod) AS data_type, + a.attnotnull, + pg_catalog.pg_get_expr(ad.adbin, ad.adrelid) AS default_expr, + a.attidentity, a.attgenerated +FROM pg_catalog.pg_attribute a +LEFT JOIN pg_catalog.pg_attrdef ad ON ad.adrelid = a.attrelid AND ad.adnum = a.attnum +WHERE a.attrelid = $1::regclass AND a.attnum > 0 AND NOT a.attisdropped +ORDER BY a.attnum; +` + +type column struct { + name string + dataType string + notNull bool + defaultExpr sql.NullString + identityKind string // "" | "a" | "d" + generated string // "" | "s" +} + +// columnDef renders this column's fragment of a CREATE TABLE statement +// (everything after the column name that isn't a table-level constraint). +func (c column) def() string { + part := quoteIdent(c.name) + " " + c.dataType + + switch { + case c.identityKind == "a": + part += " GENERATED ALWAYS AS IDENTITY" + case c.identityKind == "d": + part += " GENERATED BY DEFAULT AS IDENTITY" + case c.generated == "s": + part += fmt.Sprintf(" GENERATED ALWAYS AS (%s) STORED", c.defaultExpr.String) + case c.defaultExpr.Valid: + part += " DEFAULT " + c.defaultExpr.String + } + + if c.notNull { + part += " NOT NULL" + } + return part +} + +// isStoredGenerated reports whether this column is a `GENERATED ALWAYS AS +// (...) STORED` column, which Postgres computes itself and which therefore +// must be excluded from COPY's column list and data SELECT. +func (c column) isStoredGenerated() bool { + return c.generated == "s" +} + +func getColumns(q queryer, qualifiedTable string) ([]column, error) { + rows, err := q.Query(getColumnsQuery, qualifiedTable) + if err != nil { + return nil, err + } + defer rows.Close() + + var columns []column + for rows.Next() { + var c column + var attnum int + if err := rows.Scan(&c.name, &attnum, &c.dataType, &c.notNull, &c.defaultExpr, &c.identityKind, &c.generated); err != nil { + return nil, err + } + columns = append(columns, c) + } + return columns, rows.Err() +} + +// createTableDDL builds a `CREATE TABLE` statement for the given table +// (schema-qualified as "public"."") and returns the ordered list of +// column names that should be included in COPY data (excluding any stored +// generated columns, which Postgres computes itself). +func createTableDDL(q queryer, table string) (ddl string, copyColumns []string, err error) { + qualifiedTable := `"public".` + quoteIdent(table) + + columns, err := getColumns(q, qualifiedTable) + if err != nil { + return "", nil, err + } + + defs := make([]string, 0, len(columns)) + for _, c := range columns { + defs = append(defs, " "+c.def()) + if !c.isStoredGenerated() { + copyColumns = append(copyColumns, c.name) + } + } + + ddl = fmt.Sprintf("CREATE TABLE %s (\n%s\n);\n", qualifiedTable, strings.Join(defs, ",\n")) + return ddl, copyColumns, nil +} + +const getConstraintsQuery = ` +SELECT con.conname, con.contype, pg_catalog.pg_get_constraintdef(con.oid, true) +FROM pg_catalog.pg_constraint con +WHERE con.conrelid = $1::regclass AND con.contype IN ('p','u','c','f') +ORDER BY con.contype, con.conname; +` + +// constraintTypeRank orders constraint emission as PK -> UNIQUE -> CHECK -> +// FK. FKs must come last so all referenced tables/rows already exist by the +// time the constraint is added; PK/UNIQUE come first since they back the +// indexes FKs need. +var constraintTypeRank = map[string]int{"p": 0, "u": 1, "c": 2, "f": 3} + +type constraint struct { + name string + conType string // "p" | "u" | "c" | "f" + definition string +} + +func (c constraint) alterStatement(qualifiedTable string) string { + return fmt.Sprintf("ALTER TABLE %s ADD CONSTRAINT %s %s;", qualifiedTable, quoteIdent(c.name), c.definition) +} + +func getConstraints(q queryer, qualifiedTable string) ([]constraint, error) { + rows, err := q.Query(getConstraintsQuery, qualifiedTable) + if err != nil { + return nil, err + } + defer rows.Close() + + var constraints []constraint + for rows.Next() { + var c constraint + if err := rows.Scan(&c.name, &c.conType, &c.definition); err != nil { + return nil, err + } + constraints = append(constraints, c) + } + return constraints, rows.Err() +} + +// orderConstraints returns constraints sorted into PK -> UNIQUE -> CHECK -> +// FK phase order (stable within each type), regardless of the order they +// were retrieved in. +func orderConstraints(constraints []constraint) []constraint { + ordered := make([]constraint, len(constraints)) + copy(ordered, constraints) + sort.SliceStable(ordered, func(i, j int) bool { + return constraintTypeRank[ordered[i].conType] < constraintTypeRank[ordered[j].conType] + }) + return ordered +} + +const getIndexesQuery = ` +SELECT pg_catalog.pg_get_indexdef(i.indexrelid) +FROM pg_catalog.pg_index i +WHERE i.indrelid = $1::regclass AND NOT i.indisprimary + AND NOT EXISTS (SELECT 1 FROM pg_catalog.pg_constraint c WHERE c.conindid = i.indexrelid) +ORDER BY i.indexrelid::regclass::text; +` + +const getOwnedSequencesQuery = ` +SELECT s.relname AS seq_name, a.attname AS col_name +FROM pg_catalog.pg_depend d +JOIN pg_catalog.pg_class s ON s.oid = d.objid AND s.relkind = 'S' +JOIN pg_catalog.pg_class t ON t.oid = d.refobjid +JOIN pg_catalog.pg_attribute a ON a.attrelid = t.oid AND a.attnum = d.refobjsubid +WHERE d.classid = 'pg_class'::regclass + AND d.refclassid = 'pg_class'::regclass + AND d.deptype IN ('a','i') + AND t.oid = $1::regclass +ORDER BY s.relname; +` + +type ownedSequence struct { + seqName string + colName string +} + +func getOwnedSequences(q queryer, qualifiedTable string) ([]ownedSequence, error) { + rows, err := q.Query(getOwnedSequencesQuery, qualifiedTable) + if err != nil { + return nil, err + } + defer rows.Close() + + var seqs []ownedSequence + for rows.Next() { + var s ownedSequence + if err := rows.Scan(&s.seqName, &s.colName); err != nil { + return nil, err + } + seqs = append(seqs, s) + } + return seqs, rows.Err() +} + +// readSequenceState reads a sequence's own last_value/is_called, mirroring +// what pg_dump does - this restores the exact source cursor position, which +// matters when whereClauseForTables filtered out the highest-id rows (a +// MAX(col) recomputation from dumped data would under-restore in that case). +func readSequenceState(q queryer, qualifiedSeqName string) (lastValue int64, isCalled bool, err error) { + row := q.QueryRow(fmt.Sprintf("SELECT last_value, is_called FROM %s", qualifiedSeqName)) + err = row.Scan(&lastValue, &isCalled) + return lastValue, isCalled, err +} + +// quoteLiteral escapes a value for use inside a single-quoted SQL string +// literal (as opposed to quoteIdent, which escapes identifiers). +func quoteLiteral(s string) string { + return strings.ReplaceAll(s, "'", "''") +} + +func setvalStatement(qualifiedSeqName string, lastValue int64, isCalled bool) string { + return fmt.Sprintf("SELECT pg_catalog.setval('%s', %d, %t);", quoteLiteral(qualifiedSeqName), lastValue, isCalled) +} + +// sequenceRestoreStatements finds every sequence owned by a column of the +// given table (covering both classic SERIAL and GENERATED ... AS IDENTITY) +// and returns the `setval` statements needed to restore their exact live +// state after data load. +func sequenceRestoreStatements(q queryer, qualifiedTable string) ([]string, error) { + seqs, err := getOwnedSequences(q, qualifiedTable) + if err != nil { + return nil, err + } + + statements := make([]string, 0, len(seqs)) + for _, s := range seqs { + qualifiedSeqName := `"public".` + quoteIdent(s.seqName) + lastValue, isCalled, err := readSequenceState(q, qualifiedSeqName) + if err != nil { + return nil, err + } + statements = append(statements, setvalStatement(qualifiedSeqName, lastValue, isCalled)) + } + return statements, nil +} + +// getIndexes lists secondary (non-PK, non-constraint-backed) indexes on the +// given table as complete `CREATE INDEX ...;` statements. +func getIndexes(q queryer, qualifiedTable string) ([]string, error) { + rows, err := q.Query(getIndexesQuery, qualifiedTable) + if err != nil { + return nil, err + } + defer rows.Close() + + var defs []string + for rows.Next() { + var def string + if err := rows.Scan(&def); err != nil { + return nil, err + } + defs = append(defs, def+";") + } + return defs, rows.Err() +} diff --git a/pkg/util/postgres/go_pgdump/catalog_test.go b/pkg/util/postgres/go_pgdump/catalog_test.go new file mode 100644 index 0000000..da34a15 --- /dev/null +++ b/pkg/util/postgres/go_pgdump/catalog_test.go @@ -0,0 +1,216 @@ +package go_pgdump + +import ( + "testing" + + sqlmock "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" +) + +func TestGetTablesOnlyReturnsBaseTablesInPublicSchema(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + rows := sqlmock.NewRows([]string{"relname"}). + AddRow("t_child"). + AddRow("t_parent") + + // the query itself must filter to ordinary tables (relkind = 'r') in the + // public schema, excluding views/matviews/partition children - a wrongly + // shaped query (e.g. missing the relkind filter) will not match this + // expectation and the test will fail. + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_class c.*relkind = 'r'.*NOT c\.relispartition`). + WillReturnRows(rows) + + tables, err := getTables(db) + assert.NoError(t, err) + assert.Equal(t, []string{"t_child", "t_parent"}, tables) + + assert.NoError(t, mock.ExpectationsWereMet()) +} + +var columnQueryCols = []string{"attname", "attnum", "data_type", "attnotnull", "default_expr", "attidentity", "attgenerated"} + +func expectColumnQuery(mock sqlmock.Sqlmock, rows *sqlmock.Rows) { + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_attribute a.*attisdropped`).WillReturnRows(rows) +} + +func TestCreateTableDDLPlainNotNullColumn(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + rows := sqlmock.NewRows(columnQueryCols). + AddRow("name", 1, "character varying(255)", true, nil, "", "") + expectColumnQuery(mock, rows) + + ddl, copyColumns, err := createTableDDL(db, "t") + assert.NoError(t, err) + assert.Contains(t, ddl, `"name" character varying(255) NOT NULL`) + assert.Equal(t, []string{"name"}, copyColumns) + + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestSetvalStatementEscapesSingleQuotesInSequenceName(t *testing.T) { + // a sequence name can legally contain a single quote (e.g. a table named + // `my'thing` quoted as "my'thing"); since setval's first argument is a + // single-quoted string literal (not an identifier), that quote must be + // doubled or it terminates the literal early and corrupts the statement. + got := setvalStatement(`"public"."my'seq"`, 1, true) + want := `SELECT pg_catalog.setval('"public"."my''seq"', 1, true);` + if got != want { + t.Errorf("setvalStatement(...) = %q, want %q", got, want) + } +} + +func TestSequenceRestoreStatementsUseLiveSequenceStateVerbatim(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + ownedSeqRows := sqlmock.NewRows([]string{"seq_name", "col_name"}). + AddRow("t_id_seq", "id") + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_depend d.*deptype IN`).WillReturnRows(ownedSeqRows) + + // deliberately NOT the MAX(id)+... of any data - proves we read the + // sequence's own last_value/is_called rather than recomputing from rows, + // which matters when whereClauseForTables filters out the highest ids. + seqStateRows := sqlmock.NewRows([]string{"last_value", "is_called"}). + AddRow(42, true) + mock.ExpectQuery(`^SELECT last_value, is_called FROM "public"\."t_id_seq"$`).WillReturnRows(seqStateRows) + + statements, err := sequenceRestoreStatements(db, `"public"."t"`) + assert.NoError(t, err) + assert.Equal(t, []string{ + `SELECT pg_catalog.setval('"public"."t_id_seq"', 42, true);`, + }, statements) + + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestGetIndexesExcludesConstraintBackedIndexes(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + // Only the plain secondary index should be returned - the query itself + // excludes the PK index (NOT i.indisprimary) and unique-constraint-backed + // indexes (NOT EXISTS ... conindid); a wrongly shaped query that forgot + // either filter would not match this expectation. + rows := sqlmock.NewRows([]string{"indexdef"}). + AddRow(`CREATE INDEX t_email_idx ON public.t USING btree (email)`) + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_index i.*NOT i\.indisprimary.*conindid`). + WillReturnRows(rows) + + indexDefs, err := getIndexes(db, `"public"."t"`) + assert.NoError(t, err) + assert.Equal(t, []string{`CREATE INDEX t_email_idx ON public.t USING btree (email);`}, indexDefs) + + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestConstraintsAreEmittedInPkUniqueCheckFkOrder(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + // pg_constraint's natural query order (contype, conname) is alphabetical + // (c, f, p, u) - deliberately return them scrambled/alphabetical here to + // prove the emission order is enforced by our code, not by accident of + // the query's own ORDER BY. + rows := sqlmock.NewRows([]string{"conname", "contype", "definition"}). + AddRow("t_status_check", "c", "CHECK ((status = ANY (ARRAY['a','b'])))"). + AddRow("t_parent_id_fkey", "f", `FOREIGN KEY (parent_id) REFERENCES "public".t_parent(id)`). + AddRow("t_pkey", "p", "PRIMARY KEY (id)"). + AddRow("t_email_key", "u", "UNIQUE (email)") + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_constraint con.*contype IN`).WillReturnRows(rows) + + constraints, err := getConstraints(db, `"public"."t"`) + assert.NoError(t, err) + + ordered := orderConstraints(constraints) + var statements []string + for _, c := range ordered { + statements = append(statements, c.alterStatement(`"public"."t"`)) + } + + assert.Equal(t, []string{ + `ALTER TABLE "public"."t" ADD CONSTRAINT "t_pkey" PRIMARY KEY (id);`, + `ALTER TABLE "public"."t" ADD CONSTRAINT "t_email_key" UNIQUE (email);`, + `ALTER TABLE "public"."t" ADD CONSTRAINT "t_status_check" CHECK ((status = ANY (ARRAY['a','b'])));`, + `ALTER TABLE "public"."t" ADD CONSTRAINT "t_parent_id_fkey" FOREIGN KEY (parent_id) REFERENCES "public".t_parent(id);`, + }, statements) + + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestCreateTableDDLIdentityColumns(t *testing.T) { + tests := []struct { + name string + identityKind string + defaultExpr interface{} + want string + }{ + {name: "GENERATED ALWAYS AS IDENTITY, no default emitted", identityKind: "a", defaultExpr: nil, want: `"id" integer GENERATED ALWAYS AS IDENTITY`}, + {name: "GENERATED BY DEFAULT AS IDENTITY, default suppressed", identityKind: "d", defaultExpr: "nextval('t_id_seq'::regclass)", want: `"id" integer GENERATED BY DEFAULT AS IDENTITY`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + rows := sqlmock.NewRows(columnQueryCols). + AddRow("id", 1, "integer", true, tt.defaultExpr, tt.identityKind, "") + expectColumnQuery(mock, rows) + + ddl, copyColumns, err := createTableDDL(db, "t") + assert.NoError(t, err) + assert.Contains(t, ddl, tt.want) + assert.NotContains(t, ddl, "DEFAULT nextval", "identity columns must not also emit a DEFAULT clause") + assert.Equal(t, []string{"id"}, copyColumns) + + assert.NoError(t, mock.ExpectationsWereMet()) + }) + } +} + +func TestCreateTableDDLStoredGeneratedColumnExcludedFromCopy(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + rows := sqlmock.NewRows(columnQueryCols). + AddRow("first_name", 1, "text", false, nil, "", ""). + AddRow("last_name", 2, "text", false, nil, "", ""). + AddRow("full_name", 3, "text", false, "first_name || ' ' || last_name", "", "s") + expectColumnQuery(mock, rows) + + ddl, copyColumns, err := createTableDDL(db, "t") + assert.NoError(t, err) + assert.Contains(t, ddl, `"full_name" text GENERATED ALWAYS AS (first_name || ' ' || last_name) STORED`) + assert.Equal(t, []string{"first_name", "last_name"}, copyColumns, + "stored generated columns must be excluded from the COPY column list - Postgres computes them itself") + + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestCreateTableDDLColumnWithDefault(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + rows := sqlmock.NewRows(columnQueryCols). + AddRow("balance", 1, "numeric(10,2)", false, "0", "", "") + expectColumnQuery(mock, rows) + + ddl, copyColumns, err := createTableDDL(db, "t") + assert.NoError(t, err) + assert.Contains(t, ddl, `"balance" numeric(10,2) DEFAULT 0`) + assert.Equal(t, []string{"balance"}, copyColumns) + + assert.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/pkg/util/postgres/go_pgdump/copy.go b/pkg/util/postgres/go_pgdump/copy.go new file mode 100644 index 0000000..7b39dab --- /dev/null +++ b/pkg/util/postgres/go_pgdump/copy.go @@ -0,0 +1,88 @@ +package go_pgdump + +import ( + "database/sql" + "fmt" + "io" + "strings" +) + +// quoteIdent quotes a Postgres identifier (table/column/sequence name) for +// safe use in generated DDL, doubling any embedded double quotes. +func quoteIdent(name string) string { + return `"` + strings.ReplaceAll(name, `"`, `""`) + `"` +} + +// copyEscape renders a single value in COPY text format: SQL NULL becomes the +// `\N` marker, everything else has backslash, tab, newline and carriage +// return escaped so the value can never be confused with the NULL marker or +// a row/column delimiter. +func copyEscape(v sql.NullString) string { + if !v.Valid { + return `\N` + } + + s := v.String + s = strings.ReplaceAll(s, `\`, `\\`) + s = strings.ReplaceAll(s, "\t", `\t`) + s = strings.ReplaceAll(s, "\n", `\n`) + s = strings.ReplaceAll(s, "\r", `\r`) + return s +} + +// streamCopy writes a `COPY ... FROM stdin; ... \.` text-format block for +// the given table/columns to w. Every column is cast to ::text server-side +// so that bytea/json/array/uuid/timestamptz/numeric/boolean values all +// round-trip using only generic COPY-text escaping - no per-Postgres-type Go +// formatter is needed. An empty whereClause defaults to "TRUE" (dump +// everything), mirroring the MySQL dumper's WhereClauseForTables handling. +func streamCopy(w io.Writer, q queryer, table string, copyColumns []string, whereClause string) error { + if whereClause == "" { + whereClause = "TRUE" + } + + qualifiedTable := `"public".` + quoteIdent(table) + + quotedCols := make([]string, len(copyColumns)) + castCols := make([]string, len(copyColumns)) + for i, col := range copyColumns { + quotedCols[i] = quoteIdent(col) + castCols[i] = quoteIdent(col) + "::text" + } + + selectSQL := fmt.Sprintf("SELECT %s FROM %s WHERE %s", strings.Join(castCols, ", "), qualifiedTable, whereClause) + rows, err := q.Query(selectSQL) + if err != nil { + return err + } + defer rows.Close() + + if _, err := fmt.Fprintf(w, "COPY %s (%s) FROM stdin;\n", qualifiedTable, strings.Join(quotedCols, ",")); err != nil { + return err + } + + values := make([]sql.NullString, len(copyColumns)) + dest := make([]interface{}, len(copyColumns)) + for i := range values { + dest[i] = &values[i] + } + + for rows.Next() { + if err := rows.Scan(dest...); err != nil { + return err + } + escaped := make([]string, len(values)) + for i, v := range values { + escaped[i] = copyEscape(v) + } + if _, err := fmt.Fprintf(w, "%s\n", strings.Join(escaped, "\t")); err != nil { + return err + } + } + if err := rows.Err(); err != nil { + return err + } + + _, err = fmt.Fprint(w, "\\.\n") + return err +} diff --git a/pkg/util/postgres/go_pgdump/copy_test.go b/pkg/util/postgres/go_pgdump/copy_test.go new file mode 100644 index 0000000..1c1dc03 --- /dev/null +++ b/pkg/util/postgres/go_pgdump/copy_test.go @@ -0,0 +1,92 @@ +package go_pgdump + +import ( + "bytes" + "database/sql" + "testing" + + sqlmock "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" +) + +func TestQuoteIdent(t *testing.T) { + tests := []struct { + name string + in string + want string + }{ + {name: "plain name", in: "users", want: `"users"`}, + {name: "name containing a double quote", in: `fo"o`, want: `"fo""o"`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := quoteIdent(tt.in) + if got != tt.want { + t.Errorf("quoteIdent(%q) = %q, want %q", tt.in, got, tt.want) + } + }) + } +} + +func TestCopyEscape(t *testing.T) { + tests := []struct { + name string + in sql.NullString + want string + }{ + {name: "SQL NULL", in: sql.NullString{Valid: false}, want: `\N`}, + {name: "empty string is not NULL", in: sql.NullString{String: "", Valid: true}, want: ""}, + {name: "embedded tab", in: sql.NullString{String: "a\tb", Valid: true}, want: `a\tb`}, + {name: "embedded newline", in: sql.NullString{String: "a\nb", Valid: true}, want: `a\nb`}, + {name: "embedded carriage return", in: sql.NullString{String: "a\rb", Valid: true}, want: `a\rb`}, + {name: "embedded backslash", in: sql.NullString{String: `a\b`, Valid: true}, want: `a\\b`}, + {name: "literal backslash-N is escaped, not mistaken for NULL", in: sql.NullString{String: `\N`, Valid: true}, want: `\\N`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := copyEscape(tt.in) + if got != tt.want { + t.Errorf("copyEscape(%+v) = %q, want %q", tt.in, got, tt.want) + } + }) + } +} + +func TestStreamCopyEmitsExactBlockAndCastsColumnsToText(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + rows := sqlmock.NewRows([]string{"a", "b"}). + AddRow("1", nil). + AddRow("2", "has\ta\ttab") + mock.ExpectQuery(`^SELECT "a"::text, "b"::text FROM "public"\."t" WHERE TRUE$`).WillReturnRows(rows) + + var buf bytes.Buffer + err = streamCopy(&buf, db, "t", []string{"a", "b"}, "") + assert.NoError(t, err) + + assert.Equal(t, "COPY \"public\".\"t\" (\"a\",\"b\") FROM stdin;\n"+ + "1\t\\N\n"+ + "2\thas\\ta\\ttab\n"+ + "\\.\n", buf.String()) + + assert.NoError(t, mock.ExpectationsWereMet()) +} + +func TestStreamCopyUsesCustomWhereClause(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + rows := sqlmock.NewRows([]string{"a"}) + mock.ExpectQuery(`^SELECT "a"::text FROM "public"\."t" WHERE FALSE$`).WillReturnRows(rows) + + var buf bytes.Buffer + err = streamCopy(&buf, db, "t", []string{"a"}, "FALSE") + assert.NoError(t, err) + + assert.NoError(t, mock.ExpectationsWereMet()) +} diff --git a/pkg/util/postgres/go_pgdump/pgdump.go b/pkg/util/postgres/go_pgdump/pgdump.go new file mode 100644 index 0000000..075570b --- /dev/null +++ b/pkg/util/postgres/go_pgdump/pgdump.go @@ -0,0 +1,205 @@ +package go_pgdump + +import ( + "context" + "database/sql" + "fmt" + "io" +) + +// headerSQL is written verbatim at the top of the dump so it replays +// cleanly via `psql -f`. +const headerSQL = `SET statement_timeout = 0; +SET client_encoding = 'UTF8'; +SET standard_conforming_strings = on; +SET check_function_bodies = false; +SET client_min_messages = warning; +SET row_security = off; +` + +// Data configures and drives a Postgres dump, mirroring the shape of +// pkg/util/mysql/go_mysqldump.Data. +type Data struct { + Out io.Writer + Connection *sql.DB + WhereClauseForTables map[string]string + IgnoreTables []string + + tx *sql.Tx +} + +// NewDumper registers a new dumper against an already-open connection. +func NewDumper(db *sql.DB, out io.Writer) *Data { + return &Data{Connection: db, Out: out} +} + +// Close closes the dumper's output (if closable) and its DB connection. +func (d *Data) Close() error { + defer func() { + d.Connection = nil + d.Out = nil + }() + if out, ok := d.Out.(io.Closer); ok { + _ = out.Close() + } + return d.Connection.Close() +} + +func (d *Data) isIgnoredTable(name string) bool { + for _, t := range d.IgnoreTables { + if t == name { + return true + } + } + return false +} + +func qualify(table string) string { + return `"public".` + quoteIdent(table) +} + +// qualifiedConstraint pairs a constraint with the (already schema-qualified) +// table it belongs to, so constraints from every table can be grouped and +// emitted together by type across the whole dump. +type qualifiedConstraint struct { + qualifiedTable string + c constraint +} + +// Dump writes a complete, replayable SQL dump of the database to d.Out. +// +// Phases, in order: (1) DROP + CREATE TABLE for every table (columns only - +// no constraints yet), (2) COPY data for every table, (3) PK, then UNIQUE +// constraints, then secondary indexes, then CHECK constraints, then FOREIGN +// KEY constraints last (deferred so referenced tables/rows already exist +// and no cross-table topological sort is needed), (4) sequence `setval` +// restoration. This mirrors real pg_dump's schema/data/post-data staging. +func (d *Data) Dump() error { + tx, err := d.Connection.BeginTx(context.Background(), &sql.TxOptions{ + Isolation: sql.LevelRepeatableRead, + ReadOnly: true, + }) + if err != nil { + return err + } + defer tx.Rollback() + d.tx = tx + + for _, stmt := range []string{ + "SET TIME ZONE 'UTC'", + "SET extra_float_digits = 3", + "SET bytea_output = 'hex'", + } { + if _, err := tx.Exec(stmt); err != nil { + return err + } + } + + if _, err := fmt.Fprint(d.Out, headerSQL); err != nil { + return err + } + + allTables, err := getTables(tx) + if err != nil { + return err + } + + tables := make([]string, 0, len(allTables)) + for _, t := range allTables { + if !d.isIgnoredTable(t) { + tables = append(tables, t) + } + } + + copyColumnsByTable := make(map[string][]string, len(tables)) + + // Phase 1: DROP + CREATE TABLE (columns only). + for _, t := range tables { + ddl, copyColumns, err := createTableDDL(tx, t) + if err != nil { + return err + } + copyColumnsByTable[t] = copyColumns + + if _, err := fmt.Fprintf(d.Out, "DROP TABLE IF EXISTS %s CASCADE;\n", qualify(t)); err != nil { + return err + } + if _, err := fmt.Fprint(d.Out, ddl); err != nil { + return err + } + } + + // Phase 2: data. + for _, t := range tables { + if err := streamCopy(d.Out, tx, t, copyColumnsByTable[t], d.WhereClauseForTables[t]); err != nil { + return err + } + } + + // Phase 3: constraints + indexes, gathered across all tables first so + // they can be emitted grouped by type (PK/UNIQUE/index/CHECK/FK) rather + // than per table. + var allConstraints []qualifiedConstraint + var allIndexes []string + for _, t := range tables { + qualifiedTable := qualify(t) + + cs, err := getConstraints(tx, qualifiedTable) + if err != nil { + return err + } + for _, c := range orderConstraints(cs) { + allConstraints = append(allConstraints, qualifiedConstraint{qualifiedTable, c}) + } + + idx, err := getIndexes(tx, qualifiedTable) + if err != nil { + return err + } + allIndexes = append(allIndexes, idx...) + } + + writeConstraintsOfRank := func(rank int) error { + for _, ac := range allConstraints { + if constraintTypeRank[ac.c.conType] == rank { + if _, err := fmt.Fprintln(d.Out, ac.c.alterStatement(ac.qualifiedTable)); err != nil { + return err + } + } + } + return nil + } + + if err := writeConstraintsOfRank(constraintTypeRank["p"]); err != nil { + return err + } + if err := writeConstraintsOfRank(constraintTypeRank["u"]); err != nil { + return err + } + for _, idx := range allIndexes { + if _, err := fmt.Fprintln(d.Out, idx); err != nil { + return err + } + } + if err := writeConstraintsOfRank(constraintTypeRank["c"]); err != nil { + return err + } + if err := writeConstraintsOfRank(constraintTypeRank["f"]); err != nil { + return err + } + + // Phase 4: sequence state restoration. + for _, t := range tables { + stmts, err := sequenceRestoreStatements(tx, qualify(t)) + if err != nil { + return err + } + for _, s := range stmts { + if _, err := fmt.Fprintln(d.Out, s); err != nil { + return err + } + } + } + + return nil +} diff --git a/pkg/util/postgres/go_pgdump/pgdump_test.go b/pkg/util/postgres/go_pgdump/pgdump_test.go new file mode 100644 index 0000000..cd63d7b --- /dev/null +++ b/pkg/util/postgres/go_pgdump/pgdump_test.go @@ -0,0 +1,124 @@ +package go_pgdump + +import ( + "bytes" + "strings" + "testing" + + sqlmock "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" +) + +// TestDumpOrdersOutputIntoSchemaThenDataThenConstraintsThenSequences is the +// package's own integration test: it stitches together every catalog/copy +// piece built in the earlier slices for a small two-table schema (t_child +// has an FK to t_parent) and asserts the complete output file ordering: +// both DROP+CREATE TABLEs, then both tables' COPY blocks, then +// PK -> UNIQUE -> secondary indexes -> CHECK -> FK (deferred to last so +// referenced tables/rows already exist), then sequence setval restoration. +func TestDumpOrdersOutputIntoSchemaThenDataThenConstraintsThenSequences(t *testing.T) { + db, mock, err := sqlmock.New() + assert.NoError(t, err, "an error was not expected when opening a stub database connection") + defer db.Close() + + mock.ExpectBegin() + mock.ExpectExec(`SET TIME ZONE 'UTC'`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec(`SET extra_float_digits = 3`).WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec(`SET bytea_output = 'hex'`).WillReturnResult(sqlmock.NewResult(0, 0)) + + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_class c.*relkind = 'r'`). + WillReturnRows(sqlmock.NewRows([]string{"relname"}). + AddRow("t_child"). + AddRow("t_parent")) + + // Phase 1: column introspection (order follows getTables: t_child, t_parent) + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_attribute a.*attisdropped`). + WithArgs(`"public"."t_child"`). + WillReturnRows(sqlmock.NewRows(columnQueryCols). + AddRow("id", 1, "integer", true, nil, "a", ""). + AddRow("parent_id", 2, "integer", true, nil, "", ""). + AddRow("status", 3, "text", false, nil, "", "")) + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_attribute a.*attisdropped`). + WithArgs(`"public"."t_parent"`). + WillReturnRows(sqlmock.NewRows(columnQueryCols). + AddRow("id", 1, "integer", true, nil, "a", ""). + AddRow("name", 2, "text", true, nil, "", "")) + + // Phase 2: data + mock.ExpectQuery(`^SELECT "id"::text, "parent_id"::text, "status"::text FROM "public"\."t_child" WHERE TRUE$`). + WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id", "status"}).AddRow("1", "1", "ok")) + mock.ExpectQuery(`^SELECT "id"::text, "name"::text FROM "public"\."t_parent" WHERE TRUE$`). + WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow("1", "alpha")) + + // Phase 3: constraints + indexes (t_child then t_parent) + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_constraint con.*contype IN`). + WithArgs(`"public"."t_child"`). + WillReturnRows(sqlmock.NewRows([]string{"conname", "contype", "definition"}). + AddRow("t_child_pkey", "p", "PRIMARY KEY (id)"). + AddRow("t_child_parent_id_fkey", "f", `FOREIGN KEY (parent_id) REFERENCES "public".t_parent(id)`). + AddRow("t_child_status_check", "c", "CHECK ((status IS NOT NULL))")) + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_index i.*NOT i\.indisprimary`). + WithArgs(`"public"."t_child"`). + WillReturnRows(sqlmock.NewRows([]string{"indexdef"})) + + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_constraint con.*contype IN`). + WithArgs(`"public"."t_parent"`). + WillReturnRows(sqlmock.NewRows([]string{"conname", "contype", "definition"}). + AddRow("t_parent_pkey", "p", "PRIMARY KEY (id)"). + AddRow("t_parent_name_key", "u", "UNIQUE (name)")) + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_index i.*NOT i\.indisprimary`). + WithArgs(`"public"."t_parent"`). + WillReturnRows(sqlmock.NewRows([]string{"indexdef"}). + AddRow("CREATE INDEX t_parent_name_idx ON public.t_parent USING btree (name)")) + + // Phase 4: sequence restoration (t_child then t_parent) + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_depend d.*deptype IN`). + WithArgs(`"public"."t_child"`). + WillReturnRows(sqlmock.NewRows([]string{"seq_name", "col_name"}).AddRow("t_child_id_seq", "id")) + mock.ExpectQuery(`^SELECT last_value, is_called FROM "public"\."t_child_id_seq"$`). + WillReturnRows(sqlmock.NewRows([]string{"last_value", "is_called"}).AddRow(5, true)) + + mock.ExpectQuery(`(?s)FROM pg_catalog\.pg_depend d.*deptype IN`). + WithArgs(`"public"."t_parent"`). + WillReturnRows(sqlmock.NewRows([]string{"seq_name", "col_name"}).AddRow("t_parent_id_seq", "id")) + mock.ExpectQuery(`^SELECT last_value, is_called FROM "public"\."t_parent_id_seq"$`). + WillReturnRows(sqlmock.NewRows([]string{"last_value", "is_called"}).AddRow(2, true)) + + mock.ExpectRollback() + + var buf bytes.Buffer + dumper := NewDumper(db, &buf) + assert.NoError(t, dumper.Dump()) + assert.NoError(t, mock.ExpectationsWereMet()) + + out := buf.String() + find := func(substr string) int { + i := strings.Index(out, substr) + if i == -1 { + t.Fatalf("expected output to contain %q, got:\n%s", substr, out) + } + return i + } + + dropChild := find(`DROP TABLE IF EXISTS "public"."t_child" CASCADE;`) + dropParent := find(`DROP TABLE IF EXISTS "public"."t_parent" CASCADE;`) + copyChild := find(`COPY "public"."t_child"`) + copyParent := find(`COPY "public"."t_parent"`) + pkChild := find(`ADD CONSTRAINT "t_child_pkey" PRIMARY KEY`) + pkParent := find(`ADD CONSTRAINT "t_parent_pkey" PRIMARY KEY`) + uniqueParent := find(`ADD CONSTRAINT "t_parent_name_key" UNIQUE`) + indexParent := find(`CREATE INDEX t_parent_name_idx`) + checkChild := find(`ADD CONSTRAINT "t_child_status_check" CHECK`) + fkChild := find(`ADD CONSTRAINT "t_child_parent_id_fkey" FOREIGN KEY`) + setvalChild := find(`setval('"public"."t_child_id_seq"'`) + setvalParent := find(`setval('"public"."t_parent_id_seq"'`) + + assert.True(t, dropChild < copyChild && dropParent < copyChild, "both DROP/CREATE TABLEs must precede any COPY") + assert.True(t, copyChild < pkChild && copyParent < pkChild, "COPY data must precede constraints") + assert.True(t, pkChild < uniqueParent, "PK constraints must precede UNIQUE constraints") + assert.True(t, uniqueParent < indexParent, "UNIQUE constraints must precede secondary indexes") + assert.True(t, indexParent < checkChild, "secondary indexes must precede CHECK constraints") + assert.True(t, checkChild < fkChild, "CHECK constraints must precede FOREIGN KEY constraints") + assert.True(t, fkChild < setvalChild && fkChild < setvalParent, "FOREIGN KEY constraints must precede sequence restoration") + _ = pkParent +} diff --git a/pkg/util/postgres/pgdump.go b/pkg/util/postgres/pgdump.go new file mode 100644 index 0000000..265fd7b --- /dev/null +++ b/pkg/util/postgres/pgdump.go @@ -0,0 +1,41 @@ +package postgres + +import ( + "database/sql" + "fmt" + "io" + + _ "github.com/jackc/pgx/v5/stdlib" + "github.com/sandstorm/synco/v2/pkg/common" + go_pgdump "github.com/sandstorm/synco/v2/pkg/util/postgres/go_pgdump" +) + +// CreateDump opens a connection to the given Postgres database and writes a +// complete, replayable SQL dump to writer. +// +// Unlike the MySQL dumper, there is no manual TLS-then-fallback retry here: +// `sslmode=prefer` negotiates TLS automatically and falls back to plaintext +// on its own if the server doesn't support it. +func CreateDump(dbCredentials *common.DbCredentials, writer io.WriteCloser, whereClauseForTables map[string]string) (*sql.DB, error) { + dsn := fmt.Sprintf( + "host=%s port=%d user=%s password=%s dbname=%s sslmode=prefer", + dbCredentials.Host, dbCredentials.Port, dbCredentials.User, dbCredentials.Password, dbCredentials.DbName, + ) + + db, err := sql.Open("pgx", dsn) + if err != nil { + return nil, fmt.Errorf("error opening database: %w", err) + } + + dumper := go_pgdump.NewDumper(db, writer) + dumper.WhereClauseForTables = whereClauseForTables + if err := dumper.Dump(); err != nil { + return nil, fmt.Errorf("error dumping database: %w", err) + } + + if err := writer.Close(); err != nil { + return nil, fmt.Errorf("error closing dumper: %w", err) + } + + return db, nil +} diff --git a/test_e2e/postgresframework_test.go b/test_e2e/postgresframework_test.go new file mode 100644 index 0000000..096a9eb --- /dev/null +++ b/test_e2e/postgresframework_test.go @@ -0,0 +1,97 @@ +package test_e2e + +import ( + "io/ioutil" + "strconv" + "strings" + "testing" + "time" + + "github.com/orlangure/gnomock" + "github.com/orlangure/gnomock/preset/postgres" +) +import "github.com/rogpeppe/go-internal/testscript" + +const pgQueries = ` + create table t_parent ( + id integer generated always as identity primary key, + name varchar(255) not null + ); + create table t_child ( + id integer generated always as identity primary key, + parent_id integer not null references t_parent(id), + value text + ); + insert into t_parent (name) values ('alpha'); + insert into t_parent (name) values ('beta'); + insert into t_child (parent_id, value) values (1, 'child-of-alpha'); + insert into t_child (parent_id, value) values (2, 'child-of-beta'); +` + +func startPostgresDb(t *testing.T) (string, string) { + t.Helper() + p := postgres.Preset( + postgres.WithVersion("16"), + postgres.WithUser("admin", "password"), + postgres.WithDatabase("dummy1"), + postgres.WithQueries(pgQueries), + ) + var container *gnomock.Container + var err error + if reuseDatabaseContainer { + container, err = gnomock.Start(p, gnomock.WithDebugMode(), gnomock.WithContainerReuse(), gnomock.WithContainerName("synco-test-postgres")) + } else { + container, err = gnomock.Start(p) + t.Cleanup(func() { + _ = gnomock.Stop(container) + }) + } + + if err != nil { + panic(err) + } + return container.Host, strconv.Itoa(container.DefaultPort()) +} + +func TestPostgresFrameworkExportsDatabase(t *testing.T) { + dbHost, dbPort := startPostgresDb(t) + testscript.Run(t, testscript.Params{ + Dir: "testdata/postgresframework", + Setup: func(env *testscript.Env) error { + env.Setenv("DB_USER", "admin") + env.Setenv("DB_PASSWORD", "password") + env.Setenv("DB_NAME", "dummy1") + env.Setenv("DB_HOST", dbHost) + env.Setenv("DB_PORT", dbPort) + + return nil + }, + TestWork: true, + Cmds: map[string]func(ts *testscript.TestScript, neg bool, args []string){ + "fileContentWithTimeout": func(ts *testscript.TestScript, neg bool, args []string) { + fileName := args[0] + expectedContent := args[1] + maxDuration, err := time.ParseDuration(args[2]) + if err != nil { + ts.Fatalf("Error parsing duration: %s", err) + } + + startTime := time.Now() + for { + file, err := ioutil.ReadFile(fileName) + if err == nil && strings.TrimSpace(string(file)) == strings.TrimSpace(expectedContent) { + ts.Logf("Successful file content comparison") + // no error and matching file content -> success! + return + } + + if time.Since(startTime) > maxDuration { + ts.Fatalf("Error maxDuration") + break + } + time.Sleep(200 * time.Millisecond) + } + }, + }, + }) +} diff --git a/test_e2e/testdata/postgresframework/flow_working.txtar b/test_e2e/testdata/postgresframework/flow_working.txtar new file mode 100644 index 0000000..cfefaf2 --- /dev/null +++ b/test_e2e/testdata/postgresframework/flow_working.txtar @@ -0,0 +1,71 @@ +# prepare +chmod 755 flow + +# server part +# --all makes resource extraction walk the filesystem target folder instead of +# querying the (here non-existent) neos_flow_resourcemanagement_persistentresource table. +exec synco serve --debug --all --id abcde --password super-secret-pass --listen :8883 & +fileContentWithTimeout $WORK/Web/_Resources/synco-abcde/state Ready 10s + +# client part +# the base URL of the production server is read from .synco.yml (see below) +exec synco receive --debug synco-abcde super-secret-pass --interactive=false + +# the DB dump must have been downloaded and decrypted, and must contain a +# reconstructed schema (not just data) plus the identity-column tables we seeded. +exists dump/dbDump.sql +grep 'CREATE TABLE "public"."t_parent"' dump/dbDump.sql +grep 'CREATE TABLE "public"."t_child"' dump/dbDump.sql +grep 'COPY "public"."t_parent"' dump/dbDump.sql +grep 'COPY "public"."t_child"' dump/dbDump.sql +grep 'ADD CONSTRAINT' dump/dbDump.sql +grep 'setval' dump/dbDump.sql + +-- .synco.yml -- +hosts: + - baseUrl: http://127.0.0.1:8883 + +-- flow -- +#!/usr/bin/env bash + +# This file is a fake ./flow CLI which can fake returning persistence options +# and resource configuration. + +if [[ "$@" == "configuration:show --type Settings --path Neos.Flow.persistence.backendOptions" ]]; then + + cat << EOF +Configuration "Settings: Neos.Flow.persistence.backendOptions": + +driver: pdo_pgsql +host: $DB_HOST +dbname: $DB_NAME +user: $DB_USER +password: $DB_PASSWORD +charset: utf8 +defaultTableOptions: + charset: utf8 +port: $DB_PORT +EOF + +elif [[ "$@" == "configuration:show --type Settings --path Neos.Flow.resource" ]]; then + + cat << EOF +Configuration "Settings: Neos.Flow.resource": + +collections: + persistent: + target: localWebDirectoryPersistentResourcesTarget +targets: + localWebDirectoryPersistentResourcesTarget: + target: Neos\Flow\ResourceManagement\Target\FileSystemSymlinkTarget + targetOptions: + path: ./Web/_Resources/Persistent/ + baseUri: _Resources/Persistent/ +EOF + +else + echo "Unsupported call " $@ + exit 1 +fi + +-- Web/_Resources/.keepme --