From eff2c14a9db8aa84d5be2dca4d65ec644cf08448 Mon Sep 17 00:00:00 2001 From: Matthew Meszaros Date: Sun, 4 Oct 2026 02:46:56 -0700 Subject: [PATCH] feat: make workspace import write only rows the destination workspace owns by checking every archive key and foreign key against each table's owner scope before a batch lands, scoping overwrite updates to the destination's own rows, warning in the archive check, and documenting the rule --- .../docs/guides/workspace-export-import.mdx | 6 +- internal/app/orgtransfer/import.go | 74 +++++- .../orgtransfer/import_tenant_live_test.go | 218 ++++++++++++++++++ internal/app/orgtransfer/jobs.go | 27 ++- internal/app/orgtransfer/spec.go | 16 ++ internal/repository/pg_orgtransfer.go | 142 +++++++++++- internal/repository/pg_orgtransfer_test.go | 17 ++ 7 files changed, 480 insertions(+), 20 deletions(-) create mode 100644 internal/app/orgtransfer/import_tenant_live_test.go create mode 100644 internal/repository/pg_orgtransfer_test.go diff --git a/docs/content/docs/guides/workspace-export-import.mdx b/docs/content/docs/guides/workspace-export-import.mdx index 9bd67cbc2..a3ae6fda9 100644 --- a/docs/content/docs/guides/workspace-export-import.mdx +++ b/docs/content/docs/guides/workspace-export-import.mdx @@ -73,7 +73,11 @@ Then pick the groups to apply, decide what happens to rows that already exist, a **Keep what is here** is the default: the archive only adds rows the destination does not already have. This is the right choice for an empty destination workspace, and the safe one for a workspace already in use. -**Replace with the archive** overwrites matching rows with the archive's versions. Use it when you are re-running an import to pick up changes made on the source since the last one. It cannot be undone. +**Replace with the archive** overwrites matching rows in the destination workspace with the archive's versions. Use it when you are re-running an import to pick up changes made on the source since the last one. It cannot be undone. + +### Records of another workspace + +An import only writes the destination workspace's own records, and everything it brings in points only at that workspace's data. Records keep their identifiers across a move, so an archive holding records that already belong to another workspace on the same instance (for example, one exported from a workspace that still lives there) is refused as a whole and nothing is written. The archive check reports this before you confirm. Import such an archive into the workspace it was exported from, or into a workspace on another instance. ### How people are matched diff --git a/internal/app/orgtransfer/import.go b/internal/app/orgtransfer/import.go index 191f6b3a1..d1688f934 100644 --- a/internal/app/orgtransfer/import.go +++ b/internal/app/orgtransfer/import.go @@ -30,6 +30,11 @@ const ( maxManifestBytes = 64 << 20 ) +// ErrForeignRows refuses an archive naming records another workspace on this +// instance owns: an import only writes and references the destination's own. +var ErrForeignRows = errors.New("this archive holds records that already belong to another workspace on this instance, so nothing was imported. " + + "Import it into the workspace it was exported from, or into a workspace on another instance") + // ImportFrom applies an archive to a destination workspace. // // The whole thing runs in one transaction. An import is a rare, deliberate, @@ -253,21 +258,23 @@ func (s *service) importTable( "%s: %d column(s) in the archive do not exist here and were ignored.", t.Name, unknown)) } - var pk []string - if ic.conflict == models.OrgImportConflictOverwrite { - if pk, err = s.repo.PrimaryKeyColumns(ctx, t.Name); err != nil { - return 0, nil, err - } - if len(pk) == 0 { - warnings = append(warnings, fmt.Sprintf( - "%s has no primary key, so existing rows there were kept rather than overwritten.", t.Name)) - } + pk, err := s.repo.PrimaryKeyColumns(ctx, t.Name) + if err != nil { + return 0, nil, err + } + if len(pk) == 0 && ic.conflict == models.OrgImportConflictOverwrite { + warnings = append(warnings, fmt.Sprintf( + "%s has no primary key, so existing rows there were kept rather than overwritten.", t.Name)) } refs, err := s.referencePlan(ctx, t, destCols, ic.selected) if err != nil { return 0, nil, err } + tenantRefs, err := s.tenantReferences(ctx, t.Name, insertCols) + if err != nil { + return 0, nil, err + } if len(refs.nullify) > 0 { warnings = append(warnings, fmt.Sprintf( "%s: %d reference(s) point at data this import does not include and were cleared.", @@ -289,7 +296,14 @@ func (s *service) importTable( if len(batch) == 0 { return nil } - n, err := s.repo.InsertBatch(ctx, tx, t.Name, insertCols, batch, ic.conflict, pk) + foreign, err := s.repo.CountForeignRows(ctx, tx, t.Name, t.OwnerScope(), pk, tenantRefs, orgID, batch) + if err != nil { + return err + } + if foreign > 0 { + return fmt.Errorf("%s: %w", t.Name, ErrForeignRows) + } + n, err := s.repo.InsertBatch(ctx, tx, t.Name, insertCols, batch, ic.conflict, pk, t.OwnerScope(), orgID) if err != nil { return err } @@ -522,6 +536,46 @@ func (s *service) referencePlan( return plan, nil } +// tenantReferences are the foreign keys a written row carries into data some +// organization owns, each with the fragment selecting the destination's rows. +func (s *service) tenantReferences(ctx context.Context, table string, insertCols []string) ([]repository.TenantReference, error) { + fks, err := s.repo.ForeignKeys(ctx, table) + if err != nil { + return nil, err + } + written := make(map[string]bool, len(insertCols)) + for _, c := range insertCols { + written[c] = true + } + var out []repository.TenantReference + for _, fk := range fks { + carried := true + for _, c := range fk.Columns { + carried = carried && written[c] + } + if !carried { + continue + } + var owner string + switch fk.RefTable { + case "users": + // People are matched to destination accounts by email. + continue + case "organizations": + owner = `id = $1` + default: + dep, ok := TableByName[fk.RefTable] + if !ok { + // Instance-wide data, or a reference referencePlan clears. + continue + } + owner = dep.OwnerScope() + } + out = append(out, repository.TenantReference{ForeignKey: fk, RefOwner: owner}) + } + return out, nil +} + // mergeOrganization applies the archive's workspace settings onto the // destination org. Identity, ownership, and lifecycle columns are excluded: // an archive must not be able to hand a workspace to someone else or schedule diff --git a/internal/app/orgtransfer/import_tenant_live_test.go b/internal/app/orgtransfer/import_tenant_live_test.go new file mode 100644 index 000000000..2df3e2626 --- /dev/null +++ b/internal/app/orgtransfer/import_tenant_live_test.go @@ -0,0 +1,218 @@ +package orgtransfer + +import ( + "archive/zip" + "bytes" + "context" + "encoding/json" + "errors" + "os" + "testing" + "time" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/warmbly/warmbly/internal/infrastructure/db" + "github.com/warmbly/warmbly/internal/models" + "github.com/warmbly/warmbly/internal/repository" +) + +// An import only writes rows of the destination workspace, and every reference +// it writes resolves to a row of that workspace. +// +// WARMBLY_TEST_DB=postgres://warmbly:warmbly@localhost:15432/?sslmode=disable \ +// go test ./internal/app/orgtransfer/ -run LiveImport -v + +type tenantFixture struct { + owner, org, contact, campaign uuid.UUID +} + +func liveImportService(t *testing.T) (Service, *pgxpool.Pool) { + t.Helper() + dsn := os.Getenv("WARMBLY_TEST_DB") + if dsn == "" { + t.Skip("WARMBLY_TEST_DB not set") + } + handle, err := db.New(context.Background(), dsn) + if err != nil { + t.Fatalf("connect: %v", err) + } + t.Cleanup(func() { handle.Pool.Close() }) + return NewService(repository.NewOrgTransferRepository(handle), nil, nil, nil, InstanceInfo{}), handle.Pool +} + +func newTenantFixture(t *testing.T, pool *pgxpool.Pool) tenantFixture { + t.Helper() + ctx := context.Background() + f := tenantFixture{owner: uuid.New(), org: uuid.New(), contact: uuid.New(), campaign: uuid.New()} + tag := f.org.String()[:8] + for _, q := range []struct { + sql string + args []any + }{ + {`INSERT INTO users (id, first_name, last_name, email, password_hash) VALUES ($1, 'Owner', 'Live', $2, 'x')`, + []any{f.owner, "xfer-" + tag + "@test.local"}}, + {`INSERT INTO organizations (id, name, slug, owner_user_id) VALUES ($1, 'Transfer', $2, $3)`, + []any{f.org, "xfer-" + tag, f.owner}}, + {`INSERT INTO organization_members (organization_id, user_id, role, accepted_at) VALUES ($1, $2, 'owner', NOW())`, + []any{f.org, f.owner}}, + {`INSERT INTO contacts (id, user_id, organization_id, email, first_name, last_name, company, phone, custom_fields) + VALUES ($1, $2, $3, $4, 'Kept', 'Here', '', '', '{}'::jsonb)`, + []any{f.contact, f.owner, f.org, "kept-" + tag + "@test.local"}}, + {`INSERT INTO campaigns (id, user_id, organization_id, name, description, days, updated_at, created_at) + VALUES ($1, $2, $3, 'Kept', '', 62, NOW(), NOW())`, + []any{f.campaign, f.owner, f.org}}, + } { + if _, err := pool.Exec(ctx, q.sql, q.args...); err != nil { + t.Fatalf("fixture: %v", err) + } + } + t.Cleanup(func() { + c := context.Background() + for _, q := range []string{ + `DELETE FROM campaign_leads WHERE campaign_id IN (SELECT id FROM campaigns WHERE organization_id = $1)`, + `DELETE FROM campaigns WHERE organization_id = $1`, + `DELETE FROM contacts WHERE organization_id = $1`, + `DELETE FROM organization_members WHERE organization_id = $1`, + `DELETE FROM organizations WHERE id = $1`, + } { + if _, err := pool.Exec(c, q, f.org); err != nil { + t.Errorf("cleanup: %v", err) + } + } + if _, err := pool.Exec(c, `DELETE FROM users WHERE id = $1`, f.owner); err != nil { + t.Errorf("cleanup: %v", err) + } + }) + return f +} + +// tenantArchive builds an archive holding the given rows per table. +func tenantArchive(t *testing.T, tables map[string][]map[string]any) *bytes.Reader { + t.Helper() + var buf bytes.Buffer + zw := zip.NewWriter(&buf) + m := Manifest{ + Kind: ArchiveKind, + FormatVersion: models.OrgTransferFormatVersion, + OrganizationID: uuid.New(), + ExportedAt: time.Now(), + } + for _, tbl := range Tables { + rows, ok := tables[tbl.Name] + if !ok { + continue + } + w, err := zw.Create(dataPath(tbl.Name)) + if err != nil { + t.Fatal(err) + } + cols := map[string]bool{} + for _, r := range rows { + line, err := json.Marshal(r) + if err != nil { + t.Fatal(err) + } + if _, err := w.Write(append(line, '\n')); err != nil { + t.Fatal(err) + } + for c := range r { + cols[c] = true + } + } + mt := ManifestTable{Name: tbl.Name, Group: tbl.Group, Rows: int64(len(rows))} + for c := range cols { + mt.Columns = append(mt.Columns, c) + } + m.Tables = append(m.Tables, mt) + } + w, err := zw.Create(manifestPath) + if err != nil { + t.Fatal(err) + } + if err := json.NewEncoder(w).Encode(m); err != nil { + t.Fatal(err) + } + if err := zw.Close(); err != nil { + t.Fatal(err) + } + return bytes.NewReader(buf.Bytes()) +} + +func contactRow(id, org uuid.UUID, email string) map[string]any { + return map[string]any{ + "id": id, "user_id": uuid.New(), "organization_id": org, "email": email, + "first_name": "Archive", "last_name": "Row", "company": "", "phone": "", "custom_fields": map[string]any{}, + } +} + +func TestLiveImportRefusesRowsOwnedByAnotherWorkspace(t *testing.T) { + svc, pool := liveImportService(t) + ctx := context.Background() + other := newTenantFixture(t, pool) + dest := newTenantFixture(t, pool) + + for _, conflict := range []models.OrgImportConflict{models.OrgImportConflictOverwrite, models.OrgImportConflictSkip} { + archive := tenantArchive(t, map[string][]map[string]any{ + "contacts": {contactRow(other.contact, other.org, "moved@test.local")}, + }) + _, err := svc.ImportFrom(ctx, dest.org, archive, ImportOptions{Conflict: conflict, ActorUserID: dest.owner}, nil) + if !errors.Is(err, ErrForeignRows) { + t.Fatalf("%s: want ErrForeignRows, got %v", conflict, err) + } + var org uuid.UUID + var email string + if err := pool.QueryRow(ctx, `SELECT organization_id, email FROM contacts WHERE id = $1`, other.contact).Scan(&org, &email); err != nil { + t.Fatal(err) + } + if org != other.org || email == "moved@test.local" { + t.Fatalf("%s: another workspace's contact was changed: org %s email %s", conflict, org, email) + } + } +} + +func TestLiveImportRefusesReferencesIntoAnotherWorkspace(t *testing.T) { + svc, pool := liveImportService(t) + ctx := context.Background() + other := newTenantFixture(t, pool) + dest := newTenantFixture(t, pool) + + fresh := uuid.New() + archive := tenantArchive(t, map[string][]map[string]any{ + "contacts": {contactRow(fresh, uuid.New(), "fresh-"+fresh.String()[:8]+"@test.local")}, + "campaign_leads": {{"campaign_id": other.campaign, "contact_id": fresh, "source": "manual"}}, + }) + _, err := svc.ImportFrom(ctx, dest.org, archive, ImportOptions{Conflict: models.OrgImportConflictSkip, ActorUserID: dest.owner}, nil) + if !errors.Is(err, ErrForeignRows) { + t.Fatalf("want ErrForeignRows, got %v", err) + } + var n int + if err := pool.QueryRow(ctx, `SELECT count(*) FROM campaign_leads WHERE campaign_id = $1`, other.campaign).Scan(&n); err != nil { + t.Fatal(err) + } + if n != 0 { + t.Fatalf("a lead landed on another workspace's campaign") + } +} + +func TestLiveImportOverwritesItsOwnRows(t *testing.T) { + svc, pool := liveImportService(t) + ctx := context.Background() + dest := newTenantFixture(t, pool) + + archive := tenantArchive(t, map[string][]map[string]any{ + "contacts": {contactRow(dest.contact, uuid.New(), "rerun-"+dest.contact.String()[:8]+"@test.local")}, + "campaign_leads": {{"campaign_id": dest.campaign, "contact_id": dest.contact, "source": "manual"}}, + }) + if _, err := svc.ImportFrom(ctx, dest.org, archive, ImportOptions{Conflict: models.OrgImportConflictOverwrite, ActorUserID: dest.owner}, nil); err != nil { + t.Fatalf("re-running an import into its own workspace: %v", err) + } + var email string + if err := pool.QueryRow(ctx, `SELECT email FROM contacts WHERE id = $1 AND organization_id = $2`, dest.contact, dest.org).Scan(&email); err != nil { + t.Fatal(err) + } + if email != "rerun-"+dest.contact.String()[:8]+"@test.local" { + t.Fatalf("own row not overwritten: %s", email) + } +} diff --git a/internal/app/orgtransfer/jobs.go b/internal/app/orgtransfer/jobs.go index 4e9e30c26..a053a3726 100644 --- a/internal/app/orgtransfer/jobs.go +++ b/internal/app/orgtransfer/jobs.go @@ -264,6 +264,7 @@ func (s *service) Preflight( } // What already exists here, and what this instance cannot take. + foreignTables := 0 for _, mt := range manifest.Tables { t, known := TableByName[mt.Name] if !known || t.ImportSkip { @@ -289,7 +290,7 @@ func (s *service) Preflight( if !ok { continue } - n, err := s.countConflicts(ctx, mt.Name, pk, entry) + n, foreign, err := s.countConflicts(ctx, orgID, t, pk, entry) if err != nil { // A conflict count is advisory; failing the whole preflight over // one unreadable table helps nobody. @@ -298,8 +299,15 @@ func (s *service) Preflight( if n > 0 { out.Conflicts[mt.Name] = n } + if foreign > 0 { + foreignTables++ + } } + if foreignTables > 0 { + out.Warnings = append(out.Warnings, "Some records in this archive already belong to another workspace on this instance, so the import will be refused. "+ + "Import it into the workspace it was exported from, or into a workspace on another instance.") + } if len(out.SkippedTables) > 0 { out.Warnings = append(out.Warnings, fmt.Sprintf( "%d table(s) in this archive do not exist on this instance and will be skipped. It was probably exported from a newer release.", @@ -313,18 +321,27 @@ func (s *service) Preflight( // scanning a million-row inbox table twice for a preview is not worth it. const conflictSampleRows = 2000 -func (s *service) countConflicts(ctx context.Context, table string, pk []string, entry *zip.File) (int64, error) { +// The second count is the sampled keys another workspace already holds. +func (s *service) countConflicts(ctx context.Context, orgID uuid.UUID, t *Table, pk []string, entry *zip.File) (int64, int64, error) { rc, err := entry.Open() if err != nil { - return 0, err + return 0, 0, err } defer rc.Close() rows, err := readRows(rc, conflictSampleRows) if err != nil { - return 0, err + return 0, 0, err } - return s.repo.CountExisting(ctx, table, pk, rows) + existing, err := s.repo.CountExisting(ctx, t.Name, pk, rows) + if err != nil { + return 0, 0, err + } + foreign, err := s.repo.CountForeignRows(ctx, nil, t.Name, t.OwnerScope(), pk, nil, orgID, rows) + if err != nil { + return 0, 0, err + } + return existing, foreign, nil } func (s *service) RequestImport( diff --git a/internal/app/orgtransfer/spec.go b/internal/app/orgtransfer/spec.go index 3bd265ce1..e06b11cf6 100644 --- a/internal/app/orgtransfer/spec.go +++ b/internal/app/orgtransfer/spec.go @@ -62,6 +62,10 @@ type Table struct { // organization. $1 is the organization id. Scope string + // Owner selects every row the organization owns, when that is wider than + // what Scope exports. Empty means Scope. $1 is the organization id. + Owner string + // Secrets are columns holding ciphertext that must be re-keyed. Secrets []SecretColumn @@ -83,6 +87,15 @@ type Table struct { Note string } +// OwnerScope is the WHERE fragment deciding whether an existing row belongs to +// the organization, which an import checks before writing or referencing it. +func (t *Table) OwnerScope() string { + if t.Owner != "" { + return t.Owner + } + return t.Scope +} + // Scope fragments. Written as subqueries rather than joins so every scope is a // plain WHERE clause and the reader can stay a single generic SELECT. const ( @@ -169,6 +182,7 @@ var Tables = []Table{ Name: "domain_redirects", Group: models.OrgDataGroupCore, // A row Cloud serves for a linked instance belongs to that link, which does not travel. Scope: `organization_id = $1 AND linked_instance_id IS NULL`, + Owner: scopeOrg, // DNS points at the source (or at Cloud for it) until moved, so the destination serves it itself once its own check passes. ResetOnImport: []string{"verified", "verified_at", "last_checked_at", "last_error", "served_by", "remote_host", "remote_records", "linked_instance_id", "reach_status", "reach_hint", "reach_detail", "reach_proxy", "reach_checked_at"}, @@ -700,6 +714,7 @@ var Tables = []Table{ { Name: "inbox_tag_results", Group: models.OrgDataGroupInbox, Scope: scopeOrg + ` AND status = 'complete'`, + Owner: scopeOrg, Note: "Automatic tagging verdicts, including the raw probabilities. They travel because retuning the weights " + "against stored answers is free while re-running the model over the history is not. Below email_accounts, " + "which it references.", @@ -728,6 +743,7 @@ var Tables = []Table{ // A placement probe's task stays behind: a pending one would send from // the destination to the source instance's seeds. Scope: `email_account_id IN ` + orgMailboxes + ` AND task_type <> 'placement'`, + Owner: `email_account_id IN ` + orgMailboxes, // The handle belongs to the source instance's queue. ResetOnImport: []string{"cloud_task_name"}, }, diff --git a/internal/repository/pg_orgtransfer.go b/internal/repository/pg_orgtransfer.go index 496258c25..279c50387 100644 --- a/internal/repository/pg_orgtransfer.go +++ b/internal/repository/pg_orgtransfer.go @@ -5,6 +5,8 @@ import ( "encoding/json" "errors" "fmt" + "regexp" + "strconv" "strings" "time" @@ -40,6 +42,9 @@ type OrgTransferRepository interface { // part of this import, and asking the catalog means a constraint added // later is covered without a code change. ForeignKeyRefs(ctx context.Context, table string) (map[string]string, error) + // ForeignKeys lists the table's foreign keys with their column pairs, so a + // composite key is checked as one reference. + ForeignKeys(ctx context.Context, table string) ([]ForeignKey, error) // ---- export ---- @@ -54,8 +59,14 @@ type OrgTransferRepository interface { // CountExisting reports how many of the supplied primary-key values are // already present, for the preflight report. CountExisting(ctx context.Context, table string, pk []string, rows []json.RawMessage) (int64, error) - // InsertBatch writes rows into table, honouring the conflict strategy. - InsertBatch(ctx context.Context, tx pgx.Tx, table string, cols []string, rows []json.RawMessage, conflict models.OrgImportConflict, pk []string) (int64, error) + // CountForeignRows counts rows whose primary key is already held by a row + // outside owner, or whose reference resolves to a row outside its + // target's owner. Every owner is a WHERE fragment with $1 the organization + // id. tx may be nil outside an import. + CountForeignRows(ctx context.Context, tx pgx.Tx, table, owner string, pk []string, refs []TenantReference, orgID uuid.UUID, rows []json.RawMessage) (int64, error) + // InsertBatch writes rows into table, honouring the conflict strategy. An + // overwrite only updates an existing row that owner selects for orgID. + InsertBatch(ctx context.Context, tx pgx.Tx, table string, cols []string, rows []json.RawMessage, conflict models.OrgImportConflict, pk []string, owner string, orgID uuid.UUID) (int64, error) // MergeOrganization applies the archive's organization row onto an // existing workspace, restricted to cols. MergeOrganization(ctx context.Context, tx pgx.Tx, orgID uuid.UUID, cols []string, row json.RawMessage) error @@ -110,6 +121,20 @@ type ColumnInfo struct { // Writable reports whether the column accepts a supplied value. func (c ColumnInfo) Writable() bool { return !c.IsGenerated && !c.IsIdentity } +// ForeignKey is one declared foreign key; Columns[i] references RefColumns[i]. +type ForeignKey struct { + Columns []string + RefTable string + RefColumns []string +} + +// TenantReference is a foreign key whose target rows belong to an +// organization, with the WHERE fragment ($1 the organization id) selecting them. +type TenantReference struct { + ForeignKey + RefOwner string +} + type orgTransferRepository struct { DB *db.DB } @@ -227,6 +252,40 @@ func (r *orgTransferRepository) ForeignKeyRefs(ctx context.Context, table string return out, rows.Err() } +func (r *orgTransferRepository) ForeignKeys(ctx context.Context, table string) ([]ForeignKey, error) { + rows, err := r.DB.Query(ctx, ` + SELECT rc.relname, + ARRAY(SELECT a.attname::text FROM unnest(con.conkey) WITH ORDINALITY k(num, ord) + JOIN pg_attribute a ON a.attrelid = con.conrelid AND a.attnum = k.num + ORDER BY k.ord), + ARRAY(SELECT a.attname::text FROM unnest(con.confkey) WITH ORDINALITY k(num, ord) + JOIN pg_attribute a ON a.attrelid = con.confrelid AND a.attnum = k.num + ORDER BY k.ord) + FROM pg_constraint con + JOIN pg_class c ON c.oid = con.conrelid + JOIN pg_namespace n ON n.oid = c.relnamespace + JOIN pg_class rc ON rc.oid = con.confrelid + WHERE con.contype = 'f' + AND n.nspname = 'public' + AND c.relname = $1 + ORDER BY con.conname + `, table) + if err != nil { + return nil, err + } + defer rows.Close() + + var out []ForeignKey + for rows.Next() { + var fk ForeignKey + if err := rows.Scan(&fk.RefTable, &fk.Columns, &fk.RefColumns); err != nil { + return nil, err + } + out = append(out, fk) + } + return out, rows.Err() +} + // ---------- export ---------- // StreamScoped reads a table with to_jsonb so every column type is rendered by @@ -328,6 +387,72 @@ func (r *orgTransferRepository) CountExisting(ctx context.Context, table string, return n, nil } +func (r *orgTransferRepository) CountForeignRows( + ctx context.Context, + tx pgx.Tx, + table, owner string, + pk []string, + refs []TenantReference, + orgID uuid.UUID, + rows []json.RawMessage, +) (int64, error) { + if len(rows) == 0 { + return 0, nil + } + // Each owner names its table's columns unqualified, so it sits alone in a + // subquery's FROM and resolves against that row, never the archive row. + var checks []string + if len(pk) > 0 && owner != "" { + checks = append(checks, `EXISTS (SELECT 1 FROM public.`+quoteIdent(table)+` AS existing + WHERE `+pairClause("existing", pk, "archive_row", pk)+` AND (`+owner+`) IS NOT TRUE)`) + } + for _, ref := range refs { + if len(ref.Columns) == 0 || len(ref.Columns) != len(ref.RefColumns) || ref.RefOwner == "" { + continue + } + checks = append(checks, `EXISTS (SELECT 1 FROM public.`+quoteIdent(ref.RefTable)+` AS target + WHERE `+pairClause("target", ref.RefColumns, "archive_row", ref.Columns)+` AND (`+ref.RefOwner+`) IS NOT TRUE)`) + } + if len(checks) == 0 { + return 0, nil + } + payload, err := json.Marshal(rows) + if err != nil { + return 0, err + } + q := `SELECT count(*) + FROM jsonb_populate_recordset(NULL::public.` + quoteIdent(table) + `, $2::jsonb) AS archive_row + WHERE ` + strings.Join(checks, ` OR `) + var row pgx.Row + if tx != nil { + row = tx.QueryRow(ctx, q, orgID, payload) + } else { + row = r.DB.QueryRow(ctx, q, orgID, payload) + } + var n int64 + if err := row.Scan(&n); err != nil { + return 0, fmt.Errorf("check ownership of %s: %w", table, err) + } + return n, nil +} + +// pairClause builds `l.a = r.x AND ...`; a NULL on either side matches nothing. +func pairClause(left string, leftCols []string, right string, rightCols []string) string { + parts := make([]string, len(leftCols)) + for i := range leftCols { + parts[i] = left + "." + quoteIdent(leftCols[i]) + " = " + right + "." + quoteIdent(rightCols[i]) + } + return strings.Join(parts, " AND ") +} + +// orgParam is an owner fragment's organization placeholder, never the start of $10. +var orgParam = regexp.MustCompile(`\$1\b`) + +// bindOrgParam renumbers an owner fragment's $1 to $n. +func bindOrgParam(fragment string, n int) string { + return orgParam.ReplaceAllLiteralString(fragment, "$"+strconv.Itoa(n)) +} + // matchClause builds `e.a IS NOT DISTINCT FROM c.a AND ...` for a key. func matchClause(pk []string) string { parts := make([]string, len(pk)) @@ -345,6 +470,8 @@ func (r *orgTransferRepository) InsertBatch( rows []json.RawMessage, conflict models.OrgImportConflict, pk []string, + owner string, + orgID uuid.UUID, ) (int64, error) { if len(rows) == 0 || len(cols) == 0 { return 0, nil @@ -354,8 +481,9 @@ func (r *orgTransferRepository) InsertBatch( return 0, err } + args := []any{payload} quoted := quoteIdents(cols) - q := `INSERT INTO public.` + quoteIdent(table) + ` (` + quoted + `) + q := `INSERT INTO public.` + quoteIdent(table) + ` AS dest (` + quoted + `) SELECT ` + quoted + ` FROM jsonb_populate_recordset(NULL::public.` + quoteIdent(table) + `, $1::jsonb)` @@ -377,6 +505,12 @@ func (r *orgTransferRepository) InsertBatch( q += ` ON CONFLICT DO NOTHING` } else { q += ` ON CONFLICT (` + quoteIdents(pk) + `) DO UPDATE SET ` + strings.Join(sets, ", ") + if owner != "" { + // The owner sits in its own FROM, since here an unqualified column is ambiguous with EXCLUDED. + q += ` WHERE EXISTS (SELECT 1 FROM public.` + quoteIdent(table) + ` AS owned + WHERE ` + pairClause("owned", pk, "dest", pk) + ` AND (` + bindOrgParam(owner, 2) + `))` + args = append(args, orgID) + } } default: // Untargeted, so it covers every unique constraint on the table, not @@ -385,7 +519,7 @@ func (r *orgTransferRepository) InsertBatch( q += ` ON CONFLICT DO NOTHING` } - tag, err := tx.Exec(ctx, q, payload) + tag, err := tx.Exec(ctx, q, args...) if err != nil { return 0, fmt.Errorf("insert into %s: %w", table, err) } diff --git a/internal/repository/pg_orgtransfer_test.go b/internal/repository/pg_orgtransfer_test.go new file mode 100644 index 000000000..de8483123 --- /dev/null +++ b/internal/repository/pg_orgtransfer_test.go @@ -0,0 +1,17 @@ +package repository + +import "testing" + +func TestBindOrgParamRenumbersOnlyTheOrganizationPlaceholder(t *testing.T) { + cases := map[string]string{ + `organization_id = $1`: `organization_id = $2`, + `email_id IN (SELECT id FROM email_accounts WHERE organization_id = $1)`: `email_id IN (SELECT id FROM email_accounts WHERE organization_id = $2)`, + `org_id = $1 AND x = $10`: `org_id = $2 AND x = $10`, + `a = $1 OR b = $1`: `a = $2 OR b = $2`, + } + for in, want := range cases { + if got := bindOrgParam(in, 2); got != want { + t.Errorf("bindOrgParam(%q) = %q, want %q", in, got, want) + } + } +}