mirror of
https://github.com/warmbly/warmbly.git
synced 2026-10-05 00:02:12 +00:00
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:
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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(
|
||||
|
||||
@@ -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"},
|
||||
},
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user