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

This commit is contained in:
Matthew Meszaros
2026-10-04 03:03:34 -07:00
parent 43a9d004f2
commit eff2c14a9d
7 changed files with 480 additions and 20 deletions
@@ -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
+64 -10
View File
@@ -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
@@ -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/<db>?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)
}
}
+22 -5
View File
@@ -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(
+16
View File
@@ -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"},
},
+138 -4
View File
@@ -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)
}
@@ -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)
}
}
}